187 lines
5.3 KiB
Go
187 lines
5.3 KiB
Go
|
|
package queue
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"database/sql"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"floorvisualizer/internal/logger"
|
||
|
|
"os"
|
||
|
|
"path/filepath"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"github.com/redis/go-redis/v9"
|
||
|
|
|
||
|
|
"floorvisualizer/internal/handler"
|
||
|
|
"floorvisualizer/internal/openrouter"
|
||
|
|
)
|
||
|
|
|
||
|
|
// StartWorkers launches n goroutines that consume jobs from the Redis queue.
|
||
|
|
func StartWorkers(ctx context.Context, n int, rdb *redis.Client, client *openrouter.Client, db *sql.DB) {
|
||
|
|
for i := 0; i < n; i++ {
|
||
|
|
go func(workerID int) {
|
||
|
|
logger.Info("[worker %d] started", workerID)
|
||
|
|
for {
|
||
|
|
select {
|
||
|
|
case <-ctx.Done():
|
||
|
|
logger.Info("[worker %d] shutting down", workerID)
|
||
|
|
return
|
||
|
|
default:
|
||
|
|
}
|
||
|
|
jobID, err := DequeueJob(ctx, rdb)
|
||
|
|
if err != nil {
|
||
|
|
if ctx.Err() != nil {
|
||
|
|
return
|
||
|
|
}
|
||
|
|
logger.Error("[worker %d] dequeue error: %v", workerID, err)
|
||
|
|
time.Sleep(time.Second)
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
logger.Info("[worker %d] picked job %s", workerID, jobID)
|
||
|
|
processJob(ctx, rdb, client, db, jobID)
|
||
|
|
}
|
||
|
|
}(i + 1)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func processJob(ctx context.Context, rdb *redis.Client, client *openrouter.Client, db *sql.DB, jobID string) {
|
||
|
|
SetJobStatus(ctx, rdb, jobID, "processing")
|
||
|
|
|
||
|
|
// Load payload
|
||
|
|
payload, err := GetJobPayload(ctx, rdb, jobID)
|
||
|
|
if err != nil {
|
||
|
|
SetJobError(ctx, rdb, jobID, "Failed to load job payload")
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
inputPath := payload["input_path"]
|
||
|
|
floorID := payload["floor_id"]
|
||
|
|
patternCode := payload["pattern_code"]
|
||
|
|
roomCode := payload["room_code"]
|
||
|
|
|
||
|
|
// Look up floor and pattern from DB (same logic as before)
|
||
|
|
floor := handler.FindFloorBySKUPublic(db, floorID)
|
||
|
|
pattern := handler.FindPatternByCodePublic(patternCode)
|
||
|
|
room := handler.FindRoomByCodePublic(roomCode)
|
||
|
|
|
||
|
|
if floor == nil || pattern == nil {
|
||
|
|
SetJobError(ctx, rdb, jobID, "Invalid floor or pattern")
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
maskPath := filepath.Join("outputs", fmt.Sprintf("mask_%s.png", jobID))
|
||
|
|
outputPath := filepath.Join("outputs", fmt.Sprintf("result_%s.png", jobID))
|
||
|
|
|
||
|
|
jobCtx, cancel := context.WithTimeout(ctx, 180*time.Second)
|
||
|
|
defer cancel()
|
||
|
|
|
||
|
|
// Step 0: content check
|
||
|
|
SetJobProgress(jobCtx, rdb, jobID, "check", "Checking content...")
|
||
|
|
logger.Info("[%s] Content check", jobID)
|
||
|
|
if passed, _ := validateContent(client, jobCtx, inputPath); !passed {
|
||
|
|
SetJobError(jobCtx, rdb, jobID, "Content check failed: Please upload an indoor room photo")
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 1: mask
|
||
|
|
SetJobProgress(jobCtx, rdb, jobID, "mask", "Generating floor mask...")
|
||
|
|
logger.Info("[%s] Mask generation", jobID)
|
||
|
|
maskPrompt := handler.BuildMaskPromptPublic(room)
|
||
|
|
if _, err := client.RecognizeImage(jobCtx, openrouter.ImageRecognitionRequest{
|
||
|
|
Prompt: maskPrompt, InputImagePath: inputPath, OutputPath: maskPath,
|
||
|
|
}); err != nil {
|
||
|
|
SetJobError(jobCtx, rdb, jobID, "Mask generation failed: "+err.Error())
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Step 2: inpaint
|
||
|
|
roomLabel := ""
|
||
|
|
if room != nil {
|
||
|
|
roomLabel = " · " + room.Name
|
||
|
|
}
|
||
|
|
SetJobProgress(jobCtx, rdb, jobID, "inpaint", fmt.Sprintf("Applying %s %s%s...", floor.Name, pattern.Name, roomLabel))
|
||
|
|
logger.Info("[%s] Inpainting: %s + %s", jobID, floor.Name, pattern.Name)
|
||
|
|
|
||
|
|
floorPrompt := handler.BuildFloorPromptPublic(*floor, *pattern, room)
|
||
|
|
result, err := client.GenerateImageToFile(jobCtx, openrouter.ImageGenerationRequest{
|
||
|
|
Prompt: floorPrompt, OutputPath: outputPath, InputImagePath: inputPath,
|
||
|
|
MaskImagePath: maskPath, ImageSize: "2K",
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
SetJobError(jobCtx, rdb, jobID, "Floor generation failed: "+err.Error())
|
||
|
|
return
|
||
|
|
}
|
||
|
|
|
||
|
|
// Clean up intermediate files
|
||
|
|
os.Remove(inputPath)
|
||
|
|
os.Remove(maskPath)
|
||
|
|
|
||
|
|
_ = result
|
||
|
|
SetJobResult(jobCtx, rdb, jobID, "/"+filepath.ToSlash(outputPath))
|
||
|
|
logger.Info("[%s] Done", jobID)
|
||
|
|
}
|
||
|
|
|
||
|
|
// validateContent checks if the image is an indoor room photo.
|
||
|
|
func validateContent(client *openrouter.Client, ctx context.Context, imagePath string) (bool, int) {
|
||
|
|
imgData, err := os.ReadFile(imagePath)
|
||
|
|
if err != nil {
|
||
|
|
return false, 0
|
||
|
|
}
|
||
|
|
dataURI := "data:image/png;base64," + base64enc(imgData)
|
||
|
|
resp, err := client.Chat(ctx, openrouter.ChatRequest{
|
||
|
|
Model: "google/gemini-2.5-flash",
|
||
|
|
Messages: []openrouter.Message{{Role: "user", Content: []openrouter.ContentPart{
|
||
|
|
{Type: "image_url", ImageURL: &openrouter.ImageURL{URL: dataURI}},
|
||
|
|
{Type: "text", Text: `Score this image 1-10 for indoor floor replacement suitability.
|
||
|
|
8-10: Clear indoor room, floor visible. 5-7: Indoor but poor angle. 1-4: Outdoor/pets/objects/no floor.
|
||
|
|
Return ONLY a number (1-10).`},
|
||
|
|
}}},
|
||
|
|
})
|
||
|
|
if err != nil || len(resp.Choices) == 0 {
|
||
|
|
return true, 0
|
||
|
|
}
|
||
|
|
score := parseScore(resp.Choices[0].Message.Content)
|
||
|
|
return score >= 6, score
|
||
|
|
}
|
||
|
|
|
||
|
|
func parseScore(raw string) int {
|
||
|
|
for _, s := range []string{"10", "9", "8", "7", "6", "5", "4", "3", "2", "1"} {
|
||
|
|
if len(raw) >= len(s) && raw[:len(s)] == s {
|
||
|
|
v := 0
|
||
|
|
fmt.Sscanf(s, "%d", &v)
|
||
|
|
return v
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return 0
|
||
|
|
}
|
||
|
|
|
||
|
|
func base64enc(data []byte) string {
|
||
|
|
const tbl = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"
|
||
|
|
var b []byte
|
||
|
|
for i := 0; i < len(data); i += 3 {
|
||
|
|
b0, b1, b2 := data[i], byte(0), byte(0)
|
||
|
|
if i+1 < len(data) {
|
||
|
|
b1 = data[i+1]
|
||
|
|
}
|
||
|
|
if i+2 < len(data) {
|
||
|
|
b2 = data[i+2]
|
||
|
|
}
|
||
|
|
b = append(b, tbl[b0>>2])
|
||
|
|
b = append(b, tbl[((b0&3)<<4)|(b1>>4)])
|
||
|
|
if i+1 < len(data) {
|
||
|
|
b = append(b, tbl[((b1&15)<<2)|(b2>>6)])
|
||
|
|
} else {
|
||
|
|
b = append(b, '=')
|
||
|
|
}
|
||
|
|
if i+2 < len(data) {
|
||
|
|
b = append(b, tbl[b2&63])
|
||
|
|
} else {
|
||
|
|
b = append(b, '=')
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return string(b)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Make sure io is used (for any future imports)
|
||
|
|
var _ = io.Discard
|