package pi import ( "context" "errors" "io" "os" "path/filepath" "strings" "testing" "github.com/aritmolab/thothii/tools/thothctl/internal/compose" ) func TestUpdateBuildsPinnedVersionRecreatesOnlyCoreAndPersistsRecoveryState(t *testing.T) { fake := newFakeRunner() dir := t.TempDir() result, err := Update(context.Background(), fake, Request{ StatePath: filepath.Join(dir, ".thothctl", "update-state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true, }) if err != nil { t.Fatalf("Update() error = %v", err) } if result.Phase != PhaseVerified { t.Fatalf("phase = %q, want %q", result.Phase, PhaseVerified) } assertCalled(t, fake.calls, "compose build --pull --build-arg PI_VERSION=0.81.0 core") assertCalled(t, fake.calls, "compose up --detach --no-deps --force-recreate core") assertNotCalled(t, fake.calls, "frontend") if got := string(readStateBytes(t, result.StatePath)); strings.Contains(got, "secret") || !strings.Contains(got, `"phase": "verified"`) { t.Fatalf("state = %q, want credential-free verified metadata", got) } if got := string(readStateBytes(t, result.StatePath)); strings.Contains(got, "llm.example.invalid") { t.Fatalf("state = %q, want an endpoint-free configuration digest", got) } } func TestUpdatePullsOnlyDigestPinnedSource(t *testing.T) { fake := newFakeRunner() digest := "registry.example.invalid/thothii-core@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" _, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: PullSource, Image: digest, Confirm: true}) if err == nil { t.Fatal("Update() error = nil, want Pi version verification failure from unchanged fake image") } assertCalled(t, fake.calls, "pull "+digest) assertCalled(t, fake.calls, "image tag "+digest+" thothii-core:local") fake = newFakeRunner() _, err = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: PullSource, Image: "registry.example.invalid/thothii-core:latest", Confirm: true}) if err == nil || !strings.Contains(err.Error(), "immutable sha256 digest") { t.Fatalf("Update() error = %v, want digest-pinning rejection", err) } assertNotCalled(t, fake.calls, "compose") } func TestUpdateIsNoOpWhenDesiredVersionAlreadyRuns(t *testing.T) { fake := newFakeRunner() fake.version = "0.80.3" result, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.80.3", Source: BuildSource, Confirm: true}) if err != nil { t.Fatalf("Update() error = %v", err) } if result.Phase != PhaseNoop { t.Fatalf("phase = %q, want %q", result.Phase, PhaseNoop) } assertNotCalled(t, fake.calls, "compose build") } func TestUpdateRollsBackAfterPostRecreateFailures(t *testing.T) { for _, failure := range []string{"health", "version", "smoke"} { t.Run(failure, func(t *testing.T) { fake := newFakeRunner() fake.fail = failure statePath := filepath.Join(t.TempDir(), "state.json") result, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true}) if err == nil { t.Fatal("Update() error = nil, want verification failure") } if result.Phase != PhaseRolledBack { t.Fatalf("phase = %q, want %q", result.Phase, PhaseRolledBack) } assertCalled(t, fake.calls, "image tag sha256:old thothii-core:local") assertCalled(t, fake.calls, "compose up --detach --no-deps --force-recreate core") if strings.Join(fake.volumes, ",") != "settings,pi-state,sessions,workspace-registry" { t.Fatalf("volumes changed: %v", fake.volumes) } if got := string(readStateBytes(t, statePath)); !strings.Contains(got, `"phase": "rolled_back"`) { t.Fatalf("state = %q, want rollback metadata", got) } }) } } func TestUpdateDoesNotRecreateWhenPreflightOrBuildFails(t *testing.T) { for _, failure := range []string{"preflight", "build"} { t.Run(failure, func(t *testing.T) { fake := newFakeRunner() fake.fail = failure result, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true}) if err == nil { t.Fatal("Update() error = nil, want failure") } if result.Phase == PhaseRolledBack { t.Fatalf("pre-recreate failure unexpectedly rolled back: %+v", result) } assertNotCalled(t, fake.calls, "force-recreate") }) } } func TestUpdateRequiresConfirmationAndDrainsActiveSessions(t *testing.T) { fake := newFakeRunner() _, err := Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource}) if !errors.Is(err, ErrConfirmationRequired) { t.Fatalf("Update() error = %v, want confirmation error", err) } assertNotCalled(t, fake.calls, "compose") fake = newFakeRunner() fake.activeSessions = true _, err = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true}) if !errors.Is(err, ErrActiveSessions) { t.Fatalf("Update() error = %v, want active-sessions error", err) } fake = newFakeRunner() fake.activeSessions = true _, err = Update(context.Background(), fake, Request{StatePath: filepath.Join(t.TempDir(), "state.json"), Version: "0.81.0", Source: BuildSource, Confirm: true, Drain: true}) if err != nil { t.Fatalf("Update() with drain error = %v", err) } assertCalled(t, fake.calls, "compose exec -T core tht session list --json") } func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) { fake := newFakeRunner() statePath := filepath.Join(t.TempDir(), "state.json") writeStateForTest(t, statePath, State{Version: 1, Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}}}) result, err := Rollback(context.Background(), fake, statePath, true) if err != nil { t.Fatalf("Rollback() error = %v", err) } if result.Phase != PhaseRolledBack { t.Fatalf("phase = %q, want %q", result.Phase, PhaseRolledBack) } assertCalled(t, fake.calls, "image tag sha256:old thothii-core:local") assertCalled(t, fake.calls, "compose up --detach --no-deps --force-recreate core") } func TestUpdateRefusesToOverwriteInterruptedRecoveryState(t *testing.T) { fake := newFakeRunner() statePath := filepath.Join(t.TempDir(), "state.json") writeStateForTest(t, statePath, State{Phase: PhaseRecreated, Previous: Image{ID: "sha256:old", Reference: "thothii-core:local", Volumes: []string{"settings"}}}) _, err := Update(context.Background(), fake, Request{StatePath: statePath, Version: "0.81.0", Source: BuildSource, Confirm: true}) if !errors.Is(err, ErrInterruptedUpdate) { t.Fatalf("Update() error = %v, want interrupted update error", err) } assertNotCalled(t, fake.calls, "compose") } type fakeRunner struct { calls []string fail string version string activeSessions bool built bool volumes []string } func newFakeRunner() *fakeRunner { return &fakeRunner{version: "0.80.3", volumes: []string{"settings", "pi-state", "sessions", "workspace-registry"}} } func (f *fakeRunner) Run(_ context.Context, args []string, _ io.Reader) (compose.Result, error) { call := strings.Join(args, " ") f.calls = append(f.calls, call) if f.fail == "preflight" && strings.Contains(call, "config --format json") { return compose.Result{ExitCode: 1}, errors.New("provider token=secret") } if f.fail == "build" && strings.Contains(call, "compose build") { return compose.Result{ExitCode: 1}, errors.New("build token=secret") } if f.fail == "health" && f.built && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/health") { return compose.Result{ExitCode: 1}, errors.New("health token=secret") } if f.fail == "version" && f.built && strings.Contains(call, "pi --version") && strings.Contains(call, "exec") { return compose.Result{ExitCode: 1}, errors.New("version token=secret") } if f.fail == "smoke" && strings.Contains(call, "curl -fsS http://127.0.0.1:8787/models") { return compose.Result{ExitCode: 1}, errors.New("smoke token=secret") } switch { case strings.Contains(call, "config --format json"): return compose.Result{Stdout: `{"services":{"core":{"image":"thothii-core:local","environment":{"THT_LLM_URL":"https://llm.example.invalid"}}}}`}, nil case strings.Contains(call, "ps -q core"): return compose.Result{Stdout: "core-container\n"}, nil case strings.Contains(call, "inspect --format {{.Image}}"): return compose.Result{Stdout: "sha256:old\n"}, nil case strings.Contains(call, "inspect --format {{range .Mounts}}"): return compose.Result{Stdout: strings.Join(f.volumes, "\n") + "\n"}, nil case strings.Contains(call, "tht session list --json"): if f.activeSessions { f.activeSessions = false return compose.Result{Stdout: `[{"status":"open","archived":false}]`}, nil } return compose.Result{Stdout: "[]"}, nil case strings.Contains(call, "compose build"): f.built = true f.version = "0.81.0" return compose.Result{}, nil case strings.Contains(call, "pi --version"): return compose.Result{Stdout: f.version + "\n"}, nil case strings.Contains(call, "/models"): return compose.Result{Stdout: `{"models":[{"id":"model","provider":"provider"}]}`}, nil case strings.Contains(call, "/settings"): return compose.Result{Stdout: `{"provider":"provider","model":"model","thinking":"medium"}`}, nil case strings.Contains(call, "/health"): return compose.Result{Stdout: `{"status":"ok"}`}, nil } return compose.Result{}, nil } func assertCalled(t *testing.T, calls []string, want string) { t.Helper() for _, call := range calls { if strings.Contains(call, want) { return } } t.Fatalf("calls %v did not include %q", calls, want) } func assertNotCalled(t *testing.T, calls []string, prohibited string) { t.Helper() for _, call := range calls { if strings.Contains(call, prohibited) { t.Fatalf("calls %v unexpectedly included %q", calls, prohibited) } } } func readStateBytes(t *testing.T, path string) []byte { t.Helper() contents, err := os.ReadFile(path) if err != nil { t.Fatal(err) } return contents } func writeStateForTest(t *testing.T, path string, state State) { t.Helper() if err := writeState(path, state); err != nil { t.Fatal(err) } }