From 9883def0c202d65305877bb41bdbbac2ad2b824f Mon Sep 17 00:00:00 2001 From: Emil Date: Sun, 2 Aug 2026 19:23:44 +0300 Subject: [PATCH] Add an embedded SQLite storage backend for single-binary deployments --- coordinator/cmd/coordinator/main.go | 106 +++++- coordinator/go.mod | 8 + coordinator/go.sum | 17 + coordinator/internal/infra/config.go | 20 +- .../internal/storage/sqlite/artifact_repo.go | 91 +++++ .../internal/storage/sqlite/job_repo.go | 168 +++++++++ .../internal/storage/sqlite/migrate.go | 85 +++++ .../storage/sqlite/migrations/0001_core.sql | 96 +++++ coordinator/internal/storage/sqlite/sqlite.go | 172 +++++++++ .../internal/storage/sqlite/sqlite_test.go | 349 ++++++++++++++++++ .../internal/storage/sqlite/stats_repo.go | 63 ++++ .../internal/storage/sqlite/task_repo.go | 314 ++++++++++++++++ .../storage/sqlite/task_result_repo.go | 41 ++ .../internal/storage/sqlite/ui_read_repo.go | 145 ++++++++ .../internal/storage/sqlite/worker_repo.go | 92 +++++ 15 files changed, 1744 insertions(+), 23 deletions(-) create mode 100644 coordinator/internal/storage/sqlite/artifact_repo.go create mode 100644 coordinator/internal/storage/sqlite/job_repo.go create mode 100644 coordinator/internal/storage/sqlite/migrate.go create mode 100644 coordinator/internal/storage/sqlite/migrations/0001_core.sql create mode 100644 coordinator/internal/storage/sqlite/sqlite.go create mode 100644 coordinator/internal/storage/sqlite/sqlite_test.go create mode 100644 coordinator/internal/storage/sqlite/stats_repo.go create mode 100644 coordinator/internal/storage/sqlite/task_repo.go create mode 100644 coordinator/internal/storage/sqlite/task_result_repo.go create mode 100644 coordinator/internal/storage/sqlite/ui_read_repo.go create mode 100644 coordinator/internal/storage/sqlite/worker_repo.go diff --git a/coordinator/cmd/coordinator/main.go b/coordinator/cmd/coordinator/main.go index dbfc41b..1f0deea 100644 --- a/coordinator/cmd/coordinator/main.go +++ b/coordinator/cmd/coordinator/main.go @@ -14,6 +14,7 @@ import ( "github.com/emil28092005/SciMesh/coordinator/internal/metrics" "github.com/emil28092005/SciMesh/coordinator/internal/storage/blob" "github.com/emil28092005/SciMesh/coordinator/internal/storage/postgres" + "github.com/emil28092005/SciMesh/coordinator/internal/storage/sqlite" httptransport "github.com/emil28092005/SciMesh/coordinator/internal/transport/http" "github.com/emil28092005/SciMesh/coordinator/internal/usecase" "github.com/emil28092005/SciMesh/coordinator/internal/workloads" @@ -44,6 +45,24 @@ func main() { } } +// storageDeps carries the engine-specific database handles and the repository +// implementations. The usecases below only ever see the ports. +type storageDeps struct { + tx usecase.TxManager + taskRepo usecase.TaskRepository + jobRepo usecase.JobRepository + workerRepo usecase.WorkerRepository + artifactRepo usecase.ArtifactRepository + uiReadRepo usecase.UIReadRepository + taskResultRepo usecase.TaskResultRepository + statsRepo interface { + Counts(ctx context.Context) (tasks, jobs, workers map[string]int, err error) + } + ready func(ctx context.Context) error + migrate func(ctx context.Context, log *slog.Logger) error + close func() +} + func run() error { // Bootstrap logger, used only until config says where logs should go. It // writes to stderr so it never contaminates the configured stdout stream. @@ -66,17 +85,25 @@ func run() error { ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) defer stop() - pool, err := infra.NewPool(ctx, cfg, log) + var deps *storageDeps + switch cfg.DatabaseEngine { + case "sqlite": + deps, err = openSQLite(ctx, cfg, log) + case "postgres": + deps, err = openPostgres(ctx, cfg, log) + default: + err = fmt.Errorf("SCIMESH_DB must be sqlite or postgres") + } if err != nil { - log.Error("connect database", "err", err) + log.Error("init storage", "err", err) return err } - defer pool.Close() + defer deps.close() // A downloaded binary provisions its own schema; AUTO_MIGRATE=false keeps // out-of-band migration workflows (the migrate CLI, CI, managed databases). if cfg.AutoMigrate { - if err := postgres.Migrate(ctx, cfg.DatabaseURL, log); err != nil { + if err := deps.migrate(ctx, log); err != nil { log.Error("apply migrations", "err", err) return err } @@ -88,16 +115,9 @@ func run() error { return err } - var ( - clk = infra.NewClock() - tx = postgres.NewTxManager(pool) - taskRepo = postgres.NewTaskRepo(pool) - jobRepo = postgres.NewJobRepo(pool) - workerRepo = postgres.NewWorkerRepo(pool) - artifactRepo = postgres.NewArtifactRepo(pool) - uiReadRepo = postgres.NewUIReadRepo(pool) - taskResultRepo = postgres.NewTaskResultRepo(pool) - ) + clk := infra.NewClock() + tx, taskRepo, jobRepo, workerRepo, artifactRepo, uiReadRepo, taskResultRepo := + deps.tx, deps.taskRepo, deps.jobRepo, deps.workerRepo, deps.artifactRepo, deps.uiReadRepo, deps.taskResultRepo catalog, err := workloads.Load() if err != nil { @@ -125,7 +145,7 @@ func run() error { } // Background reapers are tracked so shutdown can wait for them. Without this - // the process would exit mid-UPDATE, and the deferred pool.Close() would pull + // the process would exit mid-UPDATE, and the deferred close() would pull // connections out from under them. expireLeases := usecase.NewExpireLeases(taskRepo, jobRepo, tx, clk, catalog) markOffline := usecase.NewMarkWorkersOffline(workerRepo, clk, cfg.WorkerOfflineAfter) @@ -147,16 +167,15 @@ func run() error { // Business metrics: gauges of tasks/jobs/workers by status, sampled from the // database on every Prometheus scrape. - statsRepo := postgres.NewStatsRepo(pool) m := metrics.New() m.RegisterBusiness(func(ctx context.Context) (metrics.Stats, error) { - tasks, jobs, workers, err := statsRepo.Counts(ctx) + tasks, jobs, workers, err := deps.statsRepo.Counts(ctx) return metrics.Stats{Tasks: tasks, Jobs: jobs, Workers: workers}, err }) - // pool.Ping backs /health: readiness means the database answers, not just + // deps.ready backs /health: readiness means the database answers, not just // that the process is alive. - api := httptransport.NewServer(useCases, log, cfg.RequestTimeout, cfg.HeartbeatInterval, cfg.MaxUploadBytes, cfg.JWTSecret, cfg.UserserviceURL, m, pool.Ping, cfg.PublicCoordinatorURL, cfg.PublicUserserviceURL, cfg.DocsDir) + api := httptransport.NewServer(useCases, log, cfg.RequestTimeout, cfg.HeartbeatInterval, cfg.MaxUploadBytes, cfg.JWTSecret, cfg.UserserviceURL, m, deps.ready, cfg.PublicCoordinatorURL, cfg.PublicUserserviceURL, cfg.DocsDir) err = infra.RunServer(ctx, log, cfg.Addr, api.Handler(cfg.Token, cfg.UIToken)) // Shutdown order matters, and defers alone cannot express it (they run @@ -164,7 +183,7 @@ func run() error { // // 1. stop() cancel the context, telling the reaper to finish // 2. wg.Wait() let it return from its current tick - // 3. deferred pool.Close() closes an idle pool, not a busy one + // 3. deferred close() closes an idle pool, not a busy one // // Calling stop() here also covers the path where RunServer failed on its // own: the context would never be cancelled otherwise and wg.Wait() @@ -175,3 +194,50 @@ func run() error { return err } + +// openSQLite opens the embedded database and builds the sqlite repositories. +func openSQLite(ctx context.Context, cfg infra.Config, log *slog.Logger) (*storageDeps, error) { + if err := os.MkdirAll(cfg.StorageDir, 0o750); err != nil { + return nil, fmt.Errorf("create storage dir: %w", err) + } + db, err := sqlite.Open(cfg.DBPath) + if err != nil { + return nil, err + } + closeOnce := &sync.Once{} + return &storageDeps{ + tx: sqlite.NewTxManager(db), + taskRepo: sqlite.NewTaskRepo(db), + jobRepo: sqlite.NewJobRepo(db), + workerRepo: sqlite.NewWorkerRepo(db), + artifactRepo: sqlite.NewArtifactRepo(db), + uiReadRepo: sqlite.NewUIReadRepo(db), + taskResultRepo: sqlite.NewTaskResultRepo(db), + statsRepo: sqlite.NewStatsRepo(db), + ready: func(ctx context.Context) error { return db.PingContext(ctx) }, + migrate: func(ctx context.Context, log *slog.Logger) error { return sqlite.Migrate(ctx, db, log) }, + close: func() { closeOnce.Do(func() { _ = db.Close() }) }, + }, nil +} + +// openPostgres connects to PostgreSQL and builds the postgres repositories. +func openPostgres(ctx context.Context, cfg infra.Config, log *slog.Logger) (*storageDeps, error) { + pool, err := infra.NewPool(ctx, cfg, log) + if err != nil { + return nil, err + } + closeOnce := &sync.Once{} + return &storageDeps{ + tx: postgres.NewTxManager(pool), + taskRepo: postgres.NewTaskRepo(pool), + jobRepo: postgres.NewJobRepo(pool), + workerRepo: postgres.NewWorkerRepo(pool), + artifactRepo: postgres.NewArtifactRepo(pool), + uiReadRepo: postgres.NewUIReadRepo(pool), + taskResultRepo: postgres.NewTaskResultRepo(pool), + statsRepo: postgres.NewStatsRepo(pool), + ready: func(ctx context.Context) error { return pool.Ping(ctx) }, + migrate: func(ctx context.Context, log *slog.Logger) error { return postgres.Migrate(ctx, cfg.DatabaseURL, log) }, + close: func() { closeOnce.Do(pool.Close) }, + }, nil +} diff --git a/coordinator/go.mod b/coordinator/go.mod index 7d9cb06..9b8fb31 100644 --- a/coordinator/go.mod +++ b/coordinator/go.mod @@ -16,18 +16,26 @@ require ( require ( github.com/beorn7/perks v1.0.1 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect github.com/jackc/puddle/v2 v2.2.1 // indirect github.com/lann/builder v0.0.0-20180802200727-47ae307949d0 // indirect github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/common v0.55.0 // indirect github.com/prometheus/procfs v0.21.1 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect golang.org/x/crypto v0.17.0 // indirect golang.org/x/sync v0.22.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.40.0 // indirect google.golang.org/protobuf v1.36.11 // indirect + modernc.org/libc v1.74.1 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect + modernc.org/sqlite v1.55.0 // indirect ) diff --git a/coordinator/go.sum b/coordinator/go.sum index df0b061..6395300 100644 --- a/coordinator/go.sum +++ b/coordinator/go.sum @@ -9,6 +9,8 @@ github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XL github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= @@ -29,8 +31,12 @@ github.com/lann/builder v0.0.0-20180802200727-47ae307949d0 h1:SOEGU9fKiNWd/HOJuq github.com/lann/builder v0.0.0-20180802200727-47ae307949d0/go.mod h1:dXGbAdH5GtBTC4WfIxhKZfyBF/HBFgRZSWwZ9g/He9o= github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0 h1:P6pPBnrTSX3DEVR4fDembhRWSsG5rVo6hYhAB/ADZrk= github.com/lann/ps v0.0.0-20150810152359-62de8c46ede0/go.mod h1:vmVJ0l/dxyfGW6FmdpVm2joNMFikkuWg0EoCKLGUMNw= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/prometheus/client_golang v1.19.1 h1:wZWJDwK+NameRJuPGDhlnFgx8e8HN3XHQeLaYJFJBOE= @@ -41,6 +47,8 @@ github.com/prometheus/common v0.55.0 h1:KEi6DK7lXW/m7Ig5i47x0vRzuBsHuvJdi5ee6Y3G github.com/prometheus/common v0.55.0/go.mod h1:2SECS4xJG1kd8XF9IcM1gMX6510RAEL65zxzNImwdc8= github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI= github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= @@ -51,6 +59,7 @@ golang.org/x/crypto v0.17.0 h1:r8bRNjWL3GshPW3gkd+RpvzWrZAwPS49OmTGZ/uhM4k= golang.org/x/crypto v0.17.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= @@ -63,3 +72,11 @@ gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYs gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +modernc.org/libc v1.74.1 h1:bdR4VTKFMC4966QSNZ05XLGI/VwzVa2kTUX51Dm0riQ= +modernc.org/libc v1.74.1/go.mod h1:uH4t5bOx3G3g9Xcmj10YKlTcVISlRDwv8VoQJG9n8Os= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/sqlite v1.55.0 h1:hIFh0MCH0rGinQ/4KYb5/UbCkRkb+UP+OkLCVWa5MTM= +modernc.org/sqlite v1.55.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= diff --git a/coordinator/internal/infra/config.go b/coordinator/internal/infra/config.go index 2fe1efc..92c38b0 100644 --- a/coordinator/internal/infra/config.go +++ b/coordinator/internal/infra/config.go @@ -8,6 +8,7 @@ import ( "io/fs" "math" "os" + "path/filepath" "strconv" "time" @@ -82,6 +83,12 @@ type Config struct { // On by default so a downloaded binary provisions its own database; set // AUTO_MIGRATE=false when an operator manages migrations out of band. AutoMigrate bool + // DatabaseEngine selects the storage backend: "sqlite" (embedded, the + // single-binary default) or "postgres" (cluster deployments). The + // postgres engine requires DATABASE_URL. + DatabaseEngine string + // DBPath is the sqlite database file (engine=sqlite only). + DBPath string } // Load reads the environment and fails fast on anything required-but-missing @@ -128,8 +135,16 @@ func LoadConfig() (Config, error) { WorkerOfflineAfter: 1 * time.Minute, } - if cfg.DatabaseURL == "" { - return Config{}, fmt.Errorf("DATABASE_URL is required") + cfg.DatabaseEngine = getEnv("SCIMESH_DB", "sqlite") + switch cfg.DatabaseEngine { + case "sqlite", "postgres": + default: + return Config{}, fmt.Errorf("SCIMESH_DB must be sqlite or postgres") + } + cfg.DBPath = getEnv("SCIMESH_DB_PATH", filepath.Join(cfg.StorageDir, "scimesh.db")) + + if cfg.DatabaseEngine == "postgres" && cfg.DatabaseURL == "" { + return Config{}, fmt.Errorf("DATABASE_URL is required for the postgres engine") } if cfg.UIToken != "" && cfg.Token != "" && cfg.UIToken == cfg.Token { return Config{}, fmt.Errorf("UI_AUTH_TOKEN must differ from the worker auth token") @@ -185,7 +200,6 @@ func LoadConfig() (Config, error) { } cfg.AutoMigrate = parsed } - return cfg, nil } diff --git a/coordinator/internal/storage/sqlite/artifact_repo.go b/coordinator/internal/storage/sqlite/artifact_repo.go new file mode 100644 index 0000000..5dcbc41 --- /dev/null +++ b/coordinator/internal/storage/sqlite/artifact_repo.go @@ -0,0 +1,91 @@ +package sqlite + +import ( + "context" + "database/sql" + "errors" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// ArtifactRepo implements usecase.ArtifactRepository on SQLite. +type ArtifactRepo struct { + db *sql.DB +} + +func NewArtifactRepo(db *sql.DB) *ArtifactRepo { + return &ArtifactRepo{db: db} +} + +const artifactColumns = `id, job_id, task_id, attempt, kind, filename, storage_key, + content_type, size_bytes, sha256, created_at` + +// scanArtifact maps one row onto a domain.Artifact. +func scanArtifact(row interface{ Scan(dest ...any) error }) (*domain.Artifact, error) { + var ( + a domain.Artifact + kind string + ) + var ( + taskID sql.NullString + attempt sql.NullInt64 + createdAt sql.NullInt64 + ) + if err := row.Scan( + &a.ID, &a.JobID, &taskID, &attempt, &kind, &a.Filename, &a.StorageKey, + &a.ContentType, &a.SizeBytes, &a.SHA256, &createdAt, + ); err != nil { + return nil, err + } + a.CreatedAt = decodeTime(createdAt.Int64) + a.Kind = domain.ArtifactKind(kind) + if taskID.Valid { + if id, err := uuid.Parse(taskID.String); err == nil { + a.TaskID = &id + } + } + if attempt.Valid { + value := int(attempt.Int64) + a.Attempt = &value + } + return &a, nil +} + +func (r *ArtifactRepo) Insert(ctx context.Context, a *domain.Artifact) error { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +INSERT INTO artifacts (id, job_id, task_id, attempt, kind, filename, storage_key, + content_type, size_bytes, sha256, created_at) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + a.ID.String(), a.JobID.String(), nullableUUID(a.TaskID), nullableInt(a.Attempt), + string(a.Kind), a.Filename, a.StorageKey, a.ContentType, a.SizeBytes, a.SHA256, + encodeTime(a.CreatedAt)) + return err +} + +func (r *ArtifactRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Artifact, error) { + row := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT "+artifactColumns+" FROM artifacts WHERE id = ?", id.String()) + artifact, err := scanArtifact(row) + return artifact, mapErrNoRows(err, domain.ErrArtifactNotFound) +} + +func (r *ArtifactRepo) FindPartialResult(ctx context.Context, taskID uuid.UUID, attempt int) (*domain.Artifact, error) { + row := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT "+artifactColumns+" FROM artifacts WHERE task_id = ? AND attempt = ? AND kind = ?", + taskID.String(), attempt, string(domain.ArtifactPartialResult)) + artifact, err := scanArtifact(row) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return artifact, err +} + +// nullableInt renders a nilable int as its value, or NULL. +func nullableInt(n *int) any { + if n == nil { + return nil + } + return *n +} diff --git a/coordinator/internal/storage/sqlite/job_repo.go b/coordinator/internal/storage/sqlite/job_repo.go new file mode 100644 index 0000000..1616a54 --- /dev/null +++ b/coordinator/internal/storage/sqlite/job_repo.go @@ -0,0 +1,168 @@ +package sqlite + +import ( + "context" + "database/sql" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// JobRepo implements usecase.JobRepository on SQLite. +type JobRepo struct { + db *sql.DB +} + +func NewJobRepo(db *sql.DB) *JobRepo { + return &JobRepo{db: db} +} + +const jobColumns = `id, workload, input_uri, parameters, status, created_at, completed_at, + input_artifact_id, result_artifact_id, error_code, error_message, reducer_started_at, owner_id` + +// scanJob maps one row onto a domain.Job. Scanned values follow the sqlite +// column order exactly: ids are TEXT, parameters JSON TEXT, timestamps unix +// nanoseconds (nullable), statuses plain strings. +func scanJob(row interface{ Scan(dest ...any) error }) (*domain.Job, error) { + var ( + j domain.Job + status string + params string + ) + var ( + createdAt sql.NullInt64 + completedAt, reducerStartedAt sql.NullInt64 + inputArtifact, resultArtifact, ownerID sql.NullString + errorCode, errorMessage sql.NullString + ) + if err := row.Scan( + &j.ID, &j.Workload, &j.InputURI, ¶ms, &status, &createdAt, + &completedAt, &inputArtifact, &resultArtifact, &errorCode, &errorMessage, + &reducerStartedAt, &ownerID, + ); err != nil { + return nil, err + } + if err := decodeJSON(params, &j.Parameters); err != nil { + return nil, err + } + j.Status = domain.JobStatus(status) + j.CreatedAt = decodeTime(createdAt.Int64) + if completedAt.Valid { + value := decodeTime(completedAt.Int64) + j.CompletedAt = &value + } + if reducerStartedAt.Valid { + value := decodeTime(reducerStartedAt.Int64) + j.ReducerStartedAt = &value + } + if inputArtifact.Valid { + if id, err := uuid.Parse(inputArtifact.String); err == nil { + j.InputArtifactID = &id + } + } + if resultArtifact.Valid { + if id, err := uuid.Parse(resultArtifact.String); err == nil { + j.ResultArtifactID = &id + } + } + if ownerID.Valid { + if id, err := uuid.Parse(ownerID.String); err == nil { + j.OwnerID = &id + } + } + if errorCode.Valid { + j.ErrorCode = &errorCode.String + } + if errorMessage.Valid { + j.ErrorMessage = &errorMessage.String + } + return &j, nil +} + +// Insert runs inside the caller's transaction alongside the job's tasks. +func (r *JobRepo) Insert(ctx context.Context, j *domain.Job) error { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +INSERT INTO jobs (id, workload, input_uri, parameters, status, created_at, owner_id) +VALUES (?, ?, ?, ?, ?, ?, ?)`, + j.ID.String(), j.Workload, j.InputURI, encodeJSON(j.Parameters), string(j.Status), + encodeTime(j.CreatedAt), nullableUUID(j.OwnerID)) + return err +} + +func (r *JobRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Job, error) { + row := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT "+jobColumns+" FROM jobs WHERE id = ?", id.String()) + job, err := scanJob(row) + return job, mapErrNoRows(err, domain.ErrJobNotFound) +} + +func (r *JobRepo) ClaimReduction(ctx context.Context, id uuid.UUID, startedAt time.Time) (bool, error) { + res, err := conn(ctx, r.db).ExecContext(ctx, ` +UPDATE jobs SET reducer_started_at = ? +WHERE id = ? AND status = ? AND reducer_started_at IS NULL`, + encodeTime(startedAt), id.String(), string(domain.JobReducing)) + if err != nil { + return false, err + } + affected, err := res.RowsAffected() + return affected == 1, err +} + +func (r *JobRepo) CompleteWithResult(ctx context.Context, id, resultArtifactID uuid.UUID, completedAt time.Time) error { + res, err := conn(ctx, r.db).ExecContext(ctx, ` +UPDATE jobs SET status = ?, result_artifact_id = ?, completed_at = ?, + reducer_started_at = NULL, error_code = NULL, error_message = NULL +WHERE id = ? AND status = ?`, + string(domain.JobCompleted), resultArtifactID.String(), encodeTime(completedAt), + id.String(), string(domain.JobReducing)) + if err != nil { + return err + } + affected, err := res.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return domain.ErrJobNotFound + } + return nil +} + +func (r *JobRepo) FailReduction(ctx context.Context, id uuid.UUID, code, message string, completedAt time.Time) error { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +UPDATE jobs SET status = ?, completed_at = ?, error_code = ?, error_message = ?, + reducer_started_at = NULL +WHERE id = ? AND status = ?`, + string(domain.JobFailed), encodeTime(completedAt), code, message, + id.String(), string(domain.JobReducing)) + return err +} + +func (r *JobRepo) UpdateStatus(ctx context.Context, id uuid.UUID, + status domain.JobStatus, completedAt *time.Time) error { + + res, err := conn(ctx, r.db).ExecContext(ctx, + "UPDATE jobs SET status = ?, completed_at = ? WHERE id = ?", + string(status), encodeTimePtr(completedAt), id.String()) + if err != nil { + return err + } + affected, err := res.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return domain.ErrJobNotFound + } + return nil +} + +// nullableUUID renders a nilable UUID as its text, or NULL. +func nullableUUID(id *uuid.UUID) any { + if id == nil { + return nil + } + return id.String() +} diff --git a/coordinator/internal/storage/sqlite/migrate.go b/coordinator/internal/storage/sqlite/migrate.go new file mode 100644 index 0000000..831d411 --- /dev/null +++ b/coordinator/internal/storage/sqlite/migrate.go @@ -0,0 +1,85 @@ +package sqlite + +import ( + "context" + "database/sql" + "embed" + "fmt" + "log/slog" + "regexp" + "sort" + "strconv" +) + +//go:embed migrations/*.sql +var migrationFiles embed.FS + +var migrationNamePattern = regexp.MustCompile(`^([0-9]+)_[a-z0-9_]+\.sql$`) + +// Migrate applies every embedded migration above the current PRAGMA +// user_version watermark, each inside its own transaction. It is idempotent: +// the watermark only advances after a migration commits. +func Migrate(ctx context.Context, db *sql.DB, log *slog.Logger) error { + entries, err := migrationFiles.ReadDir("migrations") + if err != nil { + return fmt.Errorf("read embedded migrations: %w", err) + } + type file struct { + version int + name string + } + var files []file + byVersion := map[int]string{} + for _, entry := range entries { + match := migrationNamePattern.FindStringSubmatch(entry.Name()) + if match == nil { + continue + } + version, err := strconv.Atoi(match[1]) + if err != nil { + return fmt.Errorf("migration %q has an invalid version: %w", entry.Name(), err) + } + if _, duplicate := byVersion[version]; duplicate { + return fmt.Errorf("migration version %d is duplicated", version) + } + body, err := migrationFiles.ReadFile("migrations/" + entry.Name()) + if err != nil { + return fmt.Errorf("read migration %q: %w", entry.Name(), err) + } + byVersion[version] = string(body) + files = append(files, file{version: version, name: entry.Name()}) + } + if len(files) == 0 { + return fmt.Errorf("no sqlite migrations are embedded") + } + sort.Slice(files, func(i, j int) bool { return files[i].version < files[j].version }) + + var applied int + if err := db.QueryRowContext(ctx, "PRAGMA user_version").Scan(&applied); err != nil { + return fmt.Errorf("read schema version: %w", err) + } + for _, item := range files { + if item.version <= applied { + continue + } + if log != nil { + log.Info("applying sqlite migration", "version", item.version, "file", item.name) + } + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return err + } + if _, err := tx.ExecContext(ctx, byVersion[item.version]); err != nil { + _ = tx.Rollback() + return fmt.Errorf("apply migration %s: %w", item.name, err) + } + if _, err := tx.ExecContext(ctx, fmt.Sprintf("PRAGMA user_version = %d", item.version)); err != nil { + _ = tx.Rollback() + return fmt.Errorf("advance schema version after %s: %w", item.name, err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit migration %s: %w", item.name, err) + } + } + return nil +} diff --git a/coordinator/internal/storage/sqlite/migrations/0001_core.sql b/coordinator/internal/storage/sqlite/migrations/0001_core.sql new file mode 100644 index 0000000..ae6de91 --- /dev/null +++ b/coordinator/internal/storage/sqlite/migrations/0001_core.sql @@ -0,0 +1,96 @@ +-- 0001: core schema. SQLite stores enums as TEXT with CHECK constraints and +-- JSON documents as TEXT; timestamps are unix nanoseconds (INTEGER). +CREATE TABLE IF NOT EXISTS jobs ( + id TEXT PRIMARY KEY, + workload TEXT NOT NULL, + input_uri TEXT NOT NULL DEFAULT '', + parameters TEXT NOT NULL DEFAULT '{}', + status TEXT NOT NULL DEFAULT 'pending' + CHECK (status IN ('pending','running','reducing','completed','failed','cancelled')), + created_at INTEGER NOT NULL, + completed_at INTEGER, + input_artifact_id TEXT, + result_artifact_id TEXT, + error_code TEXT, + error_message TEXT, + reducer_started_at INTEGER, + owner_id TEXT +); + +CREATE TABLE IF NOT EXISTS tasks ( + id TEXT PRIMARY KEY, + job_id TEXT NOT NULL REFERENCES jobs(id) ON DELETE CASCADE, + chunk_index INTEGER NOT NULL, + workload TEXT NOT NULL, + input_uri TEXT, + input_artifact_id TEXT, + input_sha256 TEXT NOT NULL, + parameters TEXT NOT NULL DEFAULT '{}', + status TEXT NOT NULL DEFAULT 'pending' + CHECK (status IN ('pending','leased','running','completed','failed','cancelled')), + attempt INTEGER NOT NULL DEFAULT 0 CHECK (attempt >= 0), + max_attempts INTEGER NOT NULL DEFAULT 3 CHECK (max_attempts > 0), + lease_owner TEXT, + lease_expires_at INTEGER, + result_artifact_id TEXT, + metrics TEXT, + error_code TEXT, + error_message TEXT, + created_at INTEGER NOT NULL, + started_at INTEGER, + completed_at INTEGER, + version INTEGER NOT NULL DEFAULT 0, + CONSTRAINT uq_tasks_job_chunk UNIQUE (job_id, chunk_index), + CONSTRAINT ck_tasks_completed_result CHECK ( + status <> 'completed' OR (result_artifact_id IS NOT NULL) + ), + CONSTRAINT ck_tasks_leased_owner CHECK ( + status <> 'leased' OR (lease_owner IS NOT NULL AND lease_expires_at IS NOT NULL) + ) +); + +CREATE INDEX IF NOT EXISTS ix_tasks_claim ON tasks (status, lease_expires_at, created_at); +CREATE INDEX IF NOT EXISTS ix_tasks_job ON tasks (job_id); + +CREATE TABLE IF NOT EXISTS artifacts ( + id TEXT PRIMARY KEY, + job_id TEXT NOT NULL REFERENCES jobs(id) ON DELETE CASCADE, + task_id TEXT, + attempt INTEGER, + kind TEXT NOT NULL + CHECK (kind IN ('input','shard','partial_result','final_result','log')), + filename TEXT NOT NULL, + storage_key TEXT NOT NULL, + content_type TEXT NOT NULL, + size_bytes INTEGER NOT NULL DEFAULT 0, + sha256 TEXT NOT NULL DEFAULT '', + created_at INTEGER NOT NULL, + CONSTRAINT uq_partial_result_task_attempt UNIQUE (task_id, attempt) +); + +CREATE INDEX IF NOT EXISTS ix_artifacts_job ON artifacts (job_id); + +CREATE TABLE IF NOT EXISTS workers ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + capabilities TEXT NOT NULL DEFAULT '[]', + status TEXT NOT NULL DEFAULT 'online' + CHECK (status IN ('online','busy','offline')), + owner_id TEXT, + trust_level TEXT NOT NULL DEFAULT 'trusted' + CHECK (trust_level IN ('trusted','untrusted')), + last_heartbeat_at INTEGER NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS task_results ( + task_id TEXT NOT NULL REFERENCES tasks(id) ON DELETE CASCADE, + owner_id TEXT NOT NULL, + result_sha256 TEXT NOT NULL, + result_artifact_id TEXT NOT NULL, + created_at INTEGER NOT NULL DEFAULT 0, + PRIMARY KEY (task_id, owner_id) +); + +CREATE INDEX IF NOT EXISTS ix_task_results_task ON task_results (task_id); diff --git a/coordinator/internal/storage/sqlite/sqlite.go b/coordinator/internal/storage/sqlite/sqlite.go new file mode 100644 index 0000000..55a3194 --- /dev/null +++ b/coordinator/internal/storage/sqlite/sqlite.go @@ -0,0 +1,172 @@ +// Package sqlite implements the usecase repository ports on an embedded +// SQLite database. It is the single-binary storage backend: no external +// service, one file per database, pure-Go driver (modernc.org/sqlite) so the +// static release binaries stay static. +// +// Concurrency model: SQLite allows exactly one writer. Every repository write +// runs inside a TxManager transaction, and the database is opened with a +// busy_timeout, so concurrent writers serialize instead of failing. The +// postgres claim path uses FOR UPDATE SKIP LOCKED; here the same guarantee +// comes from the write lock of the surrounding transaction — ClaimNext is +// always called inside WithinTx by the usecase layer, so SELECT + UPDATE +// cannot interleave. +package sqlite + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "time" + + _ "modernc.org/sqlite" + + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +// Open opens (and creates when missing) the database file, applying WAL, +// foreign keys, and a busy timeout. Callers own the returned handle. +func Open(path string) (*sql.DB, error) { + dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(1)&_pragma=synchronous(NORMAL)", path) + db, err := sql.Open("sqlite", dsn) + if err != nil { + return nil, fmt.Errorf("open sqlite database: %w", err) + } + if err := db.PingContext(context.Background()); err != nil { + _ = db.Close() + return nil, fmt.Errorf("ping sqlite database: %w", err) + } + return db, nil +} + +// querier is satisfied by both *sql.DB and *sql.Tx, letting every repository +// method run identically inside or outside a transaction. +type querier interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) + QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row +} + +// txKey is an unexported struct type, so no other package can collide with it +// or reach the transaction we stash in the context. +type txKey struct{} + +// TxManager implements usecase.TxManager. +type TxManager struct { + db *sql.DB +} + +func NewTxManager(db *sql.DB) *TxManager { + return &TxManager{db: db} +} + +var _ usecase.TxManager = (*TxManager)(nil) + +// WithinTx runs fn inside one transaction, committing on success and rolling +// back on any error or panic. The transaction travels in the context, the +// same pattern as the postgres backend. SQLite write transactions are +// serialized by the database's single-writer lock, so a concurrent writer +// waits on the busy timeout instead of racing. +func (m *TxManager) WithinTx(ctx context.Context, fn func(ctx context.Context) error) error { + if _, ok := ctx.Value(txKey{}).(*sql.Tx); ok { + return fn(ctx) + } + tx, err := m.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + + if err := fn(context.WithValue(ctx, txKey{}, tx)); err != nil { + return err + } + return tx.Commit() +} + +// conn returns the transaction bound to ctx, or the database when there is none. +func conn(ctx context.Context, db *sql.DB) querier { + if tx, ok := ctx.Value(txKey{}).(*sql.Tx); ok { + return tx + } + return db +} + +// --- JSON and null helpers ------------------------------------------------ + +// encodeJSON stores a Go value as JSON text, defaulting to "{}". +func encodeJSON(value any) string { + if value == nil { + return "{}" + } + encoded, err := json.Marshal(value) + if err != nil { + return "{}" + } + return string(encoded) +} + +// decodeJSON reads a JSON text column into the destination. +func decodeJSON(raw any, destination any) error { + text, ok := raw.(string) + if !ok { + return nil + } + if text == "" { + return nil + } + return json.Unmarshal([]byte(text), destination) +} + +// encodeTime stores a time as unix nanoseconds (NULL for zero time). +func encodeTime(t time.Time) any { + if t.IsZero() { + return nil + } + return t.UnixNano() +} + +// encodeTimePtr stores a nilable time as unix nanoseconds. +func encodeTimePtr(t *time.Time) any { + if t == nil { + return nil + } + return t.UnixNano() +} + +// decodeTime reads a unix-nanosecond column back into a time.Time. +func decodeTime(raw any) time.Time { + switch v := raw.(type) { + case int64: + return time.Unix(0, v).UTC() + case int: + return time.Unix(0, int64(v)).UTC() + } + return time.Time{} +} + +// nullIfEmpty maps "" to SQL NULL. +func nullIfEmpty(s string) any { + if s == "" { + return nil + } + return s +} + +// mapErrNoRows translates sql.ErrNoRows into the domain not-found errors. +func mapErrNoRows(err error, notFound error) error { + if errors.Is(err, sql.ErrNoRows) { + return notFound + } + return err +} + +var ( + _ usecase.JobRepository = (*JobRepo)(nil) + _ usecase.TaskRepository = (*TaskRepo)(nil) + _ usecase.ArtifactRepository = (*ArtifactRepo)(nil) + _ usecase.WorkerRepository = (*WorkerRepo)(nil) + _ usecase.TaskResultRepository = (*TaskResultRepo)(nil) + _ usecase.UIReadRepository = (*UIReadRepo)(nil) + _ usecase.TxManager = (*TxManager)(nil) +) diff --git a/coordinator/internal/storage/sqlite/sqlite_test.go b/coordinator/internal/storage/sqlite/sqlite_test.go new file mode 100644 index 0000000..78eb679 --- /dev/null +++ b/coordinator/internal/storage/sqlite/sqlite_test.go @@ -0,0 +1,349 @@ +package sqlite + +import ( + "context" + "database/sql" + "errors" + "path/filepath" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +// newTestDB opens an isolated on-disk database and applies the migrations. +func newTestDB(t *testing.T) *sql.DB { + t.Helper() + db, err := Open(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatalf("open: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + if err := Migrate(context.Background(), db, nil); err != nil { + t.Fatalf("migrate: %v", err) + } + return db +} + +func fixedTime() time.Time { + return time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC) +} + +func seedJob(t *testing.T, db *sql.DB, n int) *domain.Job { + t.Helper() + ctx := context.Background() + tx := NewTxManager(db) + chunks := make([]domain.ChunkSpec, 0, n) + for i := 0; i < n; i++ { + chunks = append(chunks, domain.ChunkSpec{ + ChunkIndex: i, + InputURI: "s3://chunk-" + string(rune('a'+i)), + InputSHA256: "sha-" + string(rune('a'+i)), + }) + } + job, tasks, err := domain.NewJobWithTasks("similarity_search", "s3://ds", nil, chunks, fixedTime()) + if err != nil { + t.Fatal(err) + } + if err := tx.WithinTx(ctx, func(ctx context.Context) error { + if err := NewJobRepo(db).Insert(ctx, job); err != nil { + return err + } + return NewTaskRepo(db).InsertBatch(ctx, tasks) + }); err != nil { + t.Fatalf("seed: %v", err) + } + return job +} + +func TestMigrateIsIdempotent(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + if err := Migrate(ctx, db, nil); err != nil { + t.Fatalf("second migrate: %v", err) + } + var version int + if err := db.QueryRowContext(ctx, "PRAGMA user_version").Scan(&version); err != nil { + t.Fatal(err) + } + if version != 1 { + t.Errorf("user_version = %d, want 1", version) + } +} + +func TestJobRepoRoundTrip(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + job := seedJob(t, db, 1) + + got, err := NewJobRepo(db).Get(ctx, job.ID) + if err != nil { + t.Fatal(err) + } + if got.Workload != job.Workload || got.Status != domain.JobPending { + t.Errorf("job = %+v", got) + } + if err := NewJobRepo(db).UpdateStatus(ctx, job.ID, domain.JobRunning, nil); err != nil { + t.Fatal(err) + } + got, err = NewJobRepo(db).Get(ctx, job.ID) + if err != nil { + t.Fatal(err) + } + if got.Status != domain.JobRunning { + t.Errorf("status = %q, want running", got.Status) + } + if _, err := NewJobRepo(db).Get(ctx, uuid.New()); !errors.Is(err, domain.ErrJobNotFound) { + t.Errorf("missing job err = %v, want ErrJobNotFound", err) + } +} + +func TestClaimGivesEachTaskToExactlyOneWorker(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + seedJob(t, db, 3) + repo := NewTaskRepo(db) + + claimed := map[uuid.UUID]bool{} + for i := 0; i < 3; i++ { + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{ + Workloads: []string{"similarity_search"}, + Owner: "w1", + Now: fixedTime(), + LeaseUntil: fixedTime().Add(time.Minute), + }) + if err != nil { + t.Fatal(err) + } + if task == nil { + t.Fatal("claim returned nil on a non-empty queue") + } + if claimed[task.ID] { + t.Fatalf("task %s claimed twice", task.ID) + } + claimed[task.ID] = true + if task.Status != domain.TaskLeased || task.Attempt != 1 || task.LeaseOwner == nil || *task.LeaseOwner != "w1" { + t.Errorf("task = %+v", task) + } + } + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: []string{"similarity_search"}, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(time.Minute)}) + if err != nil { + t.Fatal(err) + } + if task != nil { + t.Fatal("claim must return nil on an empty queue") + } +} + +func TestUpdateRejectsStaleVersion(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + seedJob(t, db, 1) + repo := NewTaskRepo(db) + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(time.Minute)}) + if err != nil || task == nil { + t.Fatalf("claim: %v", err) + } + stale := *task + task.Status = domain.TaskRunning + task.Version++ // as a domain method would have done + if err := repo.Update(ctx, task); err != nil { + t.Fatal(err) + } + stale.Status = domain.TaskCompleted + stale.Version++ + if err := repo.Update(ctx, &stale); !errors.Is(err, domain.ErrLeaseConflict) { + t.Errorf("stale update err = %v, want ErrLeaseConflict", err) + } +} + +func TestExpireLeasesRequeuesElapsedTasks(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + seedJob(t, db, 1) + repo := NewTaskRepo(db) + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(-time.Minute)}) + if err != nil || task == nil { + t.Fatalf("claim: %v", err) + } + affected, err := repo.ExpireLeases(ctx, fixedTime()) + if err != nil { + t.Fatal(err) + } + if len(affected) != 1 { + t.Fatalf("affected = %v, want 1 job", affected) + } + task, err = repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w2", Now: fixedTime(), LeaseUntil: fixedTime().Add(time.Minute)}) + if err != nil { + t.Fatal(err) + } + if task == nil || task.Attempt != 2 { + t.Errorf("requeued task = %+v, want attempt 2", task) + } +} + +func TestExpireLeasesFailsAfterFinalAttempt(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + seedJob(t, db, 1) + repo := NewTaskRepo(db) + for attempt := 1; attempt <= 3; attempt++ { + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(-time.Minute)}) + if err != nil || task == nil { + t.Fatalf("claim %d: %v", attempt, err) + } + if _, err := repo.ExpireLeases(ctx, fixedTime()); err != nil { + t.Fatal(err) + } + } + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(time.Minute)}) + if err != nil { + t.Fatal(err) + } + if task != nil { + t.Fatal("exhausted task must not be claimable") + } +} + +func TestArtifactRepoRoundTripAndUniqueness(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + job := seedJob(t, db, 1) + repo := NewArtifactRepo(db) + attempt := 1 + artifact, err := domain.NewArtifact(job.ID, nil, domain.ArtifactShard, "shard-0.tsv", "text/tab-separated-values", fixedTime()) + if err != nil { + t.Fatal(err) + } + artifact.SetContent("abc123", 42) + if err := repo.Insert(ctx, artifact); err != nil { + t.Fatal(err) + } + got, err := repo.Get(ctx, artifact.ID) + if err != nil { + t.Fatal(err) + } + if got.SHA256 != "abc123" || got.SizeBytes != 42 { + t.Errorf("artifact = %+v", got) + } + + partialTask := uuid.New() + partial, err := domain.NewArtifact(job.ID, &partialTask, domain.ArtifactPartialResult, "p.csv", "text/csv", fixedTime()) + if err != nil { + t.Fatal(err) + } + partial.Attempt = &attempt + if err := repo.Insert(ctx, partial); err != nil { + t.Fatal(err) + } + found, err := repo.FindPartialResult(ctx, partialTask, attempt) + if err != nil || found == nil { + t.Fatalf("find partial: %v", err) + } + duplicate, err := domain.NewArtifact(job.ID, &partialTask, domain.ArtifactPartialResult, "p2.csv", "text/csv", fixedTime()) + if err != nil { + t.Fatal(err) + } + duplicate.Attempt = &attempt + if err := repo.Insert(ctx, duplicate); err == nil { + t.Fatal("duplicate partial for the same attempt must fail") + } +} + +func TestWorkerRepoRoundTripAndLiveness(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + repo := NewWorkerRepo(db) + worker, err := domain.NewWorker("w1", []string{"similarity-search"}, fixedTime()) + if err != nil { + t.Fatal(err) + } + if err := repo.Insert(ctx, worker); err != nil { + t.Fatal(err) + } + got, err := repo.Get(ctx, worker.ID) + if err != nil { + t.Fatal(err) + } + if got.Name != "w1" || len(got.Capabilities) != 1 || got.TrustLevel != domain.WorkerTrusted { + t.Errorf("worker = %+v", got) + } + if err := repo.Touch(ctx, worker.ID, fixedTime().Add(time.Hour)); err != nil { + t.Fatal(err) + } + changed, err := repo.MarkStaleOffline(ctx, fixedTime().Add(2*time.Hour)) + if err != nil { + t.Fatal(err) + } + if changed != 1 { + t.Errorf("offline changes = %d, want 1", changed) + } + got, err = repo.Get(ctx, worker.ID) + if err != nil { + t.Fatal(err) + } + if got.Status != domain.WorkerOffline { + t.Errorf("status = %q, want offline", got.Status) + } +} + +func TestTaskResultVotes(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + job := seedJob(t, db, 1) + repo := NewTaskResultRepo(db) + task, err := domain.NewTask(job.ID, 1, "similarity_search", "s3://in", "sha", nil, 3, fixedTime()) + if err != nil { + t.Fatal(err) + } + if err := NewTaskRepo(db).InsertBatch(ctx, []*domain.Task{task}); err != nil { + t.Fatal(err) + } + taskID := task.ID + artifactID := uuid.New() + ownerA, ownerB := uuid.New(), uuid.New() + if err := repo.RecordVote(ctx, taskID, ownerA, "hash", artifactID); err != nil { + t.Fatal(err) + } + if err := repo.RecordVote(ctx, taskID, ownerB, "hash", artifactID); err != nil { + t.Fatal(err) + } + if err := repo.RecordVote(ctx, taskID, ownerA, "hash2", artifactID); err != nil { + t.Fatal(err) + } + n, err := repo.CountAgreeing(ctx, taskID, "hash") + if err != nil { + t.Fatal(err) + } + if n != 1 { + t.Errorf("agreeing = %d, want 1 (owner A changed its vote)", n) + } +} + +func TestCancelByJobInvalidatesTasks(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + job := seedJob(t, db, 2) + repo := NewTaskRepo(db) + task, err := repo.ClaimNext(ctx, usecase.ClaimFilter{Workloads: nil, Owner: "w1", Now: fixedTime(), LeaseUntil: fixedTime().Add(time.Minute)}) + if err != nil || task == nil { + t.Fatalf("claim: %v", err) + } + cancelled, err := repo.CancelByJob(ctx, job.ID, fixedTime()) + if err != nil { + t.Fatal(err) + } + if cancelled != 2 { + t.Errorf("cancelled = %d, want 2", cancelled) + } + got, err := repo.Get(ctx, task.ID) + if err != nil { + t.Fatal(err) + } + if got.Status != domain.TaskCancelled || got.LeaseOwner != nil { + t.Errorf("cancelled task = %+v", got) + } +} diff --git a/coordinator/internal/storage/sqlite/stats_repo.go b/coordinator/internal/storage/sqlite/stats_repo.go new file mode 100644 index 0000000..6e108d9 --- /dev/null +++ b/coordinator/internal/storage/sqlite/stats_repo.go @@ -0,0 +1,63 @@ +package sqlite + +import ( + "context" + "database/sql" + "fmt" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// Known statuses per entity, so counts are zero-filled and every status is +// always present in the metrics (a flat 0 line beats a gap on the dashboard). +var ( + taskStatuses = []string{string(domain.TaskPending), string(domain.TaskLeased), string(domain.TaskRunning), string(domain.TaskCompleted), string(domain.TaskFailed), string(domain.TaskCancelled)} + jobStatuses = []string{string(domain.JobPending), string(domain.JobRunning), string(domain.JobReducing), string(domain.JobCompleted), string(domain.JobFailed), string(domain.JobCancelled)} + workerStatuses = []string{string(domain.WorkerOnline), string(domain.WorkerBusy), string(domain.WorkerOffline)} +) + +// StatsRepo answers the aggregate status counts the business metrics report. +type StatsRepo struct { + db *sql.DB +} + +func NewStatsRepo(db *sql.DB) *StatsRepo { + return &StatsRepo{db: db} +} + +// Counts returns status->count maps for tasks, jobs, and workers, each +// zero-filled across its known statuses. +func (r *StatsRepo) Counts(ctx context.Context) (tasks, jobs, workers map[string]int, err error) { + if tasks, err = r.countByStatus(ctx, "tasks", taskStatuses); err != nil { + return nil, nil, nil, err + } + if jobs, err = r.countByStatus(ctx, "jobs", jobStatuses); err != nil { + return nil, nil, nil, err + } + if workers, err = r.countByStatus(ctx, "workers", workerStatuses); err != nil { + return nil, nil, nil, err + } + return tasks, jobs, workers, nil +} + +func (r *StatsRepo) countByStatus(ctx context.Context, table string, known []string) (map[string]int, error) { + out := make(map[string]int, len(known)) + for _, s := range known { + out[s] = 0 // zero-fill + } + // table is a fixed internal constant, never user input — safe to format. + rows, err := r.db.QueryContext(ctx, fmt.Sprintf("SELECT status, count(*) FROM %s GROUP BY status", table)) + if err != nil { + return nil, fmt.Errorf("count %s by status: %w", table, err) + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var status string + var n int + if err := rows.Scan(&status, &n); err != nil { + return nil, err + } + out[status] = n + } + return out, rows.Err() +} diff --git a/coordinator/internal/storage/sqlite/task_repo.go b/coordinator/internal/storage/sqlite/task_repo.go new file mode 100644 index 0000000..d266921 --- /dev/null +++ b/coordinator/internal/storage/sqlite/task_repo.go @@ -0,0 +1,314 @@ +package sqlite + +import ( + "context" + "database/sql" + "errors" + "strings" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +// TaskRepo implements usecase.TaskRepository on SQLite. +type TaskRepo struct { + db *sql.DB +} + +func NewTaskRepo(db *sql.DB) *TaskRepo { + return &TaskRepo{db: db} +} + +const taskColumns = `id, job_id, chunk_index, workload, input_uri, input_artifact_id, input_sha256, + parameters, status, attempt, max_attempts, lease_owner, lease_expires_at, + result_artifact_id, metrics, error_code, error_message, + created_at, started_at, completed_at, version` + +// scanTask maps one row onto a domain.Task. +func scanTask(row interface{ Scan(dest ...any) error }) (*domain.Task, error) { + var ( + t domain.Task + status string + params string + metrics sql.NullString + ) + var ( + inputURI, leaseOwner, errorCode, errorMessage sql.NullString + inputArtifact, resultArtifact sql.NullString + leaseExpiresAt, startedAt, completedAt sql.NullInt64 + createdAt sql.NullInt64 + ) + if err := row.Scan( + &t.ID, &t.JobID, &t.ChunkIndex, &t.Workload, &inputURI, &inputArtifact, &t.InputSHA256, + ¶ms, &status, &t.Attempt, &t.MaxAttempts, &leaseOwner, &leaseExpiresAt, + &resultArtifact, &metrics, &errorCode, &errorMessage, + &createdAt, &startedAt, &completedAt, &t.Version, + ); err != nil { + return nil, err + } + if err := decodeJSON(params, &t.Parameters); err != nil { + return nil, err + } + if metrics.Valid && metrics.String != "" { + if err := decodeJSON(metrics.String, &t.Metrics); err != nil { + return nil, err + } + } + t.Status = domain.TaskStatus(status) + t.CreatedAt = decodeTime(createdAt.Int64) + if inputURI.Valid { + t.InputURI = inputURI.String + } + if leaseOwner.Valid { + t.LeaseOwner = &leaseOwner.String + } + if inputArtifact.Valid { + if id, err := uuid.Parse(inputArtifact.String); err == nil { + t.InputArtifactID = &id + } + } + if resultArtifact.Valid { + if id, err := uuid.Parse(resultArtifact.String); err == nil { + t.ResultArtifactID = &id + } + } + if leaseExpiresAt.Valid { + value := decodeTime(leaseExpiresAt.Int64) + t.LeaseExpiresAt = &value + } + if startedAt.Valid { + value := decodeTime(startedAt.Int64) + t.StartedAt = &value + } + if completedAt.Valid { + value := decodeTime(completedAt.Int64) + t.CompletedAt = &value + } + if errorCode.Valid { + t.ErrorCode = &errorCode.String + } + if errorMessage.Valid { + t.ErrorMessage = &errorMessage.String + } + return &t, nil +} + +// ClaimNext atomically leases the next eligible task. SQLite has no SKIP +// LOCKED: the guarantee comes from the surrounding transaction's write lock — +// the usecase layer always calls ClaimNext inside WithinTx, so SELECT + UPDATE +// cannot interleave with another claimant. +func (r *TaskRepo) ClaimNext(ctx context.Context, f usecase.ClaimFilter) (*domain.Task, error) { + workloadClause := "" + workloadArgs := []any{} + if len(f.Workloads) > 0 { + placeholders := make([]string, 0, len(f.Workloads)) + for _, w := range f.Workloads { + placeholders = append(placeholders, "?") + workloadArgs = append(workloadArgs, w) + } + workloadClause = " AND workload IN (" + strings.Join(placeholders, ", ") + ")" + } + voterClause := "" + var voterArg any + if f.VoterOwner != nil { + voterClause = " AND NOT EXISTS (SELECT 1 FROM task_results tr WHERE tr.task_id = tasks.id AND tr.owner_id = ?)" + voterArg = f.VoterOwner.String() + } + args := append([]any{f.Owner, f.LeaseUntil.UnixNano(), f.Now.UnixNano()}, workloadArgs...) + if f.VoterOwner != nil { + args = append(args, voterArg) + } + query := ` +UPDATE tasks SET + status = 'leased', + attempt = attempt + 1, + lease_owner = ?, + lease_expires_at = ?, + started_at = COALESCE(started_at, ?), + version = version + 1 +WHERE id IN ( + SELECT id FROM tasks + WHERE status = 'pending' AND attempt < max_attempts` + workloadClause + voterClause + ` + ORDER BY created_at, chunk_index + LIMIT 1 +) +RETURNING ` + taskColumns + + row := conn(ctx, r.db).QueryRowContext(ctx, query, args...) + task, err := scanTask(row) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, err + } + return task, nil +} + +func (r *TaskRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Task, error) { + row := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT "+taskColumns+" FROM tasks WHERE id = ?", id.String()) + task, err := scanTask(row) + return task, mapErrNoRows(err, domain.ErrTaskNotFound) +} + +// GetForUpdate reads a task. SQLite serializes writers inside a transaction, +// so no row lock is needed: the surrounding write transaction already isolates +// the read-modify-write sequence. +func (r *TaskRepo) GetForUpdate(ctx context.Context, id uuid.UUID) (*domain.Task, error) { + return r.Get(ctx, id) +} + +// Update writes the mutated entity back under optimistic concurrency. +func (r *TaskRepo) Update(ctx context.Context, t *domain.Task) error { + res, err := conn(ctx, r.db).ExecContext(ctx, ` +UPDATE tasks SET + status = ?, attempt = ?, lease_owner = ?, lease_expires_at = ?, + result_artifact_id = ?, metrics = ?, error_code = ?, error_message = ?, + started_at = ?, completed_at = ?, version = ? +WHERE id = ? AND version = ?`, + string(t.Status), t.Attempt, nullableString(t.LeaseOwner), encodeTimePtr(t.LeaseExpiresAt), + nullableUUID(t.ResultArtifactID), nullableMetrics(t.Metrics), + nullableString(t.ErrorCode), nullableString(t.ErrorMessage), + encodeTimePtr(t.StartedAt), encodeTimePtr(t.CompletedAt), t.Version, + t.ID.String(), t.Version-1) + if err != nil { + return err + } + affected, err := res.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return domain.ErrLeaseConflict + } + return nil +} + +func (r *TaskRepo) InsertBatch(ctx context.Context, tasks []*domain.Task) error { + for _, t := range tasks { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +INSERT INTO tasks (id, job_id, chunk_index, workload, input_uri, input_artifact_id, + input_sha256, parameters, status, attempt, max_attempts, created_at, version) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + t.ID.String(), t.JobID.String(), t.ChunkIndex, t.Workload, + nullIfEmpty(t.InputURI), nullableUUID(t.InputArtifactID), + t.InputSHA256, encodeJSON(t.Parameters), string(t.Status), + t.Attempt, t.MaxAttempts, encodeTime(t.CreatedAt), t.Version) + if err != nil { + return err + } + } + return nil +} + +func (r *TaskRepo) ListCompleted(ctx context.Context, jobID uuid.UUID) ([]*domain.Task, error) { + rows, err := conn(ctx, r.db).QueryContext(ctx, + "SELECT "+taskColumns+" FROM tasks WHERE job_id = ? AND status = ? ORDER BY chunk_index", + jobID.String(), string(domain.TaskCompleted)) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + var tasks []*domain.Task + for rows.Next() { + task, err := scanTask(rows) + if err != nil { + return nil, err + } + tasks = append(tasks, task) + } + return tasks, rows.Err() +} + +func (r *TaskRepo) CountByStatus(ctx context.Context, jobID uuid.UUID) (map[domain.TaskStatus]int, error) { + rows, err := conn(ctx, r.db).QueryContext(ctx, + "SELECT status, count(*) FROM tasks WHERE job_id = ? GROUP BY status", jobID.String()) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + counts := make(map[domain.TaskStatus]int) + for rows.Next() { + var ( + status string + n int + ) + if err := rows.Scan(&status, &n); err != nil { + return nil, err + } + counts[domain.TaskStatus(status)] = n + } + return counts, rows.Err() +} + +func (r *TaskRepo) CancelByJob(ctx context.Context, jobID uuid.UUID, now time.Time) (int64, error) { + res, err := conn(ctx, r.db).ExecContext(ctx, ` +UPDATE tasks SET + status = 'cancelled', + lease_owner = NULL, + lease_expires_at = NULL, + error_code = NULL, + error_message = NULL, + completed_at = ?, + version = version + 1 +WHERE job_id = ? AND status IN ('pending','leased','running')`, + now.UnixNano(), jobID.String()) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// ExpireLeases applies the lease-expiry rule set-based and returns the +// distinct jobs whose aggregate status may have changed. +func (r *TaskRepo) ExpireLeases(ctx context.Context, now time.Time) ([]uuid.UUID, error) { + rows, err := conn(ctx, r.db).QueryContext(ctx, ` +UPDATE tasks SET + status = CASE WHEN attempt < max_attempts THEN 'pending' ELSE 'failed' END, + lease_owner = NULL, + lease_expires_at = NULL, + error_code = CASE WHEN attempt >= max_attempts THEN ? ELSE error_code END, + error_message = CASE WHEN attempt >= max_attempts THEN ? ELSE error_message END, + completed_at = CASE WHEN attempt >= max_attempts THEN ? ELSE completed_at END, + version = version + 1 +WHERE status IN ('leased','running') AND lease_expires_at < ? +RETURNING job_id`, + domain.ErrCodeLeaseExpired, "lease expired after the final attempt", now.UnixNano(), now.UnixNano()) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + var affected []uuid.UUID + seen := map[uuid.UUID]bool{} + for rows.Next() { + var raw string + if err := rows.Scan(&raw); err != nil { + return nil, err + } + if id, err := uuid.Parse(raw); err == nil && !seen[id] { + seen[id] = true + affected = append(affected, id) + } + } + return affected, rows.Err() +} + +// nullableString renders a nilable string, or NULL. +func nullableString(s *string) any { + if s == nil { + return nil + } + return *s +} + +// nullableMetrics stores nil metrics as NULL, else JSON text. +func nullableMetrics(m map[string]any) any { + if m == nil { + return nil + } + return encodeJSON(m) +} diff --git a/coordinator/internal/storage/sqlite/task_result_repo.go b/coordinator/internal/storage/sqlite/task_result_repo.go new file mode 100644 index 0000000..6318754 --- /dev/null +++ b/coordinator/internal/storage/sqlite/task_result_repo.go @@ -0,0 +1,41 @@ +package sqlite + +import ( + "context" + "database/sql" + "time" + + "github.com/google/uuid" +) + +// TaskResultRepo records and tallies quorum votes for untrusted task results. +type TaskResultRepo struct { + db *sql.DB +} + +func NewTaskResultRepo(db *sql.DB) *TaskResultRepo { + return &TaskResultRepo{db: db} +} + +// RecordVote stores (or replaces) one owner's vote for a task's result. +func (r *TaskResultRepo) RecordVote(ctx context.Context, taskID, ownerID uuid.UUID, sha256 string, artifactID uuid.UUID) error { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +INSERT INTO task_results (task_id, owner_id, result_sha256, result_artifact_id, created_at) +VALUES (?, ?, ?, ?, ?) +ON CONFLICT (task_id, owner_id) DO UPDATE SET + result_sha256 = excluded.result_sha256, + result_artifact_id = excluded.result_artifact_id, + created_at = excluded.created_at`, + taskID.String(), ownerID.String(), sha256, artifactID.String(), time.Now().UnixNano()) + return err +} + +// CountAgreeing returns how many distinct owners have voted for the given +// result hash on this task. +func (r *TaskResultRepo) CountAgreeing(ctx context.Context, taskID uuid.UUID, sha256 string) (int, error) { + var n int + err := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT count(DISTINCT owner_id) FROM task_results WHERE task_id = ? AND result_sha256 = ?", + taskID.String(), sha256).Scan(&n) + return n, err +} diff --git a/coordinator/internal/storage/sqlite/ui_read_repo.go b/coordinator/internal/storage/sqlite/ui_read_repo.go new file mode 100644 index 0000000..5806321 --- /dev/null +++ b/coordinator/internal/storage/sqlite/ui_read_repo.go @@ -0,0 +1,145 @@ +package sqlite + +import ( + "context" + "database/sql" + "fmt" + "strings" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// UIReadRepo contains bounded, deterministic read queries for the operator UI. +type UIReadRepo struct{ db *sql.DB } + +func NewUIReadRepo(db *sql.DB) *UIReadRepo { return &UIReadRepo{db: db} } + +func (r *UIReadRepo) GetJob(ctx context.Context, id uuid.UUID) (*domain.Job, error) { + return NewJobRepo(r.db).Get(ctx, id) +} + +func (r *UIReadRepo) ListJobs(ctx context.Context, owner *uuid.UUID, limit int) ([]domain.Job, error) { + if limit < 1 || limit > 100 { + return nil, domain.ErrInvalidInput + } + query := "SELECT " + jobColumns + " FROM jobs" + args := []any{} + if owner != nil { + query += " WHERE owner_id = ?" + args = append(args, owner.String()) + } + query += " ORDER BY created_at DESC, id DESC LIMIT ?" + args = append(args, limit) + rows, err := conn(ctx, r.db).QueryContext(ctx, query, args...) + if err != nil { + return nil, fmt.Errorf("list jobs: %w", err) + } + defer func() { _ = rows.Close() }() + jobs := make([]domain.Job, 0) + for rows.Next() { + job, err := scanJob(rows) + if err != nil { + return nil, err + } + jobs = append(jobs, *job) + } + return jobs, rows.Err() +} + +func (r *UIReadRepo) ListTasksByJob(ctx context.Context, jobID uuid.UUID) ([]domain.Task, error) { + rows, err := conn(ctx, r.db).QueryContext(ctx, + "SELECT "+taskColumns+" FROM tasks WHERE job_id = ? ORDER BY chunk_index ASC", jobID.String()) + if err != nil { + return nil, fmt.Errorf("list tasks: %w", err) + } + defer func() { _ = rows.Close() }() + tasks := make([]domain.Task, 0) + for rows.Next() { + task, err := scanTask(rows) + if err != nil { + return nil, err + } + tasks = append(tasks, *task) + } + return tasks, rows.Err() +} + +func (r *UIReadRepo) ListTasksByJobs(ctx context.Context, jobIDs []uuid.UUID) (map[uuid.UUID][]domain.Task, error) { + out := make(map[uuid.UUID][]domain.Task, len(jobIDs)) + if len(jobIDs) == 0 { + return out, nil + } + placeholders := make([]string, 0, len(jobIDs)) + args := make([]any, 0, len(jobIDs)) + for _, id := range jobIDs { + placeholders = append(placeholders, "?") + args = append(args, id.String()) + } + rows, err := conn(ctx, r.db).QueryContext(ctx, + "SELECT "+taskColumns+" FROM tasks WHERE job_id IN ("+strings.Join(placeholders, ", ")+") ORDER BY job_id ASC, chunk_index ASC", + args...) + if err != nil { + return nil, fmt.Errorf("list tasks by jobs: %w", err) + } + defer func() { _ = rows.Close() }() + for rows.Next() { + task, err := scanTask(rows) + if err != nil { + return nil, err + } + out[task.JobID] = append(out[task.JobID], *task) + } + return out, rows.Err() +} + +func (r *UIReadRepo) ListWorkers(ctx context.Context, limit int) ([]domain.Worker, error) { + return r.listWorkers(ctx, "", nil, limit) +} + +func (r *UIReadRepo) ListWorkersByOwner(ctx context.Context, owner uuid.UUID, limit int) ([]domain.Worker, error) { + return r.listWorkers(ctx, " WHERE owner_id = ?", []any{owner.String()}, limit) +} + +func (r *UIReadRepo) listWorkers(ctx context.Context, clause string, args []any, limit int) ([]domain.Worker, error) { + if limit < 1 || limit > 100 { + return nil, domain.ErrInvalidInput + } + query := "SELECT " + workerColumns + " FROM workers" + clause + + " ORDER BY last_heartbeat_at DESC, id DESC LIMIT ?" + fullArgs := append(args, limit) + rows, err := conn(ctx, r.db).QueryContext(ctx, query, fullArgs...) + if err != nil { + return nil, fmt.Errorf("list workers: %w", err) + } + defer func() { _ = rows.Close() }() + workers := make([]domain.Worker, 0) + for rows.Next() { + worker, err := scanWorker(rows) + if err != nil { + return nil, err + } + workers = append(workers, *worker) + } + return workers, rows.Err() +} + +func (r *UIReadRepo) ListArtifactsByJob(ctx context.Context, jobID uuid.UUID) ([]domain.Artifact, error) { + rows, err := conn(ctx, r.db).QueryContext(ctx, + "SELECT "+artifactColumns+" FROM artifacts WHERE job_id = ? ORDER BY created_at ASC, id ASC", + jobID.String()) + if err != nil { + return nil, fmt.Errorf("list artifacts: %w", err) + } + defer func() { _ = rows.Close() }() + artifacts := make([]domain.Artifact, 0) + for rows.Next() { + artifact, err := scanArtifact(rows) + if err != nil { + return nil, err + } + artifacts = append(artifacts, *artifact) + } + return artifacts, rows.Err() +} diff --git a/coordinator/internal/storage/sqlite/worker_repo.go b/coordinator/internal/storage/sqlite/worker_repo.go new file mode 100644 index 0000000..4d0bf44 --- /dev/null +++ b/coordinator/internal/storage/sqlite/worker_repo.go @@ -0,0 +1,92 @@ +package sqlite + +import ( + "context" + "database/sql" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// WorkerRepo implements usecase.WorkerRepository on SQLite. +type WorkerRepo struct { + db *sql.DB +} + +func NewWorkerRepo(db *sql.DB) *WorkerRepo { + return &WorkerRepo{db: db} +} + +const workerColumns = `id, name, capabilities, status, owner_id, trust_level, last_heartbeat_at, created_at, updated_at` + +// scanWorker maps one row onto a domain.Worker. +func scanWorker(row interface{ Scan(dest ...any) error }) (*domain.Worker, error) { + var ( + w domain.Worker + status string + trust string + caps string + ) + var ( + ownerID sql.NullString + lastHeartbeat, created, updated sql.NullInt64 + ) + if err := row.Scan( + &w.ID, &w.Name, &caps, &status, &ownerID, &trust, + &lastHeartbeat, &created, &updated, + ); err != nil { + return nil, err + } + w.LastHeartbeatAt = decodeTime(lastHeartbeat.Int64) + w.CreatedAt = decodeTime(created.Int64) + w.UpdatedAt = decodeTime(updated.Int64) + if err := decodeJSON(caps, &w.Capabilities); err != nil { + return nil, err + } + w.Status = domain.WorkerStatus(status) + w.TrustLevel = domain.WorkerTrust(trust) + if ownerID.Valid { + if id, err := uuid.Parse(ownerID.String); err == nil { + w.OwnerID = &id + } + } + return &w, nil +} + +func (r *WorkerRepo) Insert(ctx context.Context, w *domain.Worker) error { + _, err := conn(ctx, r.db).ExecContext(ctx, ` +INSERT INTO workers (id, name, capabilities, status, owner_id, trust_level, + last_heartbeat_at, created_at, updated_at) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + w.ID.String(), w.Name, encodeJSON(w.Capabilities), string(w.Status), + nullableUUID(w.OwnerID), string(w.TrustLevel), + encodeTime(w.LastHeartbeatAt), encodeTime(w.CreatedAt), encodeTime(w.UpdatedAt)) + return err +} + +func (r *WorkerRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Worker, error) { + row := conn(ctx, r.db).QueryRowContext(ctx, + "SELECT "+workerColumns+" FROM workers WHERE id = ?", id.String()) + worker, err := scanWorker(row) + return worker, mapErrNoRows(err, domain.ErrWorkerNotFound) +} + +func (r *WorkerRepo) Touch(ctx context.Context, id uuid.UUID, at time.Time) error { + _, err := conn(ctx, r.db).ExecContext(ctx, + "UPDATE workers SET last_heartbeat_at = ?, status = ?, updated_at = ? WHERE id = ?", + encodeTime(at), string(domain.WorkerOnline), encodeTime(at), id.String()) + return err +} + +func (r *WorkerRepo) MarkStaleOffline(ctx context.Context, cutoff time.Time) (int64, error) { + res, err := conn(ctx, r.db).ExecContext(ctx, + "UPDATE workers SET status = ?, updated_at = ? WHERE last_heartbeat_at < ? AND status <> ?", + string(domain.WorkerOffline), encodeTime(cutoff), encodeTime(cutoff), + string(domain.WorkerOffline)) + if err != nil { + return 0, err + } + return res.RowsAffected() +}