diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e5ed7e1..1c7d95f 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -54,3 +54,6 @@ jobs: - name: Build server run: go build ./cmd/server + + - name: Build worker + run: go build ./cmd/worker diff --git a/cmd/server/main.go b/cmd/server/main.go index 507c699..5464511 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -8,6 +8,7 @@ import ( "strings" "github.com/CeruleanFlow/cerulean/internal/dao" + "github.com/CeruleanFlow/cerulean/internal/queue" "github.com/CeruleanFlow/cerulean/internal/api" "github.com/CeruleanFlow/cerulean/internal/config" @@ -35,8 +36,24 @@ func main() { taskManager := task.NewMemoryManager() documentParser := docparser.NewPDFTextParser() searchBackend, err := buildSearchBackend(cfg) + if err != nil { + log.Fatal(err) + } + jobQueue, err := buildJobQueue(cfg) + if err != nil { + log.Fatal(err) + } + + if closer, ok := jobQueue.(interface{ Close() error }); ok { + defer func() { + if err := closer.Close(); err != nil { + log.Printf("job queue close error: %v\n", err) + } + }() + } + ragService := rag.NewService(paperRepo, searchBackend) - ingestService := ingest.NewService(paperRepo, chunkRepo, objectStore, taskManager, searchBackend, documentParser) + ingestService := ingest.NewService(paperRepo, chunkRepo, objectStore, taskManager, searchBackend, documentParser, jobQueue) router := api.NewRouter(api.RouterOptions{ Config: cfg, @@ -113,3 +130,20 @@ func buildSearchBackend(cfg config.Config) (search.Backend, error) { return nil, fmt.Errorf("unsupported CERULEAN_SEARCH_DRIVER=%q; supported: local, elastic", cfg.SearchDriver) } } + +func buildJobQueue(cfg config.Config) (queue.Queue, error) { + switch strings.ToLower(cfg.QueueDriver) { + case "", "redis": + return queue.NewRedisStreamQueue(context.Background(), queue.RedisStreamConfig{ + Addr: cfg.RedisAddr, + Password: cfg.RedisPassword, + DB: cfg.RedisDB, + Stream: cfg.QueueStream, + Group: cfg.QueueGroup, + Consumer: cfg.QueueConsumer, + }) + + default: + return nil, fmt.Errorf("unsupported CERULEAN_QUEUE_DRIVER=%q; supported: redis", cfg.QueueDriver) + } +} diff --git a/cmd/worker/main.go b/cmd/worker/main.go new file mode 100644 index 0000000..3a5d7cb --- /dev/null +++ b/cmd/worker/main.go @@ -0,0 +1,155 @@ +package main + +import ( + "context" + "fmt" + "log" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "github.com/CeruleanFlow/cerulean/internal/config" + "github.com/CeruleanFlow/cerulean/internal/dao" + "github.com/CeruleanFlow/cerulean/internal/executor" + "github.com/CeruleanFlow/cerulean/internal/ingest" + docparser "github.com/CeruleanFlow/cerulean/internal/parser" + "github.com/CeruleanFlow/cerulean/internal/pipeline" + "github.com/CeruleanFlow/cerulean/internal/queue" + "github.com/CeruleanFlow/cerulean/internal/search" + "github.com/CeruleanFlow/cerulean/internal/storage" + "github.com/CeruleanFlow/cerulean/internal/task" +) + +func main() { + cfg := config.Load() + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + q, err := queue.NewRedisStreamQueue(ctx, queue.RedisStreamConfig{ + Addr: cfg.RedisAddr, + Password: cfg.RedisPassword, + DB: cfg.RedisDB, + Stream: cfg.QueueStream, + Group: cfg.QueueGroup, + Consumer: cfg.QueueConsumer, + }) + if err != nil { + log.Fatalf("create redis queue: %v", err) + } + defer func() { + if err := q.Close(); err != nil { + log.Printf("close redis queue: %v", err) + } + }() + + database, err := dao.NewMySQLDatabase(cfg.MySQLDSN) + if err != nil { + log.Fatalf("connect mysql: %v", err) + } + + objectStore, err := buildObjectStorage(cfg) + if err != nil { + log.Fatalf("create object storage: %v", err) + } + + searchBackend, err := buildSearchBackend(cfg) + if err != nil { + log.Fatalf("create search backend: %v", err) + } + + taskManager := task.NewMemoryManager() + + documentParser := docparser.NewPDFTextParser() + + ingestService := ingest.NewService( + database.Papers, + database.Chunks, + objectStore, + taskManager, + searchBackend, + documentParser, + nil, + ) + + registry := executor.NewRegistry() + + if err := registry.Register(queue.JobTypePaperIngest, pipeline.NewPaperIngestHandler(ingestService)); err != nil { + log.Fatalf("register paper ingest handler: %v", err) + } + + if err := registry.Register(queue.JobTypePaperReindex, pipeline.NewPaperReindexHandler(ingestService)); err != nil { + log.Fatalf("register paper reindex handler: %v", err) + } + + worker, err := executor.NewWorker(q, registry, executor.WorkerOptions{ + BatchSize: cfg.WorkerBatchSize, + BlockMillis: 5000, + JobTimeout: 30 * time.Minute, + Concurrency: cfg.WorkerConcurrency, + }) + if err != nil { + log.Fatalf("create executor worker: %v", err) + } + + log.Printf( + "Cerulean worker started: redis=%s stream=%s group=%s consumer=%s batch_size=%d concurrency=%d", + cfg.RedisAddr, + cfg.QueueStream, + cfg.QueueGroup, + cfg.QueueConsumer, + cfg.WorkerBatchSize, + cfg.WorkerConcurrency, + ) + + if err := worker.Run(ctx); err != nil && ctx.Err() == nil { + log.Fatalf("worker failed: %v", err) + } +} + +func buildObjectStorage(cfg config.Config) (storage.ObjectStorage, error) { + switch strings.ToLower(cfg.StorageDriver) { + case "", "local": + return storage.NewLocalObjectStorage(cfg.LocalStorageDir) + + case "minio": + useSSL := strings.EqualFold(cfg.MinIOUseSSL, "true") + return storage.NewMinIOObjectStorage(context.Background(), storage.MinIOConfig{ + Endpoint: cfg.MinIOEndpoint, + AccessKey: cfg.MinIOAccessKey, + SecretKey: cfg.MinIOSecretKey, + Bucket: cfg.MinIOBucket, + UseSSL: useSSL, + }) + + default: + return nil, fmt.Errorf("unsupported CERULEAN_STORAGE_DRIVER=%q; supported: local, minio", cfg.StorageDriver) + } +} + +func buildSearchBackend(cfg config.Config) (search.Backend, error) { + switch strings.ToLower(cfg.SearchDriver) { + //case "", "local": + // return search.NewLocalBackend(), nil + + case "elastic", "elasticsearch", "es": + backend, err := search.NewElasticBackend(context.Background(), search.ElasticConfig{ + URL: cfg.ElasticURL, + Index: cfg.ElasticIndex, + Username: cfg.ElasticUsername, + Password: cfg.ElasticPassword, + }) + if err != nil { + return nil, err + } + if backend == nil { + return nil, fmt.Errorf("elastic backend constructor returned nil") + } + return backend, nil + + default: + return nil, fmt.Errorf("unsupported CERULEAN_SEARCH_DRIVER=%q; supported: local, elastic", cfg.SearchDriver) + } +} diff --git a/go.mod b/go.mod index 4b5a179..2ca8554 100644 --- a/go.mod +++ b/go.mod @@ -3,10 +3,12 @@ module github.com/CeruleanFlow/cerulean go 1.26.0 require ( + github.com/Haruko386/Celestial v0.0.0-20260628114458-d82cbffef886 github.com/gin-gonic/gin v1.12.0 github.com/joho/godotenv v1.5.1 github.com/ledongthuc/pdf v0.0.0-20250511090121-5959a4027728 github.com/minio/minio-go/v7 v7.2.1 + github.com/redis/go-redis/v9 v9.21.0 gorm.io/datatypes v1.2.7 gorm.io/driver/mysql v1.6.0 gorm.io/gorm v1.31.2 @@ -51,6 +53,7 @@ require ( github.com/ugorji/go/codec v1.3.1 // indirect github.com/zeebo/xxh3 v1.1.0 // indirect go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect + go.uber.org/atomic v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/arch v0.22.0 // indirect golang.org/x/crypto v0.51.0 // indirect diff --git a/go.sum b/go.sum index df595ad..3b031e4 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,11 @@ filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= +github.com/Haruko386/Celestial v0.0.0-20260628114458-d82cbffef886 h1:o9fH5b0Rn8cwO7sRxf2Gnop4y9DIyve8skHzYEdO2D4= +github.com/Haruko386/Celestial v0.0.0-20260628114458-d82cbffef886/go.mod h1:Dv2Zzto5bivYyYxYOAVPNumTAG5amWkB6nh9M0J1zQg= +github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= +github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c= +github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= +github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0= github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM= github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE= @@ -102,6 +108,8 @@ github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw= github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/redis/go-redis/v9 v9.21.0 h1:FPBE4hhbAke+TLmcY3WkpbDffJEomdqPn3HYiqAtL9E= +github.com/redis/go-redis/v9 v9.21.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU= @@ -129,6 +137,8 @@ github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s= go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= diff --git a/internal/api/handler.go b/internal/api/handler.go index 3200b26..2a30c84 100644 --- a/internal/api/handler.go +++ b/internal/api/handler.go @@ -267,17 +267,12 @@ func (h *Handler) ReindexPaper(c *gin.Context) { return } - optCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) - defer cancel() - - if err := h.ingest.ReindexPaper(optCtx, id); err != nil { + job, err := h.ingest.StartPaperReindex(c.Request.Context(), id) + if err != nil { writeError(c, http.StatusInternalServerError, err) return } - c.JSON(http.StatusOK, gin.H{ - "paper_id": id, - "status": "reindexed", - }) + c.JSON(http.StatusAccepted, job) } func (h *Handler) GetTask(c *gin.Context) { diff --git a/internal/config/config.go b/internal/config/config.go index d6e17ed..31b297c 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -2,6 +2,7 @@ package config import ( "os" + "strconv" "github.com/joho/godotenv" _ "github.com/joho/godotenv" @@ -35,6 +36,18 @@ type Config struct { LLMBaseURL string LLMModel string + + RedisAddr string + RedisPassword string + RedisDB int + + QueueDriver string + QueueStream string + QueueGroup string + QueueConsumer string + + WorkerConcurrency int + WorkerBatchSize int } func Load() Config { @@ -70,6 +83,18 @@ func Load() Config { LLMBaseURL: env("CERULEAN_LLM_BASE_URL", ""), LLMModel: env("CERULEAN_LLM_MODEL", ""), + + RedisAddr: env("CERULEAN_REDIS_ADDR", "127.0.0.1:6379"), + RedisPassword: env("CERULEAN_REDIS_PASSWORD", ""), + RedisDB: envInt("CERULEAN_REDIS_DB", 0), + + QueueDriver: env("CERULEAN_QUEUE_DRIVER", "redis"), + QueueStream: env("CERULEAN_QUEUE_STREAM", "cerulean_tasks"), + QueueConsumer: env("CERULEAN_QUEUE_CONSUMER", "worker_local_1"), + QueueGroup: env("CERULEAN_QUEUE_GROUP", "cerulean_workers"), + + WorkerBatchSize: envInt("CERULEAN_WORKER_BATCH_SIZE", 4), + WorkerConcurrency: envInt("CERULEAN_WORKER_CONCURRENCY", 16), } } @@ -79,3 +104,15 @@ func env(key, fallback string) string { } return fallback } + +func envInt(key string, fallback int) int { + value := os.Getenv(key) + if value == "" { + return fallback + } + i, err := strconv.Atoi(value) + if err != nil { + return fallback + } + return i +} diff --git a/internal/executor/executor.go b/internal/executor/executor.go new file mode 100644 index 0000000..63f6879 --- /dev/null +++ b/internal/executor/executor.go @@ -0,0 +1,17 @@ +package executor + +import ( + "context" + + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type Handler interface { + Handle(ctx context.Context, job queue.Job) error +} + +type HandlerFunc func(ctx context.Context, job queue.Job) error + +func (f HandlerFunc) Handle(ctx context.Context, job queue.Job) error { + return f(ctx, job) +} diff --git a/internal/executor/registry.go b/internal/executor/registry.go new file mode 100644 index 0000000..d0f9471 --- /dev/null +++ b/internal/executor/registry.go @@ -0,0 +1,46 @@ +package executor + +import ( + "context" + "errors" + "strings" + + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type Registry struct { + handlers map[string]Handler +} + +func NewRegistry() *Registry { + return &Registry{ + handlers: make(map[string]Handler), + } +} + +func (r *Registry) Register(jobType string, handler Handler) error { + jobType = strings.TrimSpace(jobType) + if jobType == "" { + return errors.New("job type is empty") + } + + if handler == nil { + return errors.New("handler is nil") + } + + r.handlers[jobType] = handler + return nil +} + +func (r *Registry) Execute(ctx context.Context, job queue.Job) error { + jobType := strings.TrimSpace(job.Type) + if jobType == "" { + return errors.New("job type is empty") + } + + handler, ok := r.handlers[jobType] + if !ok { + return errors.New("handler not found") + } + return handler.Handle(ctx, job) +} diff --git a/internal/executor/worker.go b/internal/executor/worker.go new file mode 100644 index 0000000..efbf56a --- /dev/null +++ b/internal/executor/worker.go @@ -0,0 +1,224 @@ +package executor + +import ( + "context" + "errors" + "log" + "time" + + celestial "github.com/Haruko386/Celestial" + + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type Worker struct { + queue queue.Queue + registry Registry + + batchSize int + blockMillis int64 + jobTimeout time.Duration + concurrency int +} + +type WorkerOptions struct { + BatchSize int + BlockMillis int64 + JobTimeout time.Duration + Concurrency int +} + +type JobResult struct { + RedisID string + TaskID string + Type string + PaperID string +} + +func NewWorker(queue queue.Queue, registry *Registry, options WorkerOptions) (*Worker, error) { + if queue == nil { + return nil, errors.New("queue is nil") + } + if registry == nil { + return nil, errors.New("registry is nil") + } + + batchSize := options.BatchSize + if batchSize == 0 { + batchSize = 16 + } + + blockMillis := options.BlockMillis + if blockMillis == 0 { + blockMillis = 5000 + } + + jobTimeout := options.JobTimeout + if jobTimeout == 0 { + jobTimeout = 30 * time.Minute + } + + concurrency := options.Concurrency + if concurrency == 0 { + concurrency = 4 + } + + return &Worker{ + queue: queue, + registry: *registry, + batchSize: batchSize, + blockMillis: blockMillis, + jobTimeout: jobTimeout, + concurrency: concurrency, + }, nil +} + +func (w *Worker) Run(ctx context.Context) error { + log.Printf( + "executor worker started: batch_size=%d block_millis=%d job_timeout=%s", + w.batchSize, + w.blockMillis, + w.jobTimeout, + ) + + for { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + messages, err := w.queue.DequeueBatch(ctx, w.batchSize, w.blockMillis) + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + + log.Printf("dequeue batch failed: %v", err) + time.Sleep(2 * time.Second) + continue + } + + if len(messages) == 0 { + continue + } + + if err := w.handleBatch(ctx, messages); err != nil { + log.Printf("handle batch finished with error: %v", err) + } + } +} + +func (w *Worker) handleBatch(ctx context.Context, messages []queue.Message) error { + if len(messages) == 0 { + return nil + } + + dispatcher := celestial.New[queue.Message, JobResult](celestial.Config{ + Workers: w.concurrency, + QueueSize: w.batchSize, + StopOnError: false, + }) + + run := dispatcher.RunSlice( + ctx, + messages, + func(ctx context.Context, worker celestial.Worker, msg queue.Message) (JobResult, error) { + log.Printf( + "celestial worker=%d picked job: redis_id=%s task_id=%s type=%s", + worker.Index, + msg.RedisID, + msg.Job.TaskID, + msg.Job.Type, + ) + + return w.handleMessage(ctx, msg) + }, + ) + + for result := range run.Results() { + if result.Err != nil { + log.Printf( + "celestial job failed: worker=%d task_index=%d err=%v", + result.WorkerIndex, + result.TaskID, + result.Err, + ) + continue + } + + log.Printf( + "celestial job done: worker=%d task_index=%d redis_id=%s task_id=%s type=%s paper_id=%s", + result.WorkerIndex, + result.TaskID, + result.Value.RedisID, + result.Value.TaskID, + result.Value.Type, + result.Value.PaperID, + ) + } + + if err := run.Err(); err != nil { + return err + } + + return nil +} + +func (w *Worker) handleMessage(ctx context.Context, msg queue.Message) (JobResult, error) { + result := JobResult{ + RedisID: msg.RedisID, + TaskID: msg.Job.TaskID, + Type: msg.Job.Type, + PaperID: msg.Job.PaperID, + } + + log.Printf( + "start job: redis_id=%s job_id=%s task_id=%s type=%s paper_id=%s attempt=%d", + msg.RedisID, + msg.Job.ID, + msg.Job.TaskID, + msg.Job.Type, + msg.Job.PaperID, + msg.Job.Attempt, + ) + + jobCtx, cancel := context.WithTimeout(ctx, w.jobTimeout) + defer cancel() + + err := w.registry.Execute(jobCtx, msg.Job) + if err != nil { + log.Printf( + "job failed: redis_id=%s task_id=%s type=%s err=%v", + msg.RedisID, + msg.Job.TaskID, + msg.Job.Type, + err, + ) + + nackCtx, cancel := context.WithTimeout(context.Background(), w.jobTimeout) + defer cancel() + + if nackErr := w.queue.Nack(nackCtx, msg, err); nackErr != nil { + log.Printf("nack job failed: redis_id=%s task_id=%s err=%v", msg.RedisID, msg.Job.TaskID, nackErr) + } + + return result, err + } + + ackCtx, cancel := context.WithTimeout(context.Background(), w.jobTimeout) + defer cancel() + + if err := w.queue.Ack(ackCtx, msg); err != nil { + log.Printf("ack job failed: redis_id=%s task_id=%s err=%v", msg.RedisID, msg.Job.TaskID, err) + return result, err + } + + log.Printf( + "job succeeded and acked: redis_id=%s task_id=%s type=%s", + msg.RedisID, + msg.Job.TaskID, + msg.Job.Type, + ) + return result, nil +} diff --git a/internal/ingest/service.go b/internal/ingest/service.go index 9fd8150..9d7d13b 100644 --- a/internal/ingest/service.go +++ b/internal/ingest/service.go @@ -10,6 +10,7 @@ import ( "time" "github.com/CeruleanFlow/cerulean/internal/domain" + "github.com/CeruleanFlow/cerulean/internal/queue" "github.com/CeruleanFlow/cerulean/internal/repository" "github.com/CeruleanFlow/cerulean/internal/search" "github.com/CeruleanFlow/cerulean/internal/storage" @@ -25,6 +26,8 @@ type Service struct { search search.Backend tasks task.Manager parser docparser.Parser + + jobQueue queue.Queue } func NewService( @@ -34,18 +37,28 @@ func NewService( tasks task.Manager, searchBackend search.Backend, parser docparser.Parser, + jobQueue queue.Queue, ) *Service { return &Service{ - papers: papers, - chunks: chunks, - store: store, - tasks: tasks, - search: searchBackend, - parser: parser, + papers: papers, + chunks: chunks, + store: store, + tasks: tasks, + search: searchBackend, + parser: parser, + jobQueue: jobQueue, } } func (s *Service) StartPaperIngest(ctx context.Context, paperID string) (task.Task, error) { + paperID = strings.TrimSpace(paperID) + if paperID == "" { + return task.Task{}, fmt.Errorf("paper id is empty") + } + if s.jobQueue == nil { + return task.Task{}, fmt.Errorf("job queue is empty") + } + paper, err := s.papers.Get(ctx, paperID) if err != nil { return task.Task{}, err @@ -75,22 +88,68 @@ func (s *Service) StartPaperIngest(ctx context.Context, paperID string) (task.Ta return task.Task{}, err } - // Keep this asynchronous so the public API shape is already compatible with - // the future PaddleOCR worker and queue based pipeline. - go func(job task.Task, paper domain.Paper) { - defer func() { - if r := recover(); r != nil { - s.fail(context.Background(), job, paper, fmt.Errorf("panic during paper ingest: %v", r)) - } - }() + redisJob := queue.Job{ + ID: job.ID, + TaskID: job.ID, + Type: queue.JobTypePaperIngest, + PaperID: paper.ID, + Attempt: 0, + CreatedAt: now, + } - taskCtx, cancel := context.WithTimeout(context.Background(), 10*time.Minute) - defer cancel() + if err := s.jobQueue.Enqueue(optCtx, redisJob); err != nil { + s.fail(context.Background(), job, paper, err) + return task.Task{}, fmt.Errorf("enqueue ingest job: %w", err) + } - if err := s.runPDFTextIngest(taskCtx, job, paper); err != nil { - s.fail(context.Background(), job, paper, err) - } - }(job, paper) + return job, nil +} + +func (s *Service) StartPaperReindex(ctx context.Context, paperID string) (task.Task, error) { + paperID = strings.TrimSpace(paperID) + if paperID == "" { + return task.Task{}, fmt.Errorf("paper id is empty") + } + if s.jobQueue == nil { + return task.Task{}, fmt.Errorf("job queue is empty") + } + + if _, err := s.papers.Get(ctx, paperID); err != nil { + return task.Task{}, fmt.Errorf("paper not found: %w", err) + } + + optCtx, cancel := context.WithTimeout(ctx, 30*time.Second) + defer cancel() + + now := time.Now() + + job := task.Task{ + ID: fmt.Sprintf("task_%d", now.UnixNano()), + PaperID: paperID, + Type: queue.JobTypePaperReindex, + Status: task.Queued, + Message: "queued paper reindex", + CreatedAt: now, + UpdatedAt: now, + } + + if err := s.tasks.Create(optCtx, job); err != nil { + return task.Task{}, fmt.Errorf("create paper reindex: %w", err) + } + + redisJob := queue.Job{ + ID: job.ID, + TaskID: job.ID, + Type: queue.JobTypePaperReindex, + PaperID: paperID, + Attempt: 0, + CreatedAt: now, + } + + if err := s.jobQueue.Enqueue(optCtx, redisJob); err != nil { + s.failTaskOnly(ctx, job, err) + return task.Task{}, fmt.Errorf("enqueue paper reindex: %w", err) + } return job, nil } @@ -177,6 +236,7 @@ func (s *Service) runPDFTextIngest(ctx context.Context, job task.Task, paper dom return nil } +// fail set failed status for paper and task func (s *Service) fail(ctx context.Context, job task.Task, paper domain.Paper, err error) { opCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() @@ -194,6 +254,18 @@ func (s *Service) fail(ctx context.Context, job task.Task, paper domain.Paper, e _ = s.tasks.Update(opCtx, job) } +func (s *Service) failTaskOnly(ctx context.Context, job task.Task, err error) { + opCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + now := time.Now() + job.Status = task.Failed + job.Message = err.Error() + job.UpdatedAt = now + + _ = s.tasks.Update(opCtx, job) +} + // downloadOriginalPDF download original PDF to tmp func (s *Service) downloadOriginalPDF(ctx context.Context, paper domain.Paper) (string, func(), error) { if s.store == nil { @@ -282,3 +354,75 @@ func (s *Service) ReindexPaper(ctx context.Context, paperID string) error { return nil } + +func (s *Service) ProcessPaperIngest(ctx context.Context, paperID, taskID string) error { + taskID = strings.TrimSpace(taskID) + paperID = strings.TrimSpace(paperID) + + if paperID == "" { + return fmt.Errorf("paper id cannot be empty") + } + if taskID == "" { + return fmt.Errorf("task id cannot be empty") + } + + job, ok := s.tasks.Get(ctx, taskID) + if !ok { + return fmt.Errorf("task %s not found", taskID) + } + + paper, err := s.papers.Get(ctx, paperID) + if err != nil { + return fmt.Errorf("get paper %s from storage: %w", paperID, err) + } + + if err := s.runPDFTextIngest(ctx, job, paper); err != nil { + s.fail(ctx, job, paper, err) + return err + } + return nil +} + +func (s *Service) ProcessPaperReindex(ctx context.Context, paperID, taskID string) error { + taskID = strings.TrimSpace(taskID) + paperID = strings.TrimSpace(paperID) + + if paperID == "" { + return fmt.Errorf("paper id cannot be empty") + } + if taskID == "" { + return fmt.Errorf("task id cannot be empty") + } + + job, ok := s.tasks.Get(ctx, taskID) + if !ok { + return fmt.Errorf("task %s not found", taskID) + } + + now := time.Now() + job.Status = task.Running + job.Message = "reindexing paper chunks to Elasticsearch" + job.UpdatedAt = now + + if err := s.tasks.Update(ctx, job); err != nil { + return fmt.Errorf("update job: %w", err) + } + + if err := s.ReindexPaper(ctx, paperID); err != nil { + s.failTaskOnly(ctx, job, err) + return err + } + + now = time.Now() + job.Status = task.Succeeded + job.Message = "reindexing paper chunks to Elasticsearch" + job.UpdatedAt = now + + finishCtx, cancel := context.WithTimeout(ctx, 15*time.Second) + defer cancel() + + if err := s.tasks.Update(finishCtx, job); err != nil { + return fmt.Errorf("update job: %w", err) + } + return nil +} diff --git a/internal/pipeline/log_job.go b/internal/pipeline/log_job.go new file mode 100644 index 0000000..3922746 --- /dev/null +++ b/internal/pipeline/log_job.go @@ -0,0 +1,27 @@ +package pipeline + +import ( + "context" + "log" + + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type LogJobHandler struct{} + +func NewLogJobHandler() *LogJobHandler { + return &LogJobHandler{} +} + +func (h *LogJobHandler) Handle(ctx context.Context, job queue.Job) error { + log.Printf( + "pipeline log handler: job_id=%s task_id=%s type=%s paper_id=%s attempt=%d", + job.ID, + job.TaskID, + job.Type, + job.PaperID, + job.Attempt, + ) + + return nil +} diff --git a/internal/pipeline/paper_ingest.go b/internal/pipeline/paper_ingest.go new file mode 100644 index 0000000..9af5521 --- /dev/null +++ b/internal/pipeline/paper_ingest.go @@ -0,0 +1,37 @@ +package pipeline + +import ( + "context" + "fmt" + "strings" + + "github.com/CeruleanFlow/cerulean/internal/ingest" + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type PaperIngestHandler struct { + ingest *ingest.Service +} + +func NewPaperIngestHandler(ingest *ingest.Service) *PaperIngestHandler { + return &PaperIngestHandler{ + ingest: ingest, + } +} + +func (h *PaperIngestHandler) Handle(ctx context.Context, job queue.Job) error { + if h.ingest == nil { + return fmt.Errorf("ingest service is nil") + } + + taskID := strings.TrimSpace(job.TaskID) + paperID := strings.TrimSpace(job.PaperID) + if taskID == "" { + return fmt.Errorf("task id is empty") + } + if paperID == "" { + return fmt.Errorf("paper id is empty") + } + + return h.ingest.ProcessPaperIngest(ctx, taskID, paperID) +} diff --git a/internal/pipeline/paper_reindex.go b/internal/pipeline/paper_reindex.go new file mode 100644 index 0000000..8f4c1cc --- /dev/null +++ b/internal/pipeline/paper_reindex.go @@ -0,0 +1,37 @@ +package pipeline + +import ( + "context" + "fmt" + "strings" + + "github.com/CeruleanFlow/cerulean/internal/ingest" + "github.com/CeruleanFlow/cerulean/internal/queue" +) + +type PaperReindexHandler struct { + ingest *ingest.Service +} + +func NewPaperReindexHandler(ingest *ingest.Service) *PaperReindexHandler { + return &PaperReindexHandler{ + ingest: ingest, + } +} + +func (h *PaperReindexHandler) Handle(ctx context.Context, job queue.Job) error { + if h.ingest == nil { + return fmt.Errorf("ingest service is nil") + } + + taskID := strings.TrimSpace(job.TaskID) + paperID := strings.TrimSpace(job.PaperID) + if taskID == "" { + return fmt.Errorf("task id is empty") + } + if paperID == "" { + return fmt.Errorf("paper id is empty") + } + + return h.ingest.ProcessPaperReindex(ctx, paperID, taskID) +} diff --git a/internal/queue/job.go b/internal/queue/job.go new file mode 100644 index 0000000..bae89a8 --- /dev/null +++ b/internal/queue/job.go @@ -0,0 +1,22 @@ +package queue + +import "time" + +const ( + JobTypePaperIngest = "paper_ingest" + JobTypePaperReindex = "paper_reindex" +) + +type Job struct { + ID string `json:"id"` + TaskID string `json:"task_id"` + Type string `json:"type"` + PaperID string `json:"paper_id"` + Attempt int `json:"attempt"` + CreatedAt time.Time `json:"created_at"` +} + +type Message struct { + RedisID string `json:"redis_id"` + Job Job `json:"job"` +} diff --git a/internal/queue/queue.go b/internal/queue/queue.go new file mode 100644 index 0000000..0cfc425 --- /dev/null +++ b/internal/queue/queue.go @@ -0,0 +1,10 @@ +package queue + +import "context" + +type Queue interface { + Enqueue(ctx context.Context, job Job) error + DequeueBatch(ctx context.Context, max int, blockMillis int64) ([]Message, error) + Ack(ctx context.Context, msg Message) error + Nack(ctx context.Context, msg Message, reason error) error +} diff --git a/internal/queue/redis_stream.go b/internal/queue/redis_stream.go new file mode 100644 index 0000000..6ea0e0f --- /dev/null +++ b/internal/queue/redis_stream.go @@ -0,0 +1,169 @@ +package queue + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "github.com/redis/go-redis/v9" +) + +type RedisStreamConfig struct { + Addr string `json:"addr"` + Password string `json:"password"` + DB int `json:"db"` + + Stream string `json:"stream"` + Group string `json:"group"` + Consumer string `json:"consumer"` +} + +type RedisStreamQueue struct { + client *redis.Client + stream string + group string + consumer string +} + +func NewRedisStreamQueue(ctx context.Context, cfg RedisStreamConfig) (*RedisStreamQueue, error) { + // precheck + stream := strings.TrimSpace(cfg.Stream) + if stream == "" { + stream = "cerulean_tasks" + } + + group := strings.TrimSpace(cfg.Group) + if group == "" { + group = "cerulean_workers" + } + + consumer := strings.TrimSpace(cfg.Consumer) + if consumer == "" { + consumer = "worker_local_1" + } + + // initialize the redis + client := redis.NewClient(&redis.Options{ + Addr: cfg.Addr, + Password: cfg.Password, + DB: cfg.DB, + }) + // return if cannot initialize redis + if err := client.Ping(ctx).Err(); err != nil { + return nil, fmt.Errorf("failed to connect to redis: %w", err) + } + + q := &RedisStreamQueue{ + client: client, + stream: stream, + group: group, + consumer: consumer, + } + + if err := q.ensureGroup(ctx); err != nil { + return nil, err + } + + return q, nil +} + +func (q *RedisStreamQueue) ensureGroup(ctx context.Context) error { + err := q.client.XGroupCreateMkStream(ctx, q.stream, q.group, "0").Err() + if err == nil { + return nil + } + + if strings.Contains(err.Error(), "BUSYGROUP") { + return nil + } + + return fmt.Errorf("failed to create group '%s': %w", q.group, err) + +} + +func (q *RedisStreamQueue) Enqueue(ctx context.Context, job Job) error { + payload, err := json.Marshal(job) + if err != nil { + return fmt.Errorf("failed to encode job: %w", err) + } + + return q.client.XAdd(ctx, &redis.XAddArgs{ + Stream: q.stream, + Values: map[string]any{ + "payload": string(payload), + "type": job.Type, + "task_id": job.TaskID, + "paper_id": job.PaperID, + }, + }).Err() +} + +func (q *RedisStreamQueue) DequeueBatch(ctx context.Context, max int, blockMillis int64) ([]Message, error) { + if max <= 0 { + max = 16 + } + + if blockMillis <= 0 { + blockMillis = 5000 + } + + streams, err := q.client.XReadGroup(ctx, &redis.XReadGroupArgs{ + Group: q.group, + Consumer: q.consumer, + Streams: []string{q.stream, ">"}, + Count: int64(max), + Block: time.Duration(blockMillis) * time.Millisecond, + }).Result() + + if err != nil { + if errors.Is(err, redis.Nil) { + return nil, nil + } + return nil, fmt.Errorf("xreadgroup: %w", err) + } + + messages := make([]Message, 0) + + for _, stream := range streams { + for _, redisMsg := range stream.Messages { + raw, ok := redisMsg.Values["payload"].(string) + if !ok { + continue + } + + var job Job + if err := json.Unmarshal([]byte(raw), &job); err != nil { + continue + } + + messages = append(messages, Message{ + RedisID: redisMsg.ID, + Job: job, + }) + } + } + return messages, nil +} + +func (q *RedisStreamQueue) Ack(ctx context.Context, msg Message) error { + if strings.TrimSpace(msg.RedisID) == "" { + return nil + } + + return q.client.XAck(ctx, q.stream, q.group, msg.RedisID).Err() +} + +func (q *RedisStreamQueue) Nack(ctx context.Context, msg Message, reason error) error { + return nil +} + +// Close the redis +func (q *RedisStreamQueue) Close() error { + if q.client == nil { + return nil + } + return q.client.Close() +} diff --git a/internal/search/elastic.go b/internal/search/elastic.go index bf39878..38e70ac 100644 --- a/internal/search/elastic.go +++ b/internal/search/elastic.go @@ -73,7 +73,22 @@ func (b *ElasticBackend) Name() string { // EnsureIndex Make sure index exists func (b *ElasticBackend) EnsureIndex(ctx context.Context) error { - req, err := http.NewRequestWithContext(ctx, http.MethodHead, b.endpoint("/"+url.PathEscape(b.index)), nil) + if b == nil { + return fmt.Errorf("elastic backend is nil") + } + if b.client == nil { + return fmt.Errorf("elastic http client is nil") + } + if strings.TrimSpace(b.index) == "" { + return fmt.Errorf("elastic index is empty") + } + + req, err := http.NewRequestWithContext( + ctx, + http.MethodHead, + b.endpoint("/"+url.PathEscape(b.index)), + nil, + ) if err != nil { return err } @@ -86,13 +101,22 @@ func (b *ElasticBackend) EnsureIndex(ctx context.Context) error { } defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { + switch resp.StatusCode { + case http.StatusOK: + // 200 表示 index 已经存在。 return nil - } - if resp.StatusCode != http.StatusNotFound { + case http.StatusNotFound: + // 404 表示 index 不存在,继续创建。 + + default: body, _ := io.ReadAll(resp.Body) - return fmt.Errorf("check elastic index %q failed: status=%d body=%s", b.index, resp.StatusCode, string(body)) + return fmt.Errorf( + "check elastic index %q failed: status=%d body=%s", + b.index, + resp.StatusCode, + string(body), + ) } mapping := map[string]any{ @@ -124,10 +148,11 @@ func (b *ElasticBackend) EnsureIndex(ctx context.Context) error { }, } - _, err = b.doJSON(ctx, http.MethodPut, b.endpoint("/"+url.PathEscape(b.index)), mapping) + _, err = b.doJSON(ctx, http.MethodPut, "/"+url.PathEscape(b.index), mapping) if err != nil { return fmt.Errorf("create elastic index %q: %w", b.index, err) } + return nil } diff --git a/server b/server new file mode 100755 index 0000000..29485ef Binary files /dev/null and b/server differ diff --git a/worker b/worker new file mode 100755 index 0000000..ab677ae Binary files /dev/null and b/worker differ