305 lines
14 KiB
Go
305 lines
14 KiB
Go
package pi
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"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 --wait --wait-timeout 45 --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 --wait --wait-timeout 45 --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 failure == "preflight" && result.Phase == PhaseRolledBack { t.Fatalf("preflight failure unexpectedly rolled back: %+v", result) }
|
|
if failure == "build" && result.Phase != PhaseRolledBack { t.Fatalf("candidate build failure must compensate: %+v", result) }
|
|
if failure == "preflight" { 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, "http://127.0.0.1:8787/sessions?scope=all")
|
|
assertCalled(t, fake.calls, "/internal/maintenance/activate")
|
|
assertCalled(t, fake.calls, "/internal/maintenance/deactivate")
|
|
}
|
|
|
|
func TestRollbackRestoresInterruptedOrPreviouslyRecordedState(t *testing.T) {
|
|
fake := newFakeRunner()
|
|
statePath := filepath.Join(t.TempDir(), "state.json")
|
|
configured, err := renderedCore(context.Background(), fake)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
previous, err := runningImage(context.Background(), fake, "thothii-core:local")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
previous.ConfigurationSHA = configured.ConfigurationSHA
|
|
writeStateForTest(t, statePath, State{Version: 1, Phase: PhaseRecreated, Previous: previous})
|
|
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 --wait --wait-timeout 45 --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"}, MountFingerprint: mountFingerprint(nil)}})
|
|
_, 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")
|
|
}
|
|
|
|
func TestRunningImageCapturesServerBindAndNamedMountIdentity(t *testing.T) {
|
|
fake := newFakeRunner()
|
|
fake.mountsJSON = `[{"Type":"bind","Source":"/srv/thothii/data","Destination":"/data","RW":true},{"Type":"bind","Source":"/srv/thothii/pi","Destination":"/home/thoth/.pi","RW":true},{"Type":"volume","Name":"sessions","Source":"/var/lib/docker/volumes/sessions/_data","Destination":"/data/sessions","RW":true}]`
|
|
image, err := runningImage(context.Background(), fake, "thothii-core:local")
|
|
if err != nil {
|
|
t.Fatalf("runningImage() error = %v", err)
|
|
}
|
|
if len(image.Mounts) != 3 || image.Mounts[0].Type != "bind" || image.Mounts[0].Destination != "/data" {
|
|
t.Fatalf("mounts = %#v", image.Mounts)
|
|
}
|
|
if strings.Contains(fmt.Sprint(image), "/srv/thothii") || image.Mounts[0].SourceSHA256 == "" || image.MountFingerprint == "" {
|
|
t.Fatalf("mount contract leaked a server source or lacks a safe fingerprint: %#v", image)
|
|
}
|
|
}
|
|
|
|
func TestCanonicalDigestReferenceRejectsCredentialsAndURLForms(t *testing.T) {
|
|
valid := "registry.example.invalid/thothii-core@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
|
|
if got, err := canonicalDigestReference(valid); err != nil || got != valid { t.Fatalf("canonicalDigestReference() = %q, %v", got, err) }
|
|
for _, invalid := range []string{"https://registry.example.invalid/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "user:pass@registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", "registry.example/a@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa?token=x"} {
|
|
if _, err := canonicalDigestReference(invalid); err == nil { t.Fatalf("accepted unsafe reference %q", invalid) }
|
|
}
|
|
}
|
|
|
|
type fakeRunner struct {
|
|
calls []string
|
|
fail string
|
|
version string
|
|
activeSessions bool
|
|
built bool
|
|
currentImage string
|
|
volumes []string
|
|
mountsJSON string
|
|
}
|
|
|
|
func newFakeRunner() *fakeRunner {
|
|
return &fakeRunner{version: "0.80.3", currentImage: "sha256:old", 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 strings.Contains(call, "image tag sha256:old") {
|
|
f.fail = ""
|
|
f.currentImage = "sha256:old"
|
|
}
|
|
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" && f.built && strings.Contains(call, "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: f.currentImage + "\n"}, nil
|
|
case strings.Contains(call, "inspect --format {{json .Mounts}}"):
|
|
if f.mountsJSON != "" {
|
|
return compose.Result{Stdout: f.mountsJSON}, nil
|
|
}
|
|
return compose.Result{Stdout: `[{"Type":"volume","Name":"settings","Source":"settings","Destination":"/data/settings","RW":true},{"Type":"volume","Name":"pi-state","Source":"pi-state","Destination":"/home/thoth/.pi","RW":true},{"Type":"volume","Name":"sessions","Source":"sessions","Destination":"/data/sessions","RW":true},{"Type":"volume","Name":"workspace-registry","Source":"workspace-registry","Destination":"/data/workspace-registry","RW":true}]`}, nil
|
|
case strings.Contains(call, "/internal/maintenance/activate"):
|
|
return compose.Result{Stdout: `{"active":true,"admissions":0}`}, nil
|
|
case strings.Contains(call, "/internal/maintenance/deactivate"):
|
|
return compose.Result{Stdout: `{"active":false,"admissions":0}`}, nil
|
|
case strings.Contains(call, "/sessions?scope=all"):
|
|
if f.activeSessions {
|
|
f.activeSessions = false
|
|
return compose.Result{Stdout: `{"sessions":[{"status":"open","archived":false}]}`}, nil
|
|
}
|
|
return compose.Result{Stdout: `{"sessions":[]}`}, nil
|
|
case strings.Contains(call, "compose build"):
|
|
f.built = true
|
|
f.version = "0.81.0"
|
|
f.currentImage = "sha256:candidate"
|
|
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)
|
|
}
|
|
}
|