From 397803fb96165a72e9ba3322419ff60fe27619e1 Mon Sep 17 00:00:00 2001 From: m Date: Wed, 29 Jul 2026 18:56:26 +0200 Subject: [PATCH] sse + ocr replanish --- backend/cmd/server/main.go | 10 +- backend/internal/config/config.go | 18 ++- .../db/migrations/014_ocr_jobs.down.sql | 1 + .../db/migrations/014_ocr_jobs.up.sql | 12 ++ backend/internal/db/models.go | 10 ++ backend/internal/db/ocr_jobs.sql.go | 122 ++++++++++++++++++ backend/internal/db/queries/ocr_jobs.sql | 22 ++++ backend/internal/handler/events.go | 34 +++++ backend/internal/handler/handler.go | 3 + backend/internal/handler/resource.go | 8 +- backend/internal/handler/router.go | 3 + backend/internal/service/conversion.go | 20 ++- backend/internal/service/events.go | 54 ++++++++ backend/internal/service/ocr.go | 95 +++++++++++++- 14 files changed, 386 insertions(+), 26 deletions(-) create mode 100644 backend/internal/db/migrations/014_ocr_jobs.down.sql create mode 100644 backend/internal/db/migrations/014_ocr_jobs.up.sql create mode 100644 backend/internal/db/ocr_jobs.sql.go create mode 100644 backend/internal/db/queries/ocr_jobs.sql create mode 100644 backend/internal/handler/events.go create mode 100644 backend/internal/service/events.go diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index 70df145..8826d55 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -30,8 +30,10 @@ func main() { queries := db.New(database) + eventBroker := service.NewEventBroker() + resourceSvc := service.NewResourceService(database, queries, cfg) - ocrSvc := service.NewOCRService(cfg, resourceSvc) + ocrSvc := service.NewOCRService(database, queries, cfg, resourceSvc, eventBroker) conversionSvc := service.NewConversionService(queries, cfg) urlSvc := service.NewURLService(cfg.HMACSecret, cfg.ServerHost, cfg.URLExpiryMinutes) authSvc, err := auth.NewAuthService(database, queries, cfg) @@ -43,13 +45,13 @@ func main() { placementSvc := service.NewPlacementService(queries) syncSvc := service.NewSyncService(queries) - ocrSvc.Start() + ocrSvc.Start(cfg.OCRWorkers) defer ocrSvc.Stop() - conversionSvc.Start() + conversionSvc.Start(cfg.ConversionWorkers) defer conversionSvc.Stop() - h := handler.New(resourceSvc, ocrSvc, urlSvc, auth.NewAuthHandler(authSvc), conversionSvc, rebacSvc, placementSvc, syncSvc) + h := handler.New(resourceSvc, ocrSvc, urlSvc, auth.NewAuthHandler(authSvc), conversionSvc, rebacSvc, placementSvc, syncSvc, eventBroker) r := gin.Default() handler.SetupRoutes(r, h, authSvc) diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index ef22da9..01fd802 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -11,12 +11,14 @@ type Config struct { OCREndpoint string UploadDir string HMACSecret string - ServerHost string - PASETOKey string - LibreOfficePath string - PdftoppmPath string - ThumbnailDir string - URLExpiryMinutes int + ServerHost string + PASETOKey string + LibreOfficePath string + PdftoppmPath string + ThumbnailDir string + URLExpiryMinutes int + OCRWorkers int + ConversionWorkers int } func Load() *Config { @@ -31,7 +33,9 @@ func Load() *Config { LibreOfficePath: envOr("LIBREOFFICE_PATH", "/usr/bin/libreoffice"), PdftoppmPath: envOr("PDFTOPPM_PATH", "/usr/bin/pdftoppm"), ThumbnailDir: envOr("THUMBNAIL_DIR", "./uploads/thumbnails"), - URLExpiryMinutes: envOrInt("URL_EXPIRY_MINUTES", 60), + URLExpiryMinutes: envOrInt("URL_EXPIRY_MINUTES", 60), + OCRWorkers: envOrInt("OCR_WORKERS", 1), + ConversionWorkers: envOrInt("CONVERSION_WORKERS", 1), } } diff --git a/backend/internal/db/migrations/014_ocr_jobs.down.sql b/backend/internal/db/migrations/014_ocr_jobs.down.sql new file mode 100644 index 0000000..a1f7ec1 --- /dev/null +++ b/backend/internal/db/migrations/014_ocr_jobs.down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS ocr_jobs; diff --git a/backend/internal/db/migrations/014_ocr_jobs.up.sql b/backend/internal/db/migrations/014_ocr_jobs.up.sql new file mode 100644 index 0000000..f960b7b --- /dev/null +++ b/backend/internal/db/migrations/014_ocr_jobs.up.sql @@ -0,0 +1,12 @@ +CREATE TABLE ocr_jobs ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + resource_id UUID NOT NULL REFERENCES resources(id) ON DELETE CASCADE, + file_path TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'processing', 'done', 'failed')), + error_message TEXT NOT NULL DEFAULT '', + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX idx_ocr_jobs_status ON ocr_jobs(status); +CREATE INDEX idx_ocr_jobs_resource ON ocr_jobs(resource_id); diff --git a/backend/internal/db/models.go b/backend/internal/db/models.go index 7e67fc4..c5e961e 100644 --- a/backend/internal/db/models.go +++ b/backend/internal/db/models.go @@ -12,6 +12,16 @@ import ( "github.com/google/uuid" ) +type OcrJob struct { + ID uuid.UUID `json:"id"` + ResourceID uuid.UUID `json:"resource_id"` + FilePath string `json:"file_path"` + Status string `json:"status"` + ErrorMessage string `json:"error_message"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + type RebacRelation struct { ID uuid.UUID `json:"id"` ResourceID uuid.UUID `json:"resource_id"` diff --git a/backend/internal/db/ocr_jobs.sql.go b/backend/internal/db/ocr_jobs.sql.go new file mode 100644 index 0000000..2086a2b --- /dev/null +++ b/backend/internal/db/ocr_jobs.sql.go @@ -0,0 +1,122 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: ocr_jobs.sql + +package db + +import ( + "context" + + "github.com/google/uuid" +) + +const createOCRJob = `-- name: CreateOCRJob :one +INSERT INTO ocr_jobs (resource_id, file_path, status) +VALUES ($1, $2, 'pending') +RETURNING id, resource_id, file_path, status, error_message, created_at, updated_at +` + +type CreateOCRJobParams struct { + ResourceID uuid.UUID `json:"resource_id"` + FilePath string `json:"file_path"` +} + +func (q *Queries) CreateOCRJob(ctx context.Context, arg CreateOCRJobParams) (OcrJob, error) { + row := q.db.QueryRowContext(ctx, createOCRJob, arg.ResourceID, arg.FilePath) + var i OcrJob + err := row.Scan( + &i.ID, + &i.ResourceID, + &i.FilePath, + &i.Status, + &i.ErrorMessage, + &i.CreatedAt, + &i.UpdatedAt, + ) + return i, err +} + +const deleteOCRJob = `-- name: DeleteOCRJob :exec +DELETE FROM ocr_jobs +WHERE id = $1 +` + +func (q *Queries) DeleteOCRJob(ctx context.Context, id uuid.UUID) error { + _, err := q.db.ExecContext(ctx, deleteOCRJob, id) + return err +} + +const getOCRJob = `-- name: GetOCRJob :one +SELECT id, resource_id, file_path, status, error_message, created_at, updated_at FROM ocr_jobs +WHERE id = $1 LIMIT 1 +` + +func (q *Queries) GetOCRJob(ctx context.Context, id uuid.UUID) (OcrJob, error) { + row := q.db.QueryRowContext(ctx, getOCRJob, id) + var i OcrJob + err := row.Scan( + &i.ID, + &i.ResourceID, + &i.FilePath, + &i.Status, + &i.ErrorMessage, + &i.CreatedAt, + &i.UpdatedAt, + ) + return i, err +} + +const listPendingOCRJobs = `-- name: ListPendingOCRJobs :many +SELECT id, resource_id, file_path, status, error_message, created_at, updated_at FROM ocr_jobs +WHERE status = 'pending' +ORDER BY created_at ASC +` + +func (q *Queries) ListPendingOCRJobs(ctx context.Context) ([]OcrJob, error) { + rows, err := q.db.QueryContext(ctx, listPendingOCRJobs) + if err != nil { + return nil, err + } + defer rows.Close() + var items []OcrJob + for rows.Next() { + var i OcrJob + if err := rows.Scan( + &i.ID, + &i.ResourceID, + &i.FilePath, + &i.Status, + &i.ErrorMessage, + &i.CreatedAt, + &i.UpdatedAt, + ); err != nil { + return nil, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const updateOCRJobStatus = `-- name: UpdateOCRJobStatus :exec +UPDATE ocr_jobs +SET status = $1, error_message = $2, updated_at = CURRENT_TIMESTAMP +WHERE id = $3 +` + +type UpdateOCRJobStatusParams struct { + Status string `json:"status"` + ErrorMessage string `json:"error_message"` + ID uuid.UUID `json:"id"` +} + +func (q *Queries) UpdateOCRJobStatus(ctx context.Context, arg UpdateOCRJobStatusParams) error { + _, err := q.db.ExecContext(ctx, updateOCRJobStatus, arg.Status, arg.ErrorMessage, arg.ID) + return err +} diff --git a/backend/internal/db/queries/ocr_jobs.sql b/backend/internal/db/queries/ocr_jobs.sql new file mode 100644 index 0000000..2398f38 --- /dev/null +++ b/backend/internal/db/queries/ocr_jobs.sql @@ -0,0 +1,22 @@ +-- name: CreateOCRJob :one +INSERT INTO ocr_jobs (resource_id, file_path, status) +VALUES ($1, $2, 'pending') +RETURNING *; + +-- name: GetOCRJob :one +SELECT * FROM ocr_jobs +WHERE id = $1 LIMIT 1; + +-- name: ListPendingOCRJobs :many +SELECT * FROM ocr_jobs +WHERE status = 'pending' +ORDER BY created_at ASC; + +-- name: UpdateOCRJobStatus :exec +UPDATE ocr_jobs +SET status = $1, error_message = $2, updated_at = CURRENT_TIMESTAMP +WHERE id = $3; + +-- name: DeleteOCRJob :exec +DELETE FROM ocr_jobs +WHERE id = $1; diff --git a/backend/internal/handler/events.go b/backend/internal/handler/events.go new file mode 100644 index 0000000..f48b7b3 --- /dev/null +++ b/backend/internal/handler/events.go @@ -0,0 +1,34 @@ +package handler + +import ( + "fmt" + "io" + + "github.com/gin-gonic/gin" + "github.com/vaultdrop/backend/internal/service" +) + +type EventHandler struct { + broker *service.EventBroker +} + +func NewEventHandler(broker *service.EventBroker) *EventHandler { + return &EventHandler{broker: broker} +} + +func (h *EventHandler) Stream(c *gin.Context) { + c.Header("Content-Type", "text/event-stream") + c.Header("Cache-Control", "no-cache") + c.Header("Connection", "keep-alive") + + ch, _ := h.broker.Subscribe(c.Request.Context()) + + c.Stream(func(w io.Writer) bool { + event, ok := <-ch + if !ok { + return false + } + _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", event.Type, event.Data) + return err == nil + }) +} diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index 6f2a1c6..fbf4e0e 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -13,6 +13,7 @@ type Handler struct { Share *ShareHandler Device *DeviceHandler Sync *SyncHandler + Events *EventHandler } func New( @@ -24,6 +25,7 @@ func New( rebacSvc *service.RebacService, placementSvc *service.PlacementService, syncSvc *service.SyncService, + eventBroker *service.EventBroker, ) *Handler { return &Handler{ Resource: &ResourceHandler{ @@ -38,5 +40,6 @@ func New( Share: &ShareHandler{rebac: rebacSvc}, Device: &DeviceHandler{placement: placementSvc}, Sync: &SyncHandler{sync: syncSvc}, + Events: NewEventHandler(eventBroker), } } diff --git a/backend/internal/handler/resource.go b/backend/internal/handler/resource.go index 214ff6b..4af5cb1 100644 --- a/backend/internal/handler/resource.go +++ b/backend/internal/handler/resource.go @@ -43,10 +43,14 @@ func (h *ResourceHandler) Upload(c *gin.Context) { return } - h.ocr.Enqueue(result.ID, result.Path) + if err := h.ocr.Enqueue(result.ID, result.Path); err != nil { + log.Printf("WARN %v", err) + } if service.IsConvertible(result.MimeType) { - h.conversion.Enqueue(result.ID, result.Path, result.MimeType) + if err := h.conversion.Enqueue(result.ID, result.Path, result.MimeType); err != nil { + log.Printf("WARN %v", err) + } } results = append(results, gin.H{ diff --git a/backend/internal/handler/router.go b/backend/internal/handler/router.go index d6a2595..1531086 100644 --- a/backend/internal/handler/router.go +++ b/backend/internal/handler/router.go @@ -38,6 +38,9 @@ func SetupRoutes(r *gin.Engine, h *Handler, authMiddleware *auth.AuthService) { // Variants protected.GET("/resources/:id/variants", h.Resource.GetVariants) + // Events (SSE) + protected.GET("/events", h.Events.Stream) + // Dedup protected.POST("/resources/dedup-check", h.Resource.CheckDuplicates) diff --git a/backend/internal/service/conversion.go b/backend/internal/service/conversion.go index 27e0d20..17a8c37 100644 --- a/backend/internal/service/conversion.go +++ b/backend/internal/service/conversion.go @@ -36,9 +36,11 @@ func NewConversionService(queries *db.Queries, cfg *config.Config) *ConversionSe } } -func (s *ConversionService) Start() { - go s.worker() - log.Println("[Conversion] Worker started") +func (s *ConversionService) Start(workerCount int) { + for i := range workerCount { + go s.worker() + log.Printf("[Conversion] Worker %d started", i) + } } func (s *ConversionService) Stop() { @@ -46,9 +48,15 @@ func (s *ConversionService) Stop() { log.Println("[Conversion] Worker stopped") } -func (s *ConversionService) Enqueue(resourceID, filePath, mimeType string) { - s.jobs <- ConversionJob{ResourceID: resourceID, FilePath: filePath, MimeType: mimeType} - log.Printf("[Conversion] Enqueued resource %s", resourceID) +func (s *ConversionService) Enqueue(resourceID, filePath, mimeType string) error { + select { + case s.jobs <- ConversionJob{ResourceID: resourceID, FilePath: filePath, MimeType: mimeType}: + log.Printf("[Conversion] Enqueued resource %s", resourceID) + return nil + default: + log.Printf("[Conversion] Queue full, dropping resource %s", resourceID) + return fmt.Errorf("conversion queue full (%d pending)", len(s.jobs)) + } } func (s *ConversionService) worker() { diff --git a/backend/internal/service/events.go b/backend/internal/service/events.go new file mode 100644 index 0000000..8f676b3 --- /dev/null +++ b/backend/internal/service/events.go @@ -0,0 +1,54 @@ +package service + +import ( + "context" + "sync" + + "github.com/google/uuid" +) + +type Event struct { + Type string + Data string +} + +type EventBroker struct { + mu sync.RWMutex + subscribers map[string]chan Event +} + +func NewEventBroker() *EventBroker { + return &EventBroker{ + subscribers: make(map[string]chan Event), + } +} + +func (b *EventBroker) Subscribe(ctx context.Context) (<-chan Event, string) { + id := uuid.New().String() + ch := make(chan Event, 16) + + b.mu.Lock() + b.subscribers[id] = ch + b.mu.Unlock() + + go func() { + <-ctx.Done() + b.mu.Lock() + delete(b.subscribers, id) + close(ch) + b.mu.Unlock() + }() + + return ch, id +} + +func (b *EventBroker) Publish(eventType, data string) { + b.mu.RLock() + defer b.mu.RUnlock() + for _, ch := range b.subscribers { + select { + case ch <- Event{Type: eventType, Data: data}: + default: + } + } +} diff --git a/backend/internal/service/ocr.go b/backend/internal/service/ocr.go index bf7ac13..3994deb 100644 --- a/backend/internal/service/ocr.go +++ b/backend/internal/service/ocr.go @@ -1,15 +1,21 @@ package service import ( + "context" + "database/sql" + "fmt" "log" "os" "strings" + "github.com/google/uuid" "github.com/vaultdrop/backend/internal/config" + "github.com/vaultdrop/backend/internal/db" "github.com/vaultdrop/backend/internal/ocr" ) type OCRJob struct { + DBID uuid.UUID ResourceID string FilePath string } @@ -17,20 +23,27 @@ type OCRJob struct { type OCRService struct { client *ocr.Client resourceSvc *ResourceService + broker *EventBroker + queries *db.Queries jobs chan OCRJob } -func NewOCRService(cfg *config.Config, resourceSvc *ResourceService) *OCRService { +func NewOCRService(database *sql.DB, queries *db.Queries, cfg *config.Config, resourceSvc *ResourceService, broker *EventBroker) *OCRService { return &OCRService{ client: ocr.NewClient(cfg.OCREndpoint), resourceSvc: resourceSvc, + broker: broker, + queries: queries, jobs: make(chan OCRJob, 100), } } -func (s *OCRService) Start() { - go s.worker() - log.Println("[OCR] Worker started") +func (s *OCRService) Start(workerCount int) { + s.replenish() + for i := range workerCount { + go s.worker() + log.Printf("[OCR] Worker %d started", i) + } } func (s *OCRService) Stop() { @@ -38,9 +51,51 @@ func (s *OCRService) Stop() { log.Println("[OCR] Worker stopped") } -func (s *OCRService) Enqueue(resourceID, filePath string) { - s.jobs <- OCRJob{ResourceID: resourceID, FilePath: filePath} - log.Printf("[OCR] Enqueued resource %s", resourceID) +func (s *OCRService) Enqueue(resourceID, filePath string) error { + ctx := context.Background() + + resourceUUID, err := uuid.Parse(resourceID) + if err != nil { + return fmt.Errorf("parse resource id: %w", err) + } + + dbJob, err := s.queries.CreateOCRJob(ctx, db.CreateOCRJobParams{ + ResourceID: resourceUUID, + FilePath: filePath, + }) + if err != nil { + return fmt.Errorf("create ocr job: %w", err) + } + + job := OCRJob{DBID: dbJob.ID, ResourceID: resourceID, FilePath: filePath} + + select { + case s.jobs <- job: + log.Printf("[OCR] Enqueued resource %s (job %s)", resourceID, dbJob.ID) + return nil + default: + log.Printf("[OCR] Queue full, resource %s persisted as pending (job %s)", resourceID, dbJob.ID) + return nil + } +} + +func (s *OCRService) replenish() { + ctx := context.Background() + pending, err := s.queries.ListPendingOCRJobs(ctx) + if err != nil { + log.Printf("[OCR] Failed to load pending jobs: %v", err) + return + } + for _, j := range pending { + job := OCRJob{DBID: j.ID, ResourceID: j.ResourceID.String(), FilePath: j.FilePath} + select { + case s.jobs <- job: + log.Printf("[OCR] Replenished job %s (resource %s)", j.ID, j.ResourceID) + default: + log.Printf("[OCR] Queue full, leaving job %s in pending", j.ID) + return + } + } } func (s *OCRService) worker() { @@ -50,17 +105,33 @@ func (s *OCRService) worker() { } func (s *OCRService) process(job OCRJob) { + ctx := context.Background() log.Printf("[OCR] Processing resource %s", job.ResourceID) + s.queries.UpdateOCRJobStatus(ctx, db.UpdateOCRJobStatusParams{ + ID: job.DBID, + Status: "processing", + }) + data, err := os.ReadFile(job.FilePath) if err != nil { log.Printf("[OCR] Failed to read resource %s: %v", job.ResourceID, err) + s.queries.UpdateOCRJobStatus(ctx, db.UpdateOCRJobStatusParams{ + ID: job.DBID, + Status: "failed", + ErrorMessage: err.Error(), + }) return } blocks, err := s.client.Recognize(data) if err != nil { log.Printf("[OCR] Failed to recognize resource %s: %v", job.ResourceID, err) + s.queries.UpdateOCRJobStatus(ctx, db.UpdateOCRJobStatusParams{ + ID: job.DBID, + Status: "failed", + ErrorMessage: err.Error(), + }) return } @@ -68,9 +139,19 @@ func (s *OCRService) process(job OCRJob) { if err := s.resourceSvc.UpdateOCRText(job.ResourceID, text); err != nil { log.Printf("[OCR] Failed to update ocr_text for resource %s: %v", job.ResourceID, err) + s.queries.UpdateOCRJobStatus(ctx, db.UpdateOCRJobStatusParams{ + ID: job.DBID, + Status: "failed", + ErrorMessage: err.Error(), + }) return } + s.queries.UpdateOCRJobStatus(ctx, db.UpdateOCRJobStatusParams{ + ID: job.DBID, + Status: "done", + }) + s.broker.Publish("ocr_done", job.ResourceID) log.Printf("[OCR] Completed resource %s (%d chars)", job.ResourceID, len(text)) }