FloorVisualizer/internal/queue/worker.go

184 lines
5.2 KiB
Go
Raw Normal View History

package queue
import (
"context"
"database/sql"
"fmt"
"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)
}