package database import ( "testing" "proto-gen/internal/models" ) func newTestDB(t *testing.T) *DB { t.Helper() db, err := New(t.TempDir() + "/test.db") if err != nil { t.Fatalf("creating test db: %v", err) } t.Cleanup(func() { db.Close() }) return db } func TestNew_CreatesTables(t *testing.T) { db := newTestDB(t) var count int err := db.conn.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name IN ('tasks', 'toolchain_versions')`).Scan(&count) if err != nil { t.Fatalf("querying tables: %v", err) } if count != 2 { t.Errorf("expected 2 tables, got %d", count) } } func TestCreateAndGetTask(t *testing.T) { db := newTestDB(t) task := &models.Task{ ID: "test-1", Status: models.StatusPending, Language: "go", ProtoRepo: "yoresee_doc", ProtoBranch: "main", TargetRepo: "yoresee_doc-gen-go", ToolchainImage: "proto-gen-go:v1.0", } if err := db.CreateTask(task); err != nil { t.Fatalf("CreateTask: %v", err) } got, err := db.GetTask("test-1") if err != nil { t.Fatalf("GetTask: %v", err) } if got.Language != "go" { t.Errorf("Language = %q, want %q", got.Language, "go") } if got.Status != models.StatusPending { t.Errorf("Status = %q, want %q", got.Status, models.StatusPending) } } func TestHasRunningTask(t *testing.T) { db := newTestDB(t) has, err := db.HasRunningTask("yoresee_doc-gen-go") if err != nil { t.Fatalf("HasRunningTask: %v", err) } if has { t.Error("expected no running task") } task := &models.Task{ ID: "test-2", Status: models.StatusRunning, Language: "go", ProtoRepo: "yoresee_doc", ProtoBranch: "main", TargetRepo: "yoresee_doc-gen-go", ToolchainImage: "proto-gen-go:v1.0", } if err := db.CreateTask(task); err != nil { t.Fatalf("CreateTask: %v", err) } has, err = db.HasRunningTask("yoresee_doc-gen-go") if err != nil { t.Fatalf("HasRunningTask: %v", err) } if !has { t.Error("expected running task") } } func TestListTasks(t *testing.T) { db := newTestDB(t) for i := 0; i < 3; i++ { status := models.StatusPending if i == 1 { status = models.StatusRunning } db.CreateTask(&models.Task{ ID: "t-" + string(rune('0'+i)), Status: status, Language: "go", ProtoRepo: "yoresee_doc", ProtoBranch: "main", TargetRepo: "repo-" + string(rune('0'+i)), ToolchainImage: "img:v1", }) } tasks, err := db.ListTasks("", 10, 0) if err != nil { t.Fatalf("ListTasks: %v", err) } if len(tasks) != 3 { t.Errorf("len(tasks) = %d, want 3", len(tasks)) } tasks, err = db.ListTasks(string(models.StatusRunning), 10, 0) if err != nil { t.Fatalf("ListTasks filtered: %v", err) } if len(tasks) != 1 { t.Errorf("len(running) = %d, want 1", len(tasks)) } }