From 1f45449491299cd2333ac95ec62574f32b9b4791 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 15:36:18 -0400 Subject: [PATCH 1/9] feat(model): share a llama-server run profile between plan and start Plan and start now use one axis.model-run/v1 profile. With no new flags the argv stays llama-server, -m, the weights path, --port, the port, and --host 127.0.0.1. A plan-default port refuses until --port is passed. --n-gpu-layers is a string and requires measured VRAM on one discrete device. --- cmd/axis/model.go | 96 +++--- cmd/axis/model_run_profile.go | 130 ++++++++ cmd/axis/model_run_profile_test.go | 212 +++++++++++++ internal/modellife/plan.go | 188 +++++++++--- internal/modellife/plan_profile_test.go | 209 +++++++++++++ internal/modelplan/profile_test.go | 104 +++++++ internal/modelplan/single_node.go | 19 +- internal/models/model_operation.go | 5 + internal/models/run_profile.go | 377 ++++++++++++++++++++++++ internal/models/run_profile_test.go | 160 ++++++++++ 10 files changed, 1423 insertions(+), 77 deletions(-) create mode 100644 cmd/axis/model_run_profile.go create mode 100644 cmd/axis/model_run_profile_test.go create mode 100644 internal/modellife/plan_profile_test.go create mode 100644 internal/modelplan/profile_test.go create mode 100644 internal/models/run_profile.go create mode 100644 internal/models/run_profile_test.go diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 284e3a74..30bce8c2 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -74,7 +74,7 @@ func modelCmd() *cobra.Command { } func modelPlanCmd() *cobra.Command { - var cacheAddr, format string + var cacheAddr, format, writeProfile string var port int var live bool cmd := &cobra.Command{ @@ -93,33 +93,42 @@ func modelPlanCmd() *cobra.Command { cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().BoolVar(&live, "live", false, "Bypass daemon cache and perform live fleet discovery") cmd.Flags().StringVar(&format, "format", "text", "Output format: text, json, or yaml") + cmd.Flags().StringVar(&writeProfile, "write-profile", "", "Write the selected axis.model-run/v1 profile to this path") return cmd } func modelStartCmd() *cobra.Command { - var node, weights, cacheAddr, format string - var port int + var node, weights, cacheAddr, format, fromPlan, nGPULayers string + var port, ctxSize, batchSize, ubatchSize, threads int var live bool cmd := &cobra.Command{ Use: "start", Short: "Start llama-server on a named node (explicit port and weights)", SilenceUsage: true, - PreRunE: validateOutputFormat(&format, "text", "json", "yaml"), + PreRunE: func(cmd *cobra.Command, args []string) error { + if err := validateOutputFormat(&format, "text", "json", "yaml")(cmd, args); err != nil { + return err + } + return requireModelStartIdentity(cmd) + }, RunE: func(cmd *cobra.Command, args []string) error { ctx, cancel := context.WithTimeout(cmd.Context(), 45*time.Second) defer cancel() return runModelStart(ctx, cmd, node, weights, port, defaultModelRunner) }, } - cmd.Flags().StringVar(&node, "node", "", "Cluster node name (required)") - cmd.Flags().StringVar(&weights, "weights", "", "GGUF path on a named local volume (required)") - cmd.Flags().IntVar(&port, "port", 0, "Listen port (required; no default)") + cmd.Flags().StringVar(&node, "node", "", "Cluster node name (required unless --from-plan supplies it)") + cmd.Flags().StringVar(&weights, "weights", "", "GGUF path on a named local volume (required unless --from-plan supplies it)") + cmd.Flags().IntVar(&port, "port", 0, "Listen port (required unless --from-plan has an explicit port)") + cmd.Flags().StringVar(&fromPlan, "from-plan", "", "Read an axis.model-run/v1 profile JSON file") + cmd.Flags().StringVar(&nGPULayers, "n-gpu-layers", "", "llama-server -ngl value: an integer >= 1, auto, or all") + cmd.Flags().IntVar(&ctxSize, "ctx-size", 0, "llama-server context length (-c); omitted when unset") + cmd.Flags().IntVar(&batchSize, "batch-size", 0, "llama-server logical batch size (-b); omitted when unset") + cmd.Flags().IntVar(&ubatchSize, "ubatch-size", 0, "llama-server physical batch size (-ub); omitted when unset") + cmd.Flags().IntVar(&threads, "threads", 0, "llama-server threads (-t); must be within observed CPU cores") cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().BoolVar(&live, "live", false, "Bypass daemon cache and perform live fleet discovery") cmd.Flags().StringVar(&format, "format", "text", "Start operation receipt format: text, json, or yaml") - _ = cmd.MarkFlagRequired("node") - _ = cmd.MarkFlagRequired("weights") - _ = cmd.MarkFlagRequired("port") return cmd } @@ -504,6 +513,12 @@ func runModelPlan(ctx context.Context, cmd *cobra.Command, specOrWeights string, return err } plan.SnapshotSource = source + if cmd.Flags().Changed("port") && plan.Selected != nil { + plan.Selected.PortSource = models.PortSourceExplicit + } + if err := writeSelectedRunProfile(cmd, plan.Selected); err != nil { + return err + } if format == "json" || format == "yaml" { if writeErr := printOutput(cmd.OutOrStdout(), plan, format); writeErr != nil { @@ -617,18 +632,25 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st if format == "" { format = "text" } + profile, err := profileForModelStart(cmd, nodeName, weights, port) + if err != nil { + return err + } snap, source, err := loadModelCommandSnapshot(ctx, live, cacheAddr, "start", false) if err != nil { return err } + if profile.SnapshotPublicationID == "" && snap.Publication != nil { + profile.SnapshotPublicationID = snap.Publication.ID + } - nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, nodeName) + nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, profile.Node) if err != nil { return err } for _, res := range nf.ResidentModels { - if res.Port == port { + if res.Port == profile.Port { receipt := models.ModelOperationReceipt{ Schema: "axis.model-operation/v1", ID: models.GenerateID("mo"), @@ -636,13 +658,13 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st Status: models.ModelOperationRejected, Disposition: "port_occupied", Node: nf.Name, - Engine: "llama.cpp", - Port: port, + Engine: profile.Engine, + Port: profile.Port, SnapshotSource: source, SnapshotAt: snap.Timestamp, StartedAt: startedAt, CompletedAt: time.Now().UTC(), - Error: fmt.Sprintf("port %d already occupied by resident model %q (%s)", port, res.Name, res.Runtime), + Error: fmt.Sprintf("port %d already occupied by resident model %q (%s)", profile.Port, res.Name, res.Runtime), } if snap.Publication != nil { receipt.PublicationID = snap.Publication.ID @@ -650,12 +672,12 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st _ = writeModelStartReceipt(cmd, receipt, format) return ExitCodeError{ Code: ExitErrCommandFail, - Message: fmt.Sprintf("refusing to start model on %s:%d: %s", nf.Name, port, receipt.Error), + Message: fmt.Sprintf("refusing to start model on %s:%d: %s", nf.Name, profile.Port, receipt.Error), } } } - plan, err := modellife.PlanStart(nf, weights, port) + plan, err := modellife.PlanStartProfile(nf, profile) if err != nil { return err } @@ -671,22 +693,27 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st } receipt := models.ModelOperationReceipt{ - Schema: "axis.model-operation/v1", - ID: models.GenerateID("mo"), - Action: models.ModelOperationStart, - Status: models.ModelOperationCompleted, - Disposition: "started", - Node: plan.Node, - Engine: "llama.cpp", - Port: plan.Port, - Model: path.Base(plan.Weights), - Weights: plan.Weights, - Volume: plan.Volume, - Executable: executable, - SnapshotSource: source, - SnapshotAt: snap.Timestamp, - StartedAt: startedAt, - CompletedAt: time.Now().UTC(), + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStart, + Status: models.ModelOperationCompleted, + Disposition: "started", + Node: plan.Node, + Engine: plan.Profile.Engine, + Port: plan.Port, + Model: path.Base(plan.Weights), + Weights: plan.Weights, + Volume: plan.Volume, + Executable: executable, + SnapshotSource: source, + SnapshotAt: snap.Timestamp, + StartedAt: startedAt, + CompletedAt: time.Now().UTC(), + SpecSource: plan.Profile.SpecSource, + DeviceKind: plan.Profile.DeviceKind, + DeviceIndex: plan.Profile.DeviceIndex, + VRAMFreeMeasured: plan.Profile.VRAMFreeMeasured, + PortSource: plan.Profile.PortSource, } if snap.Publication != nil { receipt.PublicationID = snap.Publication.ID @@ -1363,6 +1390,9 @@ func resolveModelNodeFromSnapshot(snap *models.ClusterSnapshot, name string) (mo type liveModelRunner struct{} func (liveModelRunner) Start(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig, plan modellife.StartPlan) error { + if err := modellife.ExecArgvMatchesProfile(plan); err != nil { + return err + } if len(plan.Argv) == 0 { return fmt.Errorf("empty argv") } diff --git a/cmd/axis/model_run_profile.go b/cmd/axis/model_run_profile.go new file mode 100644 index 00000000..124fd859 --- /dev/null +++ b/cmd/axis/model_run_profile.go @@ -0,0 +1,130 @@ +package main + +import ( + "encoding/json" + "fmt" + "os" + "strings" + + "github.com/spf13/cobra" + + "github.com/toasterbook88/axis/internal/models" +) + +func requireModelStartIdentity(cmd *cobra.Command) error { + if cmd.Flags().Changed("from-plan") { + return nil + } + var missing []string + for _, name := range []string{"node", "weights", "port"} { + if !cmd.Flags().Changed(name) { + missing = append(missing, name) + } + } + if len(missing) == 0 { + return nil + } + return fmt.Errorf(`required flag(s) "%s" not set`, strings.Join(missing, `", "`)) +} + +func profileForModelStart(cmd *cobra.Command, nodeName, weights string, port int) (models.ModelRunProfile, error) { + fromPlan, err := cmd.Flags().GetString("from-plan") + if err != nil { + fromPlan = "" + } + var profile models.ModelRunProfile + if strings.TrimSpace(fromPlan) != "" { + data, readErr := os.ReadFile(fromPlan) + if readErr != nil { + return models.ModelRunProfile{}, readErr + } + profile, err = models.LoadModelRunProfile(data) + if err != nil { + return models.ModelRunProfile{}, err + } + if cmd.Flags().Changed("node") { + profile.Node = nodeName + } + if cmd.Flags().Changed("weights") { + profile.WeightsPath = weights + profile.Volume = "" + } + if cmd.Flags().Changed("port") { + profile.Port = port + profile.PortSource = models.PortSourceExplicit + } + } else { + profile = models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: nodeName, + Engine: models.EngineLlamaCpp, + ToolName: models.ToolLlamaServer, + ArtifactKind: models.ArtifactWeightsPath, + WeightsPath: weights, + BindHost: "127.0.0.1", + Port: port, + PortSource: models.PortSourceExplicit, + } + } + if err := applyChangedStartFlags(cmd, &profile); err != nil { + return models.ModelRunProfile{}, err + } + return profile, nil +} + +func applyChangedStartFlags(cmd *cobra.Command, profile *models.ModelRunProfile) error { + if cmd.Flags().Changed("n-gpu-layers") { + raw, err := cmd.Flags().GetString("n-gpu-layers") + if err != nil { + return err + } + n, mode, err := models.ParseNGPULayers(raw) + if err != nil { + return err + } + profile.NGPULayers = n + profile.NGPULayersMode = mode + } + if cmd.Flags().Changed("ctx-size") { + value, err := cmd.Flags().GetInt("ctx-size") + if err != nil { + return err + } + profile.ContextTokens = &value + } + if cmd.Flags().Changed("batch-size") { + value, err := cmd.Flags().GetInt("batch-size") + if err != nil { + return err + } + profile.BatchSize = &value + } + if cmd.Flags().Changed("ubatch-size") { + value, err := cmd.Flags().GetInt("ubatch-size") + if err != nil { + return err + } + profile.UBatchSize = &value + } + if cmd.Flags().Changed("threads") { + value, err := cmd.Flags().GetInt("threads") + if err != nil { + return err + } + profile.Threads = &value + } + return nil +} + +func writeSelectedRunProfile(cmd *cobra.Command, selected *models.ModelRunProfile) error { + path, err := cmd.Flags().GetString("write-profile") + if err != nil || strings.TrimSpace(path) == "" || selected == nil { + return nil + } + data, err := json.MarshalIndent(selected, "", " ") + if err != nil { + return err + } + data = append(data, '\n') + return os.WriteFile(path, data, 0644) +} diff --git a/cmd/axis/model_run_profile_test.go b/cmd/axis/model_run_profile_test.go new file mode 100644 index 00000000..1e3fe36f --- /dev/null +++ b/cmd/axis/model_run_profile_test.go @@ -0,0 +1,212 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/config" + "github.com/toasterbook88/axis/internal/models" +) + +func TestModelStartMissingFlagsStillRequiredWithoutFromPlan(t *testing.T) { + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--format", "text"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), `required flag(s) "node"`) { + t.Fatalf("err=%v", err) + } +} + +func TestModelStartFromPlanRefusesPlanDefaultUntilPortChanged(t *testing.T) { + snap := testSnap() + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + path := writeProfile(t, models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: "storage", + Engine: "llama.cpp", + ToolName: "llama-server", + ArtifactKind: "weights-path", + WeightsPath: "/mnt/models/a.gguf", + BindHost: "127.0.0.1", + Port: 8080, + PortSource: models.PortSourcePlanDefault, + }) + runner := &fakeModelRunner{} + prev := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prev }) + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--from-plan", path, "--format", "text"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "plan-default") { + t.Fatalf("err=%v", err) + } + if len(runner.started) != 0 { + t.Fatalf("runner started=%v", runner.started) + } + + cmd = modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"--from-plan", path, "--port", "8082", "--format", "text"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + want := []string{"/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8082", "--host", "127.0.0.1"} + if len(runner.started) != 1 || !reflect.DeepEqual(runner.started[0], want) { + t.Fatalf("argv=%#v", runner.started) + } +} + +func TestModelStartFromPlanRejectsPlacementDocument(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + path := filepath.Join(t.TempDir(), "plan.json") + if err := os.WriteFile(path, []byte(`{"schema":"axis.model-plan/v1","best_candidate":"storage"}`), 0o644); err != nil { + t.Fatal(err) + } + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--from-plan", path}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "expected axis.model-run/v1") { + t.Fatalf("err=%v", err) + } +} + +func TestModelStartStringNGPULayersAndCtxWithoutFit(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Resources.CPUCores = 8 + snap.Nodes[0].Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prev := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prev }) + + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{ + "--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", + "--n-gpu-layers", "auto", "--ctx-size", "2048", "--format", "text", + }) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + want := []string{ + "/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1", + "-c", "2048", "-ngl", "auto", + } + if !reflect.DeepEqual(runner.started[0], want) { + t.Fatalf("argv=%#v", runner.started) + } + + runner.started = nil + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{ + "--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", + "--n-gpu-layers", "0", + }) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "n-gpu-layers") { + t.Fatalf("zero layers err=%v", err) + } + if len(runner.started) != 0 { + t.Fatalf("zero layers started=%v", runner.started) + } +} + +func TestModelPlanWriteProfileAndExplicitPort(t *testing.T) { + snap := &models.ClusterSnapshot{ + Publication: &models.PublicationEnvelope{ID: "pub-write"}, + Nodes: []models.NodeFacts{{ + Name: "gpu-worker", + Status: models.StatusComplete, + Tools: []models.ToolInfo{{Name: "llama-server", Path: "/usr/local/bin/llama-server"}}, + Resources: &models.Resources{ + RAMFreeMB: 32000, RAMTotalMB: 64000, + Volumes: []models.Volume{{Mount: "/data/models", Kind: "local"}}, + GPUs: []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 16000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }}, + }, + DiskWeights: []models.DiskWeight{{ + Name: "qwen2.5-7b", Path: "/data/models/qwen2.5-7b.gguf", + Bytes: 4 * 1024 * 1024 * 1024, Format: "gguf", + }}, + }}, + } + stubModelSnapshot(t, snap) + out := filepath.Join(t.TempDir(), "profile.json") + cmd := modelPlanCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"qwen2.5-7b", "--port", "9001", "--format", "json", "--write-profile", out}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + var planJSON map[string]any + if err := json.Unmarshal(buf.Bytes(), &planJSON); err != nil { + t.Fatal(err) + } + if planJSON["schema"] != "axis.model-plan/v1" || planJSON["best_candidate"] != "gpu-worker" { + t.Fatalf("plan=%s", buf.String()) + } + raw, err := os.ReadFile(out) + if err != nil { + t.Fatal(err) + } + profile, err := models.LoadModelRunProfile(raw) + if err != nil { + t.Fatalf("profile: %v\n%s", err, raw) + } + if profile.Schema != models.ModelRunSchema || profile.Port != 9001 || profile.PortSource != models.PortSourceExplicit { + t.Fatalf("profile=%+v", profile) + } + if strings.Contains(string(raw), "best_candidate") || strings.Contains(string(raw), "axis.model-plan/v1") { + t.Fatalf("write-profile must be the profile only:\n%s", raw) + } +} + +func TestRunModelStartDefaultPathStillUsesFunctionArgs(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + if err := runModelStart(context.Background(), cmd, "storage", "/mnt/models/a.gguf", 8081, runner); err != nil { + t.Fatal(err) + } + want := []string{"/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1"} + if !reflect.DeepEqual(runner.started[0], want) { + t.Fatalf("argv=%#v", runner.started) + } +} + +func writeProfile(t *testing.T, profile models.ModelRunProfile) string { + t.Helper() + data, err := json.Marshal(profile) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(t.TempDir(), "profile.json") + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + return path +} diff --git a/internal/modellife/plan.go b/internal/modellife/plan.go index 49b93b58..db7a830c 100644 --- a/internal/modellife/plan.go +++ b/internal/modellife/plan.go @@ -3,50 +3,184 @@ package modellife import ( "fmt" "path" + "reflect" + "strconv" "strings" "github.com/toasterbook88/axis/internal/models" ) // StartPlan is the argv Axis will exec. It does not launch anything. +// Argv is the projection of Profile. type StartPlan struct { Node string Port int Weights string Volume string Argv []string + Profile models.ModelRunProfile } // PlanStart validates weights sit on a named local volume and that // llama-server is an observed tool. Port must be explicit and valid. +// Optional launch fields stay unset, so the argv is the historical default. func PlanStart(node models.NodeFacts, weights string, port int) (StartPlan, error) { weights = path.Clean(strings.TrimSpace(weights)) - if port < 1 || port > 65535 { - return StartPlan{}, fmt.Errorf("port must be between 1 and 65535") + return PlanStartProfile(node, models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: node.Name, + Engine: models.EngineLlamaCpp, + ToolName: models.ToolLlamaServer, + ArtifactKind: models.ArtifactWeightsPath, + WeightsPath: weights, + BindHost: "127.0.0.1", + Port: port, + PortSource: models.PortSourceExplicit, + }) +} + +// PlanStartProfile validates profile against the observed node and derives argv. +// A non-empty refusal list, a plan-default port, or an offload without measured +// discrete VRAM returns an error and no argv. +func PlanStartProfile(node models.NodeFacts, profile models.ModelRunProfile) (StartPlan, error) { + profile = normalizeStartProfile(node, profile) + if err := profile.Validate(); err != nil { + return StartPlan{}, err } - if weights == "" || weights == "." { - return StartPlan{}, fmt.Errorf("weights path is required") + refusals := append([]string{}, profile.Refusals...) + if profile.PortSource != models.PortSourceExplicit { + refusals = append(refusals, fmt.Sprintf("port source %s requires an explicit port", profile.PortSource)) } - if !hasTool(node, "llama-server") { - return StartPlan{}, fmt.Errorf("node %s has no observed llama-server tool", node.Name) + if !hasTool(node, models.ToolLlamaServer) { + refusals = append(refusals, fmt.Sprintf("node %s has no observed llama-server tool", node.Name)) } - vol, ok := namedLocalVolume(node, weights) - if !ok { - return StartPlan{}, fmt.Errorf("weights %s are not on a named local volume", weights) + if vol, ok := namedLocalVolume(node, profile.WeightsPath); ok { + profile.Volume = vol + } else { + refusals = append(refusals, fmt.Sprintf("weights %s are not on a named local volume", profile.WeightsPath)) } - bin := toolPath(node, "llama-server") - if bin == "" { - bin = "llama-server" + if profile.Threads != nil { + cores := 0 + if node.Resources != nil { + cores = node.Resources.CPUCores + } + if cores <= 0 || *profile.Threads > cores { + refusals = append(refusals, fmt.Sprintf("threads must be between 1 and %d observed cpu cores", cores)) + } + } + if len(refusals) > 0 { + return StartPlan{}, fmt.Errorf("%s", strings.Join(refusals, "; ")) + } + argv, err := ArgvFromProfile(profile) + if err != nil { + return StartPlan{}, err } return StartPlan{ - Node: node.Name, - Port: port, - Weights: weights, - Volume: vol, - Argv: []string{bin, "-m", weights, "--port", fmt.Sprintf("%d", port), "--host", "127.0.0.1"}, + Node: profile.Node, + Port: profile.Port, + Weights: path.Clean(strings.TrimSpace(profile.WeightsPath)), + Volume: profile.Volume, + Argv: argv, + Profile: profile, }, nil } +func normalizeStartProfile(node models.NodeFacts, profile models.ModelRunProfile) models.ModelRunProfile { + if profile.Schema == "" { + profile.Schema = models.ModelRunSchema + } + if profile.Node == "" { + profile.Node = node.Name + } + if profile.Engine == "" { + profile.Engine = models.EngineLlamaCpp + } + if profile.ToolName == "" { + profile.ToolName = models.ToolLlamaServer + } + if profile.ArtifactKind == "" { + profile.ArtifactKind = models.ArtifactWeightsPath + } + if profile.BindHost == "" { + profile.BindHost = "127.0.0.1" + } + profile.WeightsPath = path.Clean(strings.TrimSpace(profile.WeightsPath)) + if hasTool(node, models.ToolLlamaServer) { + if bin := toolPath(node, models.ToolLlamaServer); bin != "" { + profile.EngineBinary = bin + } else { + profile.EngineBinary = models.ToolLlamaServer + } + } + dev := models.ObserveLaunchDevice(node) + profile.DeviceKind = dev.Kind + profile.DeviceModel = dev.Model + profile.Accelerator = dev.Accelerator + profile.MemoryTopology = dev.MemoryTopology + profile.VRAMFreeMB = dev.VRAMFreeMB + profile.VRAMFreeMeasured = dev.VRAMFreeMeasured + profile.Refusals = append([]string{}, profile.Refusals...) + if (profile.NGPULayers != nil || profile.NGPULayersMode != "") && + (dev.Kind != models.DeviceKindDiscrete || !dev.VRAMFreeMeasured) { + profile.Refusals = append(profile.Refusals, "n-gpu-layers requires measured free VRAM on one discrete device") + } + return profile +} + +// ArgvFromProfile projects a llama-server argv. Optional flags are appended +// after the fixed prefix, in the order ctx, n-gpu-layers, batch, ubatch, threads. +func ArgvFromProfile(profile models.ModelRunProfile) ([]string, error) { + if profile.Engine != models.EngineLlamaCpp { + return nil, fmt.Errorf("engine %q is not supported", profile.Engine) + } + if profile.BindHost != "127.0.0.1" { + return nil, fmt.Errorf("bind host must be 127.0.0.1") + } + if profile.Port < 1 || profile.Port > 65535 { + return nil, fmt.Errorf("port must be between 1 and 65535") + } + if profile.NGPULayers != nil && profile.NGPULayersMode != "" { + return nil, fmt.Errorf("n-gpu-layers accepts an integer or a mode, not both") + } + bin := profile.EngineBinary + if bin == "" { + bin = models.ToolLlamaServer + } + weights := path.Clean(strings.TrimSpace(profile.WeightsPath)) + argv := []string{bin, "-m", weights, "--port", strconv.Itoa(profile.Port), "--host", "127.0.0.1"} + if profile.ContextTokens != nil { + argv = append(argv, "-c", strconv.Itoa(*profile.ContextTokens)) + } + if profile.NGPULayers != nil { + argv = append(argv, "-ngl", strconv.Itoa(*profile.NGPULayers)) + } else if profile.NGPULayersMode != "" { + argv = append(argv, "-ngl", profile.NGPULayersMode) + } + if profile.BatchSize != nil { + argv = append(argv, "-b", strconv.Itoa(*profile.BatchSize)) + } + if profile.UBatchSize != nil { + argv = append(argv, "-ub", strconv.Itoa(*profile.UBatchSize)) + } + if profile.Threads != nil { + argv = append(argv, "-t", strconv.Itoa(*profile.Threads)) + } + return argv, nil +} + +// ExecArgvMatchesProfile reports whether argv is exactly the profile projection. +// A hand-built argv with an empty profile does not match. +func ExecArgvMatchesProfile(plan StartPlan) error { + want, err := ArgvFromProfile(plan.Profile) + if err != nil { + return fmt.Errorf("argv does not match profile: %w", err) + } + if !reflect.DeepEqual(plan.Argv, want) { + return fmt.Errorf("argv does not match profile") + } + return nil +} + func hasTool(node models.NodeFacts, name string) bool { for _, t := range node.Tools { if strings.EqualFold(t.Name, name) { @@ -66,23 +200,5 @@ func toolPath(node models.NodeFacts, name string) string { } func namedLocalVolume(node models.NodeFacts, weights string) (string, bool) { - if node.Resources == nil { - return "", false - } - best := "" - for _, v := range node.Resources.Volumes { - if v.Kind == "network" || v.Mount == "" { - continue - } - mount := path.Clean(v.Mount) - if weights == mount || strings.HasPrefix(weights, mount+"/") { - if len(mount) > len(best) { - best = mount - } - } - } - if best == "" { - return "", false - } - return best, true + return models.NamedLocalVolume(node, weights) } diff --git a/internal/modellife/plan_profile_test.go b/internal/modellife/plan_profile_test.go new file mode 100644 index 00000000..ae341c90 --- /dev/null +++ b/internal/modellife/plan_profile_test.go @@ -0,0 +1,209 @@ +package modellife + +import ( + "reflect" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/models" +) + +func TestPlanStartDefaultArgvIsExact(t *testing.T) { + plan, err := PlanStart(storageNode(), "/mnt/models/a.gguf", 8081) + if err != nil { + t.Fatal(err) + } + want := []string{"/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1"} + if !reflect.DeepEqual(plan.Argv, want) { + t.Fatalf("argv=%#v", plan.Argv) + } + if plan.Profile.Engine != "llama.cpp" || plan.Profile.PortSource != models.PortSourceExplicit || plan.Profile.BindHost != "127.0.0.1" { + t.Fatalf("profile=%+v", plan.Profile) + } + if plan.Profile.NGPULayers != nil || plan.Profile.ContextTokens != nil { + t.Fatalf("default profile must omit optional flags: %+v", plan.Profile) + } +} + +func TestPlanStartProfileAppendsOptionalFlagsAfterBase(t *testing.T) { + node := storageNode() + node.Resources.CPUCores = 8 + node.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + ctx, ngl, batch, ubatch, threads := 4096, 20, 512, 256, 4 + profile := models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: node.Name, + Engine: "llama.cpp", + ToolName: "llama-server", + ArtifactKind: "weights-path", + WeightsPath: "/mnt/models/a.gguf", + BindHost: "127.0.0.1", + Port: 8081, + PortSource: models.PortSourceExplicit, + ContextTokens: &ctx, + NGPULayers: &ngl, + BatchSize: &batch, + UBatchSize: &ubatch, + Threads: &threads, + } + plan, err := PlanStartProfile(node, profile) + if err != nil { + t.Fatal(err) + } + want := []string{ + "/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1", + "-c", "4096", "-ngl", "20", "-b", "512", "-ub", "256", "-t", "4", + } + if !reflect.DeepEqual(plan.Argv, want) { + t.Fatalf("argv=%#v", plan.Argv) + } + if !reflect.DeepEqual(plan.Argv, mustArgv(t, plan.Profile)) { + t.Fatal("stored profile does not project the argv that will exec") + } +} + +func TestPlanStartProfileNGPULayersModeAndRefusals(t *testing.T) { + node := storageNode() + node.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + profile := readyProfile(node) + profile.NGPULayersMode = "all" + plan, err := PlanStartProfile(node, profile) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(plan.Argv[len(plan.Argv)-2:], []string{"-ngl", "all"}) { + t.Fatalf("argv=%#v", plan.Argv) + } + + unmeasured := storageNode() + unmeasured.Resources.GPUs = []models.GPUInfo{{Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, Capabilities: []string{"cuda"}}} + profile = readyProfile(unmeasured) + profile.NGPULayersMode = "auto" + if _, err := PlanStartProfile(unmeasured, profile); err == nil || !strings.Contains(err.Error(), "measured free VRAM") { + t.Fatalf("unmeasured err=%v", err) + } + + unified := storageNode() + unified.Resources.MemoryTopology = models.MemoryTopologyUnified + unified.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + profile = readyProfile(unified) + profile.NGPULayersMode = "auto" + if _, err := PlanStartProfile(unified, profile); err == nil || !strings.Contains(err.Error(), "measured free VRAM") { + t.Fatalf("unified err=%v", err) + } + + cpu := storageNode() + profile = readyProfile(cpu) + profile.NGPULayersMode = "auto" + if _, err := PlanStartProfile(cpu, profile); err == nil || !strings.Contains(err.Error(), "measured free VRAM") { + t.Fatalf("cpu err=%v", err) + } +} + +func TestPlanStartProfileCtxSizeDoesNotCheckVRAM(t *testing.T) { + node := storageNode() + ctx := 100000 + profile := readyProfile(node) + profile.ContextTokens = &ctx + plan, err := PlanStartProfile(node, profile) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(plan.Argv[len(plan.Argv)-2:], []string{"-c", "100000"}) { + t.Fatalf("argv=%#v", plan.Argv) + } +} + +func TestPlanStartProfileRefusesPlanDefaultPortAndStaleRefusals(t *testing.T) { + node := storageNode() + profile := readyProfile(node) + profile.PortSource = models.PortSourcePlanDefault + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "plan-default") { + t.Fatalf("port source err=%v", err) + } + + profile = readyProfile(node) + profile.Refusals = []string{"weights are not on a named local volume"} + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "named local volume") { + t.Fatalf("refusals err=%v", err) + } +} + +func TestPlanStartProfileRefusesRangeAndBindHost(t *testing.T) { + node := storageNode() + node.Resources.CPUCores = 4 + profile := readyProfile(node) + profile.BindHost = "0.0.0.0" + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "127.0.0.1") { + t.Fatalf("bind err=%v", err) + } + + threads := 8 + profile = readyProfile(node) + profile.Threads = &threads + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "threads") { + t.Fatalf("threads err=%v", err) + } + + zero := 0 + profile = readyProfile(node) + profile.BatchSize = &zero + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "batch") { + t.Fatalf("batch err=%v", err) + } + profile = readyProfile(node) + profile.UBatchSize = &zero + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "ubatch") { + t.Fatalf("ubatch err=%v", err) + } + + both := 4 + profile = readyProfile(node) + profile.NGPULayers = &both + profile.NGPULayersMode = "all" + if _, err := PlanStartProfile(node, profile); err == nil { + t.Fatal("expected refusal when integer and mode are both set") + } +} + +func TestExecArgvRefusesHandMadeArgv(t *testing.T) { + err := ExecArgvMatchesProfile(StartPlan{ + Argv: []string{"llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1"}, + Port: 8081, + }) + if err == nil || !strings.Contains(err.Error(), "profile") { + t.Fatalf("err=%v", err) + } +} + +func readyProfile(node models.NodeFacts) models.ModelRunProfile { + return models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: node.Name, + Engine: "llama.cpp", + ToolName: "llama-server", + ArtifactKind: "weights-path", + WeightsPath: "/mnt/models/a.gguf", + BindHost: "127.0.0.1", + Port: 8081, + PortSource: models.PortSourceExplicit, + } +} + +func mustArgv(t *testing.T, profile models.ModelRunProfile) []string { + t.Helper() + argv, err := ArgvFromProfile(profile) + if err != nil { + t.Fatal(err) + } + return argv +} diff --git a/internal/modelplan/profile_test.go b/internal/modelplan/profile_test.go new file mode 100644 index 00000000..ab74848d --- /dev/null +++ b/internal/modelplan/profile_test.go @@ -0,0 +1,104 @@ +package modelplan + +import ( + "strings" + "testing" + "time" + + "github.com/toasterbook88/axis/internal/models" +) + +func TestPlanSingleNodeSelectedReadyProfile(t *testing.T) { + spec := models.ModelSpec{ + Schema: "axis.model-spec/v1", + ID: "ms-ready", + Name: "qwen", + Format: models.ModelFormatGGUF, + Source: "disk-weight", + Quantization: "Q4_K_M", + WeightsPath: "/mnt/models/qwen.gguf", + Memory: models.ModelMemoryRequirements{ + WeightSizeMB: 4096, + ContextOverheadMB: 512, + RuntimeOverheadMB: 256, + }, + Accelerators: []models.AcceleratorType{models.AcceleratorCUDA}, + } + snap := &models.ClusterSnapshot{ + Timestamp: time.Now().UTC(), + Publication: &models.PublicationEnvelope{ID: "pub-selected"}, + Nodes: []models.NodeFacts{{ + Name: "gpu-node", + Status: models.StatusComplete, + Tools: []models.ToolInfo{{Name: "llama-server", Path: "/usr/local/bin/llama-server"}}, + Resources: &models.Resources{ + RAMFreeMB: 64000, + RAMTotalMB: 128000, + Volumes: []models.Volume{{Mount: "/mnt/models", Kind: "local"}}, + GPUs: []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }}, + }, + }}, + } + plan, err := PlanSingleNode(snap, spec, 8080) + if err != nil { + t.Fatal(err) + } + if plan.BestCandidate != "gpu-node" || plan.Selected == nil { + t.Fatalf("best=%q selected=%v", plan.BestCandidate, plan.Selected) + } + got := plan.Selected + if got.Schema != models.ModelRunSchema || got.Node != "gpu-node" || got.Engine != "llama.cpp" { + t.Fatalf("selected=%+v", got) + } + if got.EngineBinary != "/usr/local/bin/llama-server" || got.Volume != "/mnt/models" || got.WeightsPath != "/mnt/models/qwen.gguf" { + t.Fatalf("selected=%+v", got) + } + if got.Port != 8080 || got.PortSource != models.PortSourcePlanDefault || got.BindHost != "127.0.0.1" { + t.Fatalf("selected=%+v", got) + } + if got.DeviceKind != models.DeviceKindDiscrete || !got.VRAMFreeMeasured || got.VRAMFreeMB != 20000 { + t.Fatalf("device=%+v", got) + } + if len(got.Refusals) != 0 || got.NGPULayers != nil || got.NGPULayersMode != "" || got.DeviceIndex != nil { + t.Fatalf("selected must not pin or offload: %+v", got) + } + if got.SnapshotPublicationID != "pub-selected" || got.Quantization != "Q4_K_M" || got.SpecSource != "disk-weight" { + t.Fatalf("provenance=%+v", got) + } +} + +func TestPlanSingleNodeSelectedSetWhenBestCandidateCannotLaunch(t *testing.T) { + spec := models.ModelSpec{ + Schema: "axis.model-spec/v1", + ID: "ms-bare", + Name: "qwen", + Format: models.ModelFormatGGUF, + WeightsPath: "/data/models/qwen.gguf", + Memory: models.ModelMemoryRequirements{WeightSizeMB: 4096, ContextOverheadMB: 512, RuntimeOverheadMB: 256}, + Accelerators: []models.AcceleratorType{models.AcceleratorCUDA, models.AcceleratorCPU}, + } + snap := &models.ClusterSnapshot{ + Nodes: []models.NodeFacts{{ + Name: "gpu-worker", + Status: models.StatusComplete, + Resources: &models.Resources{ + RAMFreeMB: 32000, RAMTotalMB: 64000, + GPUs: []models.GPUInfo{{Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, Capabilities: []string{"cuda"}}}, + }, + }}, + } + plan, err := PlanSingleNode(snap, spec, 8080) + if err != nil { + t.Fatal(err) + } + if plan.BestCandidate != "gpu-worker" || plan.Selected == nil { + t.Fatalf("best=%q selected nil=%v", plan.BestCandidate, plan.Selected == nil) + } + joined := strings.Join(plan.Selected.Refusals, "\n") + if !strings.Contains(joined, "llama-server") || !strings.Contains(joined, "named local volume") { + t.Fatalf("refusals=%v", plan.Selected.Refusals) + } +} diff --git a/internal/modelplan/single_node.go b/internal/modelplan/single_node.go index 7f8ad6aa..0197c6b8 100644 --- a/internal/modelplan/single_node.go +++ b/internal/modelplan/single_node.go @@ -56,6 +56,9 @@ type ModelPlacementPlan struct { Candidates []ModelCandidateScore `json:"candidates" yaml:"candidates"` Excluded []ModelExcludedCandidate `json:"excluded" yaml:"excluded"` BestCandidate string `json:"best_candidate,omitempty" yaml:"best_candidate,omitempty"` + // Selected is the launch profile for BestCandidate. It is advisory. + // Non-empty Refusals mean start must not exec that profile as written. + Selected *models.ModelRunProfile `json:"selected,omitempty" yaml:"selected,omitempty"` } // PlanSingleNode performs a dry-run evaluation of cluster nodes against available @@ -180,6 +183,13 @@ func PlanSingleNode(snapshot *models.ClusterSnapshot, spec models.ModelSpec, tar if len(plan.Candidates) > 0 { plan.BestCandidate = plan.Candidates[0].Node + for i := range snapshot.Nodes { + if snapshot.Nodes[i].Name == plan.BestCandidate { + selected := models.NewPlanProfile(snapshot.Nodes[i], spec, targetPort, plan.PublicationID) + plan.Selected = &selected + break + } + } } return plan, nil @@ -223,14 +233,7 @@ type acceleratorFit struct { // 0 from an unknown 0. A positive free value is a measurement even when a // collector that predates #442 omitted the flag. func observedFreeVRAM(gpu models.GPUInfo) (freeMB int64, measured bool) { - total := int64(gpu.VRAMMB) - if gpu.VRAMFreeMB < 0 { - return total, false - } - if gpu.VRAMFreeMeasured || gpu.VRAMFreeMB > 0 { - return int64(gpu.VRAMFreeMB), true - } - return total, false + return models.MeasuredFreeVRAM(gpu) } // evaluateNodeAccelerator reports the best single compatible device on a node. diff --git a/internal/models/model_operation.go b/internal/models/model_operation.go index ec7b3dc3..f949b718 100644 --- a/internal/models/model_operation.go +++ b/internal/models/model_operation.go @@ -52,5 +52,10 @@ type ModelOperationReceipt struct { TotalTokens int `json:"total_tokens,omitempty" yaml:"total_tokens,omitempty"` EndpointURL string `json:"endpoint_url,omitempty" yaml:"endpoint_url,omitempty"` ResponseText string `json:"response_text,omitempty" yaml:"response_text,omitempty"` + SpecSource string `json:"spec_source,omitempty" yaml:"spec_source,omitempty"` + DeviceKind string `json:"device_kind,omitempty" yaml:"device_kind,omitempty"` + DeviceIndex *int `json:"device_index,omitempty" yaml:"device_index,omitempty"` + VRAMFreeMeasured bool `json:"vram_free_measured,omitempty" yaml:"vram_free_measured,omitempty"` + PortSource string `json:"port_source,omitempty" yaml:"port_source,omitempty"` Error string `json:"error,omitempty" yaml:"error,omitempty"` } diff --git a/internal/models/run_profile.go b/internal/models/run_profile.go new file mode 100644 index 00000000..442f6a07 --- /dev/null +++ b/internal/models/run_profile.go @@ -0,0 +1,377 @@ +package models + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "path" + "strconv" + "strings" +) + +const ( + // ModelRunSchema is the JSON schema of a launch profile. + ModelRunSchema = "axis.model-run/v1" + // PortSourceExplicit is an operator-chosen port. + PortSourceExplicit = "explicit" + // PortSourcePlanDefault is the plan command's default occupancy port. + PortSourcePlanDefault = "plan-default" + // DeviceKindDiscrete is a CUDA or ROCm device with its own VRAM. + DeviceKindDiscrete = "discrete" + // DeviceKindUnified is Apple unified memory, or a node marked unified. + DeviceKindUnified = "unified" + // DeviceKindCPU is a node with no discrete or unified accelerator. + DeviceKindCPU = "cpu" + // EngineLlamaCpp is the only engine PR1 can launch. + EngineLlamaCpp = "llama.cpp" + // ToolLlamaServer is the observed tool name for llama.cpp. + ToolLlamaServer = "llama-server" + // ArtifactWeightsPath is a local weight file, not an Ollama model name. + ArtifactWeightsPath = "weights-path" +) + +// ModelRunProfile is the launch description shared by model plan and model start. +// It does not exec. Argv is derived from it. +type ModelRunProfile struct { + Schema string `json:"schema"` + Node string `json:"node"` + Engine string `json:"engine"` + EngineBinary string `json:"engine_binary,omitempty"` + ToolName string `json:"tool_name,omitempty"` + SpecID string `json:"spec_id,omitempty"` + ArtifactKind string `json:"artifact_kind,omitempty"` + WeightsPath string `json:"weights_path,omitempty"` + OllamaModel string `json:"ollama_model,omitempty"` + Format ModelFormat `json:"format,omitempty"` + Quantization string `json:"quantization,omitempty"` + Volume string `json:"volume,omitempty"` + SpecSource string `json:"spec_source,omitempty"` + DeviceKind string `json:"device_kind,omitempty"` + DeviceIndex *int `json:"device_index,omitempty"` + IndexSource string `json:"index_source,omitempty"` + DeviceModel string `json:"device_model,omitempty"` + MemoryTopology MemoryTopology `json:"memory_topology,omitempty"` + VRAMFreeMB int64 `json:"vram_free_mb,omitempty"` + VRAMFreeMeasured bool `json:"vram_free_measured,omitempty"` + Accelerator string `json:"accelerator,omitempty"` + BindHost string `json:"bind_host"` + Port int `json:"port"` + ContextTokens *int `json:"context_tokens,omitempty"` + NGPULayers *int `json:"n_gpu_layers,omitempty"` + NGPULayersMode string `json:"n_gpu_layers_mode,omitempty"` + BatchSize *int `json:"batch_size,omitempty"` + UBatchSize *int `json:"ubatch_size,omitempty"` + Threads *int `json:"threads,omitempty"` + OllamaNumCtx *int `json:"ollama_num_ctx,omitempty"` + OllamaKeepAlive string `json:"ollama_keep_alive,omitempty"` + OllamaNumGPU *int `json:"ollama_num_gpu,omitempty"` + MLXModel string `json:"mlx_model,omitempty"` + PrefillStepSize *int `json:"prefill_step_size,omitempty"` + PromptCacheBytes *int64 `json:"prompt_cache_bytes,omitempty"` + KVBits *int `json:"kv_bits,omitempty"` + PortSource string `json:"port_source,omitempty"` + SnapshotPublicationID string `json:"snapshot_publication_id,omitempty"` + Refusals []string `json:"refusals,omitempty"` +} + +// LaunchDevice is the one device a llama-server launch is allowed to name. +type LaunchDevice struct { + Kind string + Model string + Accelerator string + MemoryTopology MemoryTopology + VRAMFreeMB int64 + VRAMFreeMeasured bool +} + +// NamedLocalVolume returns the longest non-network mount that contains weights. +// The match is the historical namedLocalVolume rule: a mount of "/" does not +// prefix-match other absolute paths, because the prefix used is "//". +func NamedLocalVolume(node NodeFacts, weights string) (string, bool) { + if node.Resources == nil { + return "", false + } + weights = path.Clean(strings.TrimSpace(weights)) + best := "" + for _, vol := range node.Resources.Volumes { + if vol.Kind == "network" || vol.Mount == "" { + continue + } + mount := path.Clean(vol.Mount) + if weights == mount || strings.HasPrefix(weights, mount+"/") { + if len(mount) > len(best) { + best = mount + } + } + } + if best == "" { + return "", false + } + return best, true +} + +// MeasuredFreeVRAM is the free-VRAM reading shared with the placement planner. +// A negative free value is unmeasured and is not used. A positive free value is +// measured even when the collector omitted the flag. A flagged zero stays zero. +func MeasuredFreeVRAM(gpu GPUInfo) (freeMB int64, measured bool) { + total := int64(gpu.VRAMMB) + if gpu.VRAMFreeMB < 0 { + return total, false + } + if gpu.VRAMFreeMeasured || gpu.VRAMFreeMB > 0 { + return int64(gpu.VRAMFreeMB), true + } + return total, false +} + +// ObserveLaunchDevice classifies the node for a llama-server launch. +// Unified topology wins even when a discrete GPU is also present, and that +// result does not report measured discrete VRAM. Otherwise the best discrete +// device wins by measured-or-total free VRAM, then by total VRAM. +func ObserveLaunchDevice(node NodeFacts) LaunchDevice { + if node.Resources == nil { + return LaunchDevice{Kind: DeviceKindCPU, Accelerator: "cpu"} + } + if node.Resources.MemoryTopology == MemoryTopologyUnified { + return unifiedLaunchDevice(node.Resources) + } + if gpu, acc, ok := bestDiscreteGPU(node.Resources.GPUs); ok { + free, measured := MeasuredFreeVRAM(gpu) + return LaunchDevice{ + Kind: DeviceKindDiscrete, + Model: gpu.Model, + Accelerator: acc, + MemoryTopology: node.Resources.MemoryTopology, + VRAMFreeMB: free, + VRAMFreeMeasured: measured, + } + } + if gpu, ok := firstMetalGPU(node.Resources.GPUs); ok { + return LaunchDevice{ + Kind: DeviceKindUnified, + Model: gpu.Model, + Accelerator: "metal", + MemoryTopology: node.Resources.MemoryTopology, + } + } + return LaunchDevice{ + Kind: DeviceKindCPU, + Accelerator: "cpu", + MemoryTopology: node.Resources.MemoryTopology, + } +} + +func unifiedLaunchDevice(res *Resources) LaunchDevice { + dev := LaunchDevice{ + Kind: DeviceKindUnified, + Accelerator: "cpu", + MemoryTopology: MemoryTopologyUnified, + } + if gpu, ok := firstMetalGPU(res.GPUs); ok { + dev.Model = gpu.Model + dev.Accelerator = "metal" + return dev + } + if len(res.GPUs) > 0 { + dev.Model = res.GPUs[0].Model + } + return dev +} + +func bestDiscreteGPU(gpus []GPUInfo) (GPUInfo, string, bool) { + var best GPUInfo + var bestAcc string + var bestFree, bestTotal int64 + found := false + for _, gpu := range gpus { + acc, ok := discreteAccelerator(gpu) + if !ok { + continue + } + free, _ := MeasuredFreeVRAM(gpu) + total := int64(gpu.VRAMMB) + if !found || free > bestFree || (free == bestFree && total > bestTotal) { + best = gpu + bestAcc = acc + bestFree = free + bestTotal = total + found = true + } + } + return best, bestAcc, found +} + +func discreteAccelerator(gpu GPUInfo) (string, bool) { + if gpu.HasCapability("cuda") || strings.EqualFold(gpu.Vendor, "nvidia") { + return "cuda", true + } + if gpu.HasCapability("rocm") || strings.EqualFold(gpu.Vendor, "amd") { + return "rocm", true + } + return "", false +} + +func firstMetalGPU(gpus []GPUInfo) (GPUInfo, bool) { + for _, gpu := range gpus { + if gpu.HasCapability("metal") || strings.EqualFold(gpu.Vendor, "apple") { + return gpu, true + } + } + return GPUInfo{}, false +} + +// ParseNGPULayers parses the start flag. An integer of at least 1 and the +// tokens auto and all are the only accepted values. +func ParseNGPULayers(raw string) (*int, string, error) { + raw = strings.TrimSpace(raw) + switch raw { + case "auto", "all": + return nil, raw, nil + } + n, err := strconv.Atoi(raw) + if err != nil || n < 1 { + return nil, "", fmt.Errorf("n-gpu-layers must be an integer >= 1, auto, or all") + } + return &n, "", nil +} + +// LoadModelRunProfile decodes one axis.model-run/v1 object. +// A placement plan, or a document that wraps the profile, is rejected. +func LoadModelRunProfile(data []byte) (ModelRunProfile, error) { + var header struct { + Schema string `json:"schema"` + } + if err := json.Unmarshal(data, &header); err != nil { + return ModelRunProfile{}, err + } + if header.Schema != ModelRunSchema { + return ModelRunProfile{}, fmt.Errorf("expected axis.model-run/v1") + } + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + var profile ModelRunProfile + if err := dec.Decode(&profile); err != nil { + return ModelRunProfile{}, err + } + var extra any + if err := dec.Decode(&extra); err != io.EOF { + return ModelRunProfile{}, fmt.Errorf("expected axis.model-run/v1") + } + return profile, nil +} + +// NewPlanProfile builds the advisory profile for one node. Refusals are +// recorded on the profile. The caller does not exec it. +func NewPlanProfile(node NodeFacts, spec ModelSpec, port int, publicationID string) ModelRunProfile { + weights := path.Clean(strings.TrimSpace(spec.WeightsPath)) + profile := ModelRunProfile{ + Schema: ModelRunSchema, + Node: node.Name, + Engine: EngineLlamaCpp, + ToolName: ToolLlamaServer, + SpecID: spec.ID, + ArtifactKind: ArtifactWeightsPath, + WeightsPath: weights, + Format: spec.Format, + Quantization: spec.Quantization, + SpecSource: spec.Source, + BindHost: "127.0.0.1", + Port: port, + PortSource: PortSourcePlanDefault, + SnapshotPublicationID: publicationID, + } + if tool, ok := toolByName(node, ToolLlamaServer); ok { + profile.EngineBinary = tool.Path + } else { + profile.Refusals = append(profile.Refusals, fmt.Sprintf("node %s has no observed llama-server tool", node.Name)) + } + if weights == "" || weights == "." { + profile.Refusals = append(profile.Refusals, "weights path is required") + } else if mount, ok := NamedLocalVolume(node, weights); ok { + profile.Volume = mount + } else { + profile.Refusals = append(profile.Refusals, fmt.Sprintf("weights %s are not on a named local volume", weights)) + } + applyLaunchDevice(&profile, ObserveLaunchDevice(node)) + return profile +} + +func toolByName(node NodeFacts, name string) (ToolInfo, bool) { + for _, tool := range node.Tools { + if strings.EqualFold(tool.Name, name) { + return tool, true + } + } + return ToolInfo{}, false +} + +func applyLaunchDevice(profile *ModelRunProfile, dev LaunchDevice) { + profile.DeviceKind = dev.Kind + profile.DeviceModel = dev.Model + profile.Accelerator = dev.Accelerator + profile.MemoryTopology = dev.MemoryTopology + profile.VRAMFreeMB = dev.VRAMFreeMB + profile.VRAMFreeMeasured = dev.VRAMFreeMeasured +} + +// Validate checks the closed llama-server field set. Fact-plane refusals +// (tool, volume, measured VRAM, CPU count) are not decided here. +func (p ModelRunProfile) Validate() error { + if p.Schema != ModelRunSchema { + return fmt.Errorf("expected axis.model-run/v1") + } + if p.Engine != EngineLlamaCpp { + return fmt.Errorf("engine %q is not supported", p.Engine) + } + if p.BindHost != "127.0.0.1" { + return fmt.Errorf("bind host must be 127.0.0.1") + } + if p.Port < 1 || p.Port > 65535 { + return fmt.Errorf("port must be between 1 and 65535") + } + weights := path.Clean(strings.TrimSpace(p.WeightsPath)) + if weights == "" || weights == "." { + return fmt.Errorf("weights path is required") + } + if p.DeviceIndex != nil || p.IndexSource != "" { + return fmt.Errorf("device index is not supported") + } + if err := p.validateLaunchFields(); err != nil { + return err + } + if p.hasForeignEngineFields() { + return fmt.Errorf("only llama-server launch fields are supported") + } + return nil +} + +func (p ModelRunProfile) validateLaunchFields() error { + if p.NGPULayers != nil && p.NGPULayersMode != "" { + return fmt.Errorf("n-gpu-layers accepts an integer or a mode, not both") + } + if p.NGPULayers != nil && *p.NGPULayers < 1 { + return fmt.Errorf("n-gpu-layers must be an integer >= 1, auto, or all") + } + if p.NGPULayersMode != "" && p.NGPULayersMode != "auto" && p.NGPULayersMode != "all" { + return fmt.Errorf("n-gpu-layers must be an integer >= 1, auto, or all") + } + if p.ContextTokens != nil && *p.ContextTokens < 1 { + return fmt.Errorf("ctx-size must be >= 1") + } + if p.BatchSize != nil && *p.BatchSize < 1 { + return fmt.Errorf("batch-size must be >= 1") + } + if p.UBatchSize != nil && *p.UBatchSize < 1 { + return fmt.Errorf("ubatch-size must be >= 1") + } + if p.Threads != nil && *p.Threads < 1 { + return fmt.Errorf("threads must be >= 1") + } + return nil +} + +func (p ModelRunProfile) hasForeignEngineFields() bool { + return p.OllamaModel != "" || p.OllamaNumCtx != nil || p.OllamaKeepAlive != "" || p.OllamaNumGPU != nil || + p.MLXModel != "" || p.PrefillStepSize != nil || p.PromptCacheBytes != nil || p.KVBits != nil +} diff --git a/internal/models/run_profile_test.go b/internal/models/run_profile_test.go new file mode 100644 index 00000000..81dde732 --- /dev/null +++ b/internal/models/run_profile_test.go @@ -0,0 +1,160 @@ +package models + +import ( + "strings" + "testing" +) + +func TestNamedLocalVolumeSkipsNetworkAndKeepsLongestMount(t *testing.T) { + node := NodeFacts{Resources: &Resources{Volumes: []Volume{ + {Mount: "/", Kind: "local"}, + {Mount: "/mnt/models", Kind: "local"}, + {Mount: "/mnt/models/hot", Kind: "local"}, + {Mount: "/mnt/nas", Kind: "network"}, + {Mount: "", Kind: "local"}, + }}} + + mount, ok := NamedLocalVolume(node, "/mnt/models/hot/a.gguf") + if !ok || mount != "/mnt/models/hot" { + t.Fatalf("mount=%q ok=%v, want /mnt/models/hot", mount, ok) + } + if _, ok := NamedLocalVolume(node, "/mnt/nas/a.gguf"); ok { + t.Fatal("network mount must not win") + } + if mount, ok = NamedLocalVolume(node, "/mnt/models"); !ok || mount != "/mnt/models" { + t.Fatalf("exact mount=%q ok=%v", mount, ok) + } + // A mount of "/" uses prefix "//", so it does not match every absolute path. + if _, ok := NamedLocalVolume(node, "/etc/passwd"); ok { + t.Fatal("root mount must not match an unrelated absolute path") + } + if _, ok := NamedLocalVolume(NodeFacts{}, "/mnt/models/a.gguf"); ok { + t.Fatal("missing resources must not match") + } +} + +func TestMeasuredFreeVRAMMatchesPlannerRules(t *testing.T) { + cases := []struct { + name string + gpu GPUInfo + wantFree int64 + wantMeas bool + }{ + {name: "negative is unmeasured total", gpu: GPUInfo{VRAMMB: 8000, VRAMFreeMB: -1, VRAMFreeMeasured: true}, wantFree: 8000}, + {name: "measured zero stays zero", gpu: GPUInfo{VRAMMB: 8000, VRAMFreeMB: 0, VRAMFreeMeasured: true}, wantFree: 0, wantMeas: true}, + {name: "unflagged zero uses total", gpu: GPUInfo{VRAMMB: 8000}, wantFree: 8000}, + {name: "positive free is measured without the flag", gpu: GPUInfo{VRAMMB: 8000, VRAMFreeMB: 1000}, wantFree: 1000, wantMeas: true}, + {name: "flagged positive free", gpu: GPUInfo{VRAMMB: 8000, VRAMFreeMB: 1000, VRAMFreeMeasured: true}, wantFree: 1000, wantMeas: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + free, measured := MeasuredFreeVRAM(tc.gpu) + if free != tc.wantFree || measured != tc.wantMeas { + t.Fatalf("free=%d measured=%v, want %d %v", free, measured, tc.wantFree, tc.wantMeas) + } + }) + } +} + +func TestObserveLaunchDevice(t *testing.T) { + measured := GPUInfo{Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}} + smaller := GPUInfo{Vendor: "nvidia", Model: "RTX 3060", VRAMMB: 12288, VRAMFreeMB: 4000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}} + node := NodeFacts{Resources: &Resources{GPUs: []GPUInfo{smaller, measured}}} + got := ObserveLaunchDevice(node) + if got.Kind != DeviceKindDiscrete || got.Accelerator != "cuda" || got.Model != "RTX 4090" || !got.VRAMFreeMeasured || got.VRAMFreeMB != 20000 { + t.Fatalf("discrete device = %+v", got) + } + + unified := NodeFacts{Resources: &Resources{ + MemoryTopology: MemoryTopologyUnified, + GPUs: []GPUInfo{measured}, + }} + got = ObserveLaunchDevice(unified) + if got.Kind != DeviceKindUnified || got.VRAMFreeMeasured { + t.Fatalf("unified topology must not expose a discrete measured device: %+v", got) + } + + apple := NodeFacts{Resources: &Resources{GPUs: []GPUInfo{{ + Vendor: "apple", Model: "M3", Capabilities: []string{"metal"}, + }}}} + got = ObserveLaunchDevice(apple) + if got.Kind != DeviceKindUnified || got.Accelerator != "metal" { + t.Fatalf("apple gpu = %+v", got) + } + + got = ObserveLaunchDevice(NodeFacts{}) + if got.Kind != DeviceKindCPU || got.Accelerator != "cpu" || got.VRAMFreeMeasured { + t.Fatalf("cpu = %+v", got) + } +} + +func TestParseNGPULayers(t *testing.T) { + n, mode, err := ParseNGPULayers("12") + if err != nil || mode != "" || n == nil || *n != 12 { + t.Fatalf("n=%v mode=%q err=%v", n, mode, err) + } + n, mode, err = ParseNGPULayers("auto") + if err != nil || n != nil || mode != "auto" { + t.Fatalf("n=%v mode=%q err=%v", n, mode, err) + } + if _, _, err := ParseNGPULayers("all"); err != nil { + t.Fatal(err) + } + for _, bad := range []string{"0", "-1", "nope", ""} { + if _, _, err := ParseNGPULayers(bad); err == nil || !strings.Contains(err.Error(), "n-gpu-layers") { + t.Fatalf("%q err=%v", bad, err) + } + } +} + +func TestLoadModelRunProfileRejectsPlanDocument(t *testing.T) { + _, err := LoadModelRunProfile([]byte(`{"schema":"axis.model-plan/v1","best_candidate":"gpu"}`)) + if err == nil || !strings.Contains(err.Error(), "expected axis.model-run/v1") { + t.Fatalf("plan document err=%v", err) + } + _, err = LoadModelRunProfile([]byte(`{"selected":{"schema":"axis.model-run/v1"}}`)) + if err == nil || !strings.Contains(err.Error(), "expected axis.model-run/v1") { + t.Fatalf("wrapper err=%v", err) + } + _, err = LoadModelRunProfile([]byte(`{"schema":"axis.model-run/v1","extra":true}`)) + if err == nil || !strings.Contains(err.Error(), "unknown field") { + t.Fatalf("unknown field err=%v", err) + } +} + +func TestNewPlanProfileRecordsToolAndVolumeRefusals(t *testing.T) { + node := NodeFacts{ + Name: "gpu-node", + Resources: &Resources{ + GPUs: []GPUInfo{{Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, Capabilities: []string{"cuda"}}}, + }, + } + spec := ModelSpec{ + ID: "ms-test", + Name: "qwen", + Format: ModelFormatGGUF, + Source: "disk-weight", + WeightsPath: "/mnt/models/qwen.gguf", + Quantization: "Q4_K_M", + } + profile := NewPlanProfile(node, spec, 8080, "pub-1") + if profile.Schema != ModelRunSchema || profile.PortSource != PortSourcePlanDefault { + t.Fatalf("profile identity = %+v", profile) + } + if profile.Engine != "llama.cpp" || profile.ToolName != "llama-server" || profile.BindHost != "127.0.0.1" { + t.Fatalf("profile launch = %+v", profile) + } + if profile.SpecID != "ms-test" || profile.Quantization != "Q4_K_M" || profile.SpecSource != "disk-weight" { + t.Fatalf("spec copy = %+v", profile) + } + if profile.SnapshotPublicationID != "pub-1" || profile.NGPULayers != nil || profile.NGPULayersMode != "" { + t.Fatalf("plan must not invent an offload flag: %+v", profile) + } + joined := strings.Join(profile.Refusals, "\n") + if !strings.Contains(joined, "llama-server") || !strings.Contains(joined, "named local volume") { + t.Fatalf("refusals=%v", profile.Refusals) + } + if profile.DeviceKind != DeviceKindDiscrete || profile.VRAMFreeMeasured { + t.Fatalf("unmeasured discrete device = %+v", profile) + } +} From 8cf7c95f0b928617d208e2cef4e8cdb4f1496a13 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 16:21:16 -0400 Subject: [PATCH 2/9] feat(facts): record a GPU index and allow an explicit llama-server pin The three nvidia-smi memory queries now start with index. A four-column row records Index and IndexSource nvidia-smi only when every numeric cell parses. Legacy three-column and two-column rows, Metal, and lspci leave Index nil. --main-gpu N is emitted only when that N matches an observed nvidia-smi index. The receipt quotes the split-mode none and row help. An omitted pin does not become 0. --- cmd/axis/model.go | 8 ++- cmd/axis/model_run_profile.go | 8 +++ cmd/axis/model_run_profile_test.go | 50 ++++++++++++++ internal/facts/gpu_index_test.go | 62 +++++++++++++++++ internal/facts/gpu_query_test.go | 35 ++++++++++ internal/facts/local_gpu.go | 66 ++++++++++++++++-- internal/facts/remote.go | 6 +- internal/facts/remote_bundle.go | 2 +- internal/modellife/plan.go | 21 ++++++ internal/modellife/plan_profile_test.go | 90 +++++++++++++++++++++++++ internal/models/model_operation.go | 1 + internal/models/run_profile.go | 31 ++++++++- internal/models/types.go | 4 ++ 13 files changed, 373 insertions(+), 11 deletions(-) create mode 100644 internal/facts/gpu_index_test.go create mode 100644 internal/facts/gpu_query_test.go diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 30bce8c2..28faf7f3 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -99,7 +99,7 @@ func modelPlanCmd() *cobra.Command { func modelStartCmd() *cobra.Command { var node, weights, cacheAddr, format, fromPlan, nGPULayers string - var port, ctxSize, batchSize, ubatchSize, threads int + var port, ctxSize, batchSize, ubatchSize, threads, mainGPU int var live bool cmd := &cobra.Command{ Use: "start", @@ -126,6 +126,7 @@ func modelStartCmd() *cobra.Command { cmd.Flags().IntVar(&batchSize, "batch-size", 0, "llama-server logical batch size (-b); omitted when unset") cmd.Flags().IntVar(&ubatchSize, "ubatch-size", 0, "llama-server physical batch size (-ub); omitted when unset") cmd.Flags().IntVar(&threads, "threads", 0, "llama-server threads (-t); must be within observed CPU cores") + cmd.Flags().IntVar(&mainGPU, "main-gpu", 0, "llama-server --main-gpu; omitted unless set; must match an observed nvidia-smi index") cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().BoolVar(&live, "live", false, "Bypass daemon cache and perform live fleet discovery") cmd.Flags().StringVar(&format, "format", "text", "Start operation receipt format: text, json, or yaml") @@ -714,6 +715,7 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st DeviceIndex: plan.Profile.DeviceIndex, VRAMFreeMeasured: plan.Profile.VRAMFreeMeasured, PortSource: plan.Profile.PortSource, + DeviceNote: models.MainGPUPinNote(plan.Profile.DeviceIndex), } if snap.Publication != nil { receipt.PublicationID = snap.Publication.ID @@ -747,6 +749,10 @@ func writeModelStartReceipt(cmd *cobra.Command, receipt models.ModelOperationRec if receipt.Status == models.ModelOperationCompleted { _, err := fmt.Fprintf(cmd.OutOrStdout(), "started %s on %s:%d volume %s operation %s\n", receipt.Executable, receipt.Node, receipt.Port, receipt.Volume, receipt.ID) + if err != nil || receipt.DeviceNote == "" { + return err + } + _, err = fmt.Fprintf(cmd.OutOrStdout(), "%s\n", receipt.DeviceNote) return err } _, err := fmt.Fprintf(cmd.OutOrStdout(), "%s %s:%d: %s operation %s\n", diff --git a/cmd/axis/model_run_profile.go b/cmd/axis/model_run_profile.go index 124fd859..7b3feb9b 100644 --- a/cmd/axis/model_run_profile.go +++ b/cmd/axis/model_run_profile.go @@ -113,6 +113,14 @@ func applyChangedStartFlags(cmd *cobra.Command, profile *models.ModelRunProfile) } profile.Threads = &value } + if cmd.Flags().Changed("main-gpu") { + value, err := cmd.Flags().GetInt("main-gpu") + if err != nil { + return err + } + profile.DeviceIndex = &value + profile.IndexSource = models.IndexSourceNvidiaSMI + } return nil } diff --git a/cmd/axis/model_run_profile_test.go b/cmd/axis/model_run_profile_test.go index 1e3fe36f..efb22f73 100644 --- a/cmd/axis/model_run_profile_test.go +++ b/cmd/axis/model_run_profile_test.go @@ -183,6 +183,56 @@ func TestModelPlanWriteProfileAndExplicitPort(t *testing.T) { } } +func TestModelStartMainGPUMatchesObservedIndex(t *testing.T) { + snap := testSnap() + zero := 0 + snap.Nodes[0].Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", Index: &zero, IndexSource: models.IndexSourceNvidiaSMI, + VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prev := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prev }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{ + "--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", + "--main-gpu", "0", "--format", "json", + }) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + want := []string{ + "/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1", + "--main-gpu", "0", + } + if len(runner.started) != 1 || !reflect.DeepEqual(runner.started[0], want) { + t.Fatalf("argv=%#v", runner.started) + } + if !strings.Contains(buf.String(), "split-mode = none") || !strings.Contains(buf.String(), "split-mode = row") || !strings.Contains(buf.String(), "list-devices") { + t.Fatalf("receipt=%s", buf.String()) + } + + runner.started = nil + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{ + "--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", + "--main-gpu", "1", "--format", "text", + }) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "main-gpu") { + t.Fatalf("missing index err=%v", err) + } + if len(runner.started) != 0 { + t.Fatalf("unobserved pin started=%v", runner.started) + } +} + func TestRunModelStartDefaultPathStillUsesFunctionArgs(t *testing.T) { stubModelSnapshot(t, testSnap()) stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) diff --git a/internal/facts/gpu_index_test.go b/internal/facts/gpu_index_test.go new file mode 100644 index 00000000..8473cf48 --- /dev/null +++ b/internal/facts/gpu_index_test.go @@ -0,0 +1,62 @@ +package facts + +import ( + "strings" + "testing" +) + +func TestParseNvidiaSMIOutput_FourColumnsRecordIndex(t *testing.T) { + input := "0, NVIDIA GeForce RTX 4090, 24564, 20000\n1, NVIDIA GeForce RTX 3080, 10240, 0\n" + gpus := parseNvidiaSMIOutput(input) + if len(gpus) != 2 { + t.Fatalf("got %d gpus, want 2", len(gpus)) + } + if gpus[0].Index == nil || *gpus[0].Index != 0 { + t.Fatalf("gpu[0].Index = %v, want 0", gpus[0].Index) + } + if gpus[0].IndexSource != "nvidia-smi" { + t.Fatalf("gpu[0].IndexSource = %q", gpus[0].IndexSource) + } + if gpus[0].Model != "NVIDIA GeForce RTX 4090" || gpus[0].VRAMMB != 24564 || gpus[0].VRAMFreeMB != 20000 || !gpus[0].VRAMFreeMeasured { + t.Fatalf("gpu[0] = %+v", gpus[0]) + } + if gpus[1].Index == nil || *gpus[1].Index != 1 || gpus[1].VRAMFreeMB != 0 || !gpus[1].VRAMFreeMeasured { + t.Fatalf("gpu[1] = %+v", gpus[1]) + } +} + +func TestParseNvidiaSMIOutput_NonIntegerIndexDropsRow(t *testing.T) { + input := "nope, NVIDIA GeForce RTX 4090, 24564, 20000\n0, NVIDIA GeForce RTX 3080, not-a-number, 10\n" + gpus := parseNvidiaSMIOutput(input) + if len(gpus) != 0 { + t.Fatalf("bad cells must drop the row, got %+v", gpus) + } +} + +func TestParseNvidiaSMIOutput_LegacyWidthsLeaveIndexNil(t *testing.T) { + three := parseNvidiaSMIOutput("NVIDIA GeForce RTX 4090, 24564, 20480") + if len(three) != 1 || three[0].Index != nil || three[0].IndexSource != "" || three[0].VRAMFreeMB != 20480 || !three[0].VRAMFreeMeasured { + t.Fatalf("three-column = %+v", three) + } + two := parseNvidiaSMIOutput("NVIDIA GeForce RTX 4090, 24564") + if len(two) != 1 || two[0].Index != nil || two[0].IndexSource != "" || two[0].VRAMFreeMeasured || two[0].VRAMMB != 24564 { + t.Fatalf("two-column = %+v", two) + } +} + +func TestMetalAndLspciLeaveGPUIndexNil(t *testing.T) { + metal := parseSystemProfilerGPUs("Chipset Model: Apple M3\nVRAM (Dynamic, Max): 16384 MB\nMetal Family: Supported") + for _, gpu := range metal { + if gpu.Index != nil || gpu.IndexSource != "" { + t.Fatalf("metal gpu index = %+v", gpu) + } + } + lspci := "NVIDIA Corporation GA102 [GeForce RTX 3090]" + gpu := parseNvidiaSMIOutput(lspci) + if len(gpu) != 1 || gpu[0].Index != nil || gpu[0].IndexSource != "" { + t.Fatalf("lspci-shaped line = %+v", gpu) + } + if strings.Contains(lspci, ", ") { + t.Fatal("fixture accidentally looks like nvidia-smi csv") + } +} diff --git a/internal/facts/gpu_query_test.go b/internal/facts/gpu_query_test.go new file mode 100644 index 00000000..cc3375e9 --- /dev/null +++ b/internal/facts/gpu_query_test.go @@ -0,0 +1,35 @@ +package facts + +import ( + "strings" + "testing" +) + +func TestNvidiaSMIMemoryQueryPrependsIndex(t *testing.T) { + const want = "--query-gpu=index,name,memory.total,memory.free" + if nvidiaSMIMemoryQuery != want { + t.Fatalf("query = %q, want %q", nvidiaSMIMemoryQuery, want) + } + if strings.Contains(nvidiaSMIMemoryQuery, "utilization") || strings.Contains(nvidiaSMIMemoryQuery, "uuid") { + t.Fatalf("memory query must not be the util or resident query: %s", nvidiaSMIMemoryQuery) + } + local := strings.Join(localNvidiaSMIMemoryArgs(), " ") + if !strings.Contains(local, "nvidia-smi "+nvidiaSMIMemoryQuery) { + t.Fatalf("local args = %q", local) + } + if !strings.Contains(linuxGPUCollectCommand(), "nvidia-smi "+nvidiaSMIMemoryQuery) { + t.Fatalf("remote cmd = %q", linuxGPUCollectCommand()) + } + if !strings.Contains(remoteFactBundleScript, "nvidia-smi "+nvidiaSMIMemoryQuery) { + t.Fatal("remote bundle gpu_b64 lost the indexed memory query") + } + if !strings.Contains(darwinGPUCollectCommand(), "system_profiler SPDisplaysDataType") || strings.Contains(darwinGPUCollectCommand(), "nvidia-smi") { + t.Fatalf("darwin cmd = %q", darwinGPUCollectCommand()) + } + if !strings.Contains(remoteFactBundleScript, "system_profiler SPDisplaysDataType") { + t.Fatal("darwin bundle probe changed") + } + if !strings.Contains(LlamaServerDiscoveryScript, "--query-gpu=index,uuid") || strings.Contains(LlamaServerDiscoveryScript, "memory.total") { + t.Fatal("resident compute-apps query must stay index,uuid") + } +} diff --git a/internal/facts/local_gpu.go b/internal/facts/local_gpu.go index 01068886..0f2940a6 100644 --- a/internal/facts/local_gpu.go +++ b/internal/facts/local_gpu.go @@ -102,8 +102,26 @@ func localGPUsLinux(ctx context.Context) []models.GPUInfo { return gpus } +// nvidiaSMIMemoryQuery is the memory probe used by the local collector, the +// remote gpu command, and the remote fact bundle. The utilization probe and +// the resident index,uuid probe do not use it. +const nvidiaSMIMemoryQuery = "--query-gpu=index,name,memory.total,memory.free" + +func localNvidiaSMIMemoryArgs() []string { + return []string{"nvidia-smi", nvidiaSMIMemoryQuery, "--format=csv,noheader,nounits"} +} + +func linuxGPUCollectCommand() string { + return "nvidia-smi " + nvidiaSMIMemoryQuery + " --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //'" +} + +func darwinGPUCollectCommand() string { + return `system_profiler SPDisplaysDataType 2>/dev/null | grep -E 'Chipset Model:|VRAM|Metal' | sed 's/^ *//'` +} + func localGPUsNvidiaSMI(ctx context.Context) []models.GPUInfo { - out, err := exec.CommandContext(ctx, "nvidia-smi", "--query-gpu=name,memory.total,memory.free", "--format=csv,noheader,nounits").Output() + args := localNvidiaSMIMemoryArgs() + out, err := exec.CommandContext(ctx, args[0], args[1:]...).Output() if err != nil { return nil } @@ -118,21 +136,32 @@ func parseNvidiaSMIOutput(out string) []models.GPUInfo { continue } parts := strings.Split(line, ", ") - name := strings.TrimSpace(parts[0]) + for i := range parts { + parts[i] = strings.TrimSpace(parts[i]) + } + // Four columns are index,name,memory.total,memory.free. A cell that + // does not parse drops the row. Wider rows are not this query. + if len(parts) >= 4 { + if gpu, ok := nvidiaSMIIndexedGPU(parts); ok { + gpus = append(gpus, gpu) + } + continue + } + name := parts[0] gpu := models.GPUInfo{ Model: name, Vendor: "nvidia", Capabilities: []string{"cuda"}, } if len(parts) >= 2 { - if vram, err := strconv.Atoi(strings.TrimSpace(parts[1])); err == nil { + if vram, err := strconv.Atoi(parts[1]); err == nil { gpu.VRAMMB = vram } } // memory.free is present only when the query requests it; older callers // (and the two-column remote fallback) legitimately omit it. if len(parts) >= 3 { - if free, err := strconv.Atoi(strings.TrimSpace(parts[2])); err == nil { + if free, err := strconv.Atoi(parts[2]); err == nil { gpu.VRAMFreeMB = free gpu.VRAMFreeMeasured = true } @@ -142,6 +171,35 @@ func parseNvidiaSMIOutput(out string) []models.GPUInfo { return gpus } +func nvidiaSMIIndexedGPU(parts []string) (models.GPUInfo, bool) { + if len(parts) != 4 { + return models.GPUInfo{}, false + } + index, err := strconv.Atoi(parts[0]) + if err != nil { + return models.GPUInfo{}, false + } + total, err := strconv.Atoi(parts[2]) + if err != nil { + return models.GPUInfo{}, false + } + free, err := strconv.Atoi(parts[3]) + if err != nil { + return models.GPUInfo{}, false + } + idx := index + return models.GPUInfo{ + Model: parts[1], + Vendor: "nvidia", + Index: &idx, + IndexSource: "nvidia-smi", + VRAMMB: total, + VRAMFreeMB: free, + VRAMFreeMeasured: true, + Capabilities: []string{"cuda"}, + }, true +} + func localGPUsLspci(ctx context.Context) []models.GPUInfo { out, err := exec.CommandContext(ctx, "bash", "-c", `lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //'`).Output() if err != nil || len(out) == 0 { diff --git a/internal/facts/remote.go b/internal/facts/remote.go index e730b501..76c9a32c 100644 --- a/internal/facts/remote.go +++ b/internal/facts/remote.go @@ -323,10 +323,10 @@ func (c *RemoteCollector) remoteResources(ctx context.Context, osName, arch stri // GPU (best-effort) var gpuCmd string if osName == "darwin" { - gpuCmd = `system_profiler SPDisplaysDataType 2>/dev/null | grep -E 'Chipset Model:|VRAM|Metal' | sed 's/^ *//'` + gpuCmd = darwinGPUCollectCommand() } else { - // Try nvidia-smi first, fall back to lspci - gpuCmd = `nvidia-smi --query-gpu=name,memory.total,memory.free --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //'` + // Try nvidia-smi first, fall back to lspci. + gpuCmd = linuxGPUCollectCommand() } if out, err := c.Exec.Run(ctx, gpuCmd); err == nil { out = strings.TrimSpace(out) diff --git a/internal/facts/remote_bundle.go b/internal/facts/remote_bundle.go index 90a75a01..b89a66de 100644 --- a/internal/facts/remote_bundle.go +++ b/internal/facts/remote_bundle.go @@ -46,7 +46,7 @@ case "$(printf '%s' "$OS" | tr '[:upper:]' '[:lower:]')" in printf 'meminfo_b64=%s\n' "$(grep -E 'MemTotal|MemAvailable|MemFree' /proc/meminfo 2>/dev/null | base64 | tr -d '\n')" printf 'loadavg=%s\n' "$(cat /proc/loadavg 2>/dev/null)" printf 'pressure_b64=%s\n' "$(cat /proc/pressure/memory 2>/dev/null | base64 | tr -d '\n')" - printf 'gpu_b64=%s\n' "$(nvidia-smi --query-gpu=name,memory.total,memory.free --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //' | base64 | tr -d '\n')" + printf 'gpu_b64=%s\n' "$(nvidia-smi --query-gpu=index,name,memory.total,memory.free --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //' | base64 | tr -d '\n')" printf 'identity=%s\n' "$(cat /etc/machine-id 2>/dev/null || cat /var/lib/dbus/machine-id 2>/dev/null)" printf 'battery=%s\n' "$(cat /sys/class/power_supply/BAT0/capacity /sys/class/power_supply/BAT1/capacity /sys/class/power_supply/BATT/capacity 2>/dev/null | head -1)" printf 'power=%s\n' "$(for n in AC ADP0 ACAD Mains; do s=$(cat /sys/class/power_supply/$n/status 2>/dev/null); [ -n "$s" ] && echo "$s" && break; done)" diff --git a/internal/modellife/plan.go b/internal/modellife/plan.go index db7a830c..1504c204 100644 --- a/internal/modellife/plan.go +++ b/internal/modellife/plan.go @@ -68,6 +68,9 @@ func PlanStartProfile(node models.NodeFacts, profile models.ModelRunProfile) (St refusals = append(refusals, fmt.Sprintf("threads must be between 1 and %d observed cpu cores", cores)) } } + if profile.DeviceIndex != nil && !observedNvidiaGPUIndex(node, *profile.DeviceIndex) { + refusals = append(refusals, fmt.Sprintf("main-gpu %d is not an observed nvidia-smi index", *profile.DeviceIndex)) + } if len(refusals) > 0 { return StartPlan{}, fmt.Errorf("%s", strings.Join(refusals, "; ")) } @@ -165,9 +168,27 @@ func ArgvFromProfile(profile models.ModelRunProfile) ([]string, error) { if profile.Threads != nil { argv = append(argv, "-t", strconv.Itoa(*profile.Threads)) } + if profile.DeviceIndex != nil { + if profile.IndexSource != models.IndexSourceNvidiaSMI { + return nil, fmt.Errorf("index source %q is not nvidia-smi", profile.IndexSource) + } + argv = append(argv, "--main-gpu", strconv.Itoa(*profile.DeviceIndex)) + } return argv, nil } +func observedNvidiaGPUIndex(node models.NodeFacts, index int) bool { + if node.Resources == nil { + return false + } + for _, gpu := range node.Resources.GPUs { + if gpu.Index != nil && *gpu.Index == index && gpu.IndexSource == models.IndexSourceNvidiaSMI { + return true + } + } + return false +} + // ExecArgvMatchesProfile reports whether argv is exactly the profile projection. // A hand-built argv with an empty profile does not match. func ExecArgvMatchesProfile(plan StartPlan) error { diff --git a/internal/modellife/plan_profile_test.go b/internal/modellife/plan_profile_test.go index ae341c90..ccf0875f 100644 --- a/internal/modellife/plan_profile_test.go +++ b/internal/modellife/plan_profile_test.go @@ -8,6 +8,96 @@ import ( "github.com/toasterbook88/axis/internal/models" ) +func TestPlanStartProfileEmitsMainGPUOnlyForObservedNvidiaIndex(t *testing.T) { + node := storageNode() + zero := 0 + node.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", Index: &zero, IndexSource: "nvidia-smi", + VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + profile := readyProfile(node) + pin := 0 + profile.DeviceIndex = &pin + profile.IndexSource = models.IndexSourceNvidiaSMI + plan, err := PlanStartProfile(node, profile) + if err != nil { + t.Fatal(err) + } + want := []string{ + "/usr/local/bin/llama-server", "-m", "/mnt/models/a.gguf", "--port", "8081", "--host", "127.0.0.1", + "--main-gpu", "0", + } + if !reflect.DeepEqual(plan.Argv, want) { + t.Fatalf("argv=%#v", plan.Argv) + } + note := models.MainGPUPinNote(plan.Profile.DeviceIndex) + if !strings.Contains(note, "split-mode = none") || !strings.Contains(note, "split-mode = row") || !strings.Contains(note, "list-devices") { + t.Fatalf("note=%q", note) + } + if models.MainGPUPinNote(nil) != "" { + t.Fatal("omitted pin must not quote a default of 0") + } + + plain, err := PlanStartProfile(node, readyProfile(node)) + if err != nil { + t.Fatal(err) + } + if strings.Contains(strings.Join(plain.Argv, " "), "--main-gpu") { + t.Fatalf("omitted pin argv=%#v", plain.Argv) + } + + ctx, ngl := 2048, 12 + withFlags := readyProfile(node) + withFlags.ContextTokens = &ctx + withFlags.NGPULayers = &ngl + withFlags.DeviceIndex = &pin + withFlags.IndexSource = models.IndexSourceNvidiaSMI + plan, err = PlanStartProfile(node, withFlags) + if err != nil { + t.Fatal(err) + } + suffix := plan.Argv[len(plan.Argv)-6:] + wantSuffix := []string{"-c", "2048", "-ngl", "12", "--main-gpu", "0"} + if !reflect.DeepEqual(suffix, wantSuffix) { + t.Fatalf("suffix=%#v", plan.Argv) + } +} + +func TestPlanStartProfileRefusesMainGPUThatWasNotObserved(t *testing.T) { + node := storageNode() + one := 1 + node.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", Index: &one, IndexSource: "nvidia-smi", + VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + profile := readyProfile(node) + pin := 0 + profile.DeviceIndex = &pin + profile.IndexSource = models.IndexSourceNvidiaSMI + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "main-gpu") { + t.Fatalf("unobserved 0 err=%v", err) + } + + legacy := storageNode() + legacy.Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + profile = readyProfile(legacy) + profile.DeviceIndex = &pin + profile.IndexSource = models.IndexSourceNvidiaSMI + if _, err := PlanStartProfile(legacy, profile); err == nil || !strings.Contains(err.Error(), "main-gpu") { + t.Fatalf("nil observed index err=%v", err) + } + + profile = readyProfile(node) + profile.DeviceIndex = &one + profile.IndexSource = "hand" + if _, err := PlanStartProfile(node, profile); err == nil || !strings.Contains(err.Error(), "nvidia-smi") { + t.Fatalf("foreign source err=%v", err) + } +} + func TestPlanStartDefaultArgvIsExact(t *testing.T) { plan, err := PlanStart(storageNode(), "/mnt/models/a.gguf", 8081) if err != nil { diff --git a/internal/models/model_operation.go b/internal/models/model_operation.go index f949b718..c7f78ba2 100644 --- a/internal/models/model_operation.go +++ b/internal/models/model_operation.go @@ -57,5 +57,6 @@ type ModelOperationReceipt struct { DeviceIndex *int `json:"device_index,omitempty" yaml:"device_index,omitempty"` VRAMFreeMeasured bool `json:"vram_free_measured,omitempty" yaml:"vram_free_measured,omitempty"` PortSource string `json:"port_source,omitempty" yaml:"port_source,omitempty"` + DeviceNote string `json:"device_note,omitempty" yaml:"device_note,omitempty"` Error string `json:"error,omitempty" yaml:"error,omitempty"` } diff --git a/internal/models/run_profile.go b/internal/models/run_profile.go index 442f6a07..6ff015b1 100644 --- a/internal/models/run_profile.go +++ b/internal/models/run_profile.go @@ -29,8 +29,19 @@ const ( ToolLlamaServer = "llama-server" // ArtifactWeightsPath is a local weight file, not an Ollama model name. ArtifactWeightsPath = "weights-path" + // IndexSourceNvidiaSMI is the only index source a llama-server pin may name. + IndexSourceNvidiaSMI = "nvidia-smi" ) +// MainGPUPinNote quotes llama.cpp --main-gpu help when the operator pinned a +// device. An omitted pin returns empty so Axis does not imply a default of 0. +func MainGPUPinNote(index *int) string { + if index == nil { + return "" + } + return "main-gpu is the GPU to use for the model (with split-mode = none), or for intermediate results and KV (with split-mode = row). Axis does not emit --device; those names come from llama-server --list-devices, which Axis does not collect." +} + // ModelRunProfile is the launch description shared by model plan and model start. // It does not exec. Argv is derived from it. type ModelRunProfile struct { @@ -334,8 +345,8 @@ func (p ModelRunProfile) Validate() error { if weights == "" || weights == "." { return fmt.Errorf("weights path is required") } - if p.DeviceIndex != nil || p.IndexSource != "" { - return fmt.Errorf("device index is not supported") + if err := p.validateDevicePin(); err != nil { + return err } if err := p.validateLaunchFields(); err != nil { return err @@ -346,6 +357,22 @@ func (p ModelRunProfile) Validate() error { return nil } +func (p ModelRunProfile) validateDevicePin() error { + if p.DeviceIndex == nil && p.IndexSource == "" { + return nil + } + if p.DeviceIndex == nil || p.IndexSource == "" { + return fmt.Errorf("device index requires index source nvidia-smi") + } + if p.IndexSource != IndexSourceNvidiaSMI { + return fmt.Errorf("index source %q is not nvidia-smi", p.IndexSource) + } + if *p.DeviceIndex < 0 { + return fmt.Errorf("main-gpu must be >= 0") + } + return nil +} + func (p ModelRunProfile) validateLaunchFields() error { if p.NGPULayers != nil && p.NGPULayersMode != "" { return fmt.Errorf("n-gpu-layers accepts an integer or a mode, not both") diff --git a/internal/models/types.go b/internal/models/types.go index f2c61f9b..57086fd2 100644 --- a/internal/models/types.go +++ b/internal/models/types.go @@ -55,9 +55,13 @@ const ( // --- Observed State --- // GPUInfo describes a single GPU with vendor, model, VRAM, and capabilities. +// Index is set only when a collector parsed an nvidia-smi index. Nil means +// the row had no index fact (legacy csv, Metal, or lspci). type GPUInfo struct { Vendor string `json:"vendor" yaml:"vendor"` // apple, nvidia, amd, intel, unknown Model string `json:"model" yaml:"model"` // e.g. "Apple M3 Pro", "NVIDIA GeForce RTX 4090" + Index *int `json:"index,omitempty" yaml:"index,omitempty"` // nvidia-smi index; nil when the collector did not record one + IndexSource string `json:"index_source,omitempty" yaml:"index_source,omitempty"` // "nvidia-smi" when Index was parsed from that query VRAMMB int `json:"vram_mb,omitempty" yaml:"vram_mb,omitempty"` // 0 means unknown or unified VRAMFreeMB int `json:"vram_free_mb,omitempty" yaml:"vram_free_mb,omitempty"` // measured free VRAM; valid when VRAMFreeMeasured is true or VRAMFreeMB > 0 VRAMFreeMeasured bool `json:"vram_free_measured,omitempty" yaml:"vram_free_measured,omitempty"` // true when VRAMFreeMB was explicitly measured (even if 0 MB free) From fc6083000c85ef5ca16f17493f56c8138ad08b9f Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 16:30:47 -0400 Subject: [PATCH 3/9] feat(model): place an Ollama model on the server that is already listening --ollama-model replaces --weights and is mutually exclusive with it. Load and unload curl 127.0.0.1:11434 /api/generate, then GET /api/ps. Axis does not exec ollama serve, does not send num_gpu or main_gpu, and does not kill comm=ollama. A generation whose engine is ollama unloads on that path and never reaches the llama-server process kill. --- cmd/axis/model.go | 168 +++++++++++++++++++++++++- cmd/axis/model_ollama_test.go | 179 ++++++++++++++++++++++++++++ cmd/axis/model_run_profile.go | 47 ++++++++ internal/modellife/ollama.go | 84 +++++++++++++ internal/modellife/ollama_test.go | 115 ++++++++++++++++++ internal/models/run_profile.go | 43 ++++++- internal/models/run_profile_test.go | 49 ++++++++ 7 files changed, 677 insertions(+), 8 deletions(-) create mode 100644 cmd/axis/model_ollama_test.go create mode 100644 internal/modellife/ollama.go create mode 100644 internal/modellife/ollama_test.go diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 28faf7f3..61ba77e2 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -98,12 +98,12 @@ func modelPlanCmd() *cobra.Command { } func modelStartCmd() *cobra.Command { - var node, weights, cacheAddr, format, fromPlan, nGPULayers string - var port, ctxSize, batchSize, ubatchSize, threads, mainGPU int + var node, weights, cacheAddr, format, fromPlan, nGPULayers, ollamaModel, ollamaKeepAlive string + var port, ctxSize, batchSize, ubatchSize, threads, mainGPU, ollamaNumCtx int var live bool cmd := &cobra.Command{ Use: "start", - Short: "Start llama-server on a named node (explicit port and weights)", + Short: "Start llama-server, or place an Ollama model on the server already listening", SilenceUsage: true, PreRunE: func(cmd *cobra.Command, args []string) error { if err := validateOutputFormat(&format, "text", "json", "yaml")(cmd, args); err != nil { @@ -127,6 +127,9 @@ func modelStartCmd() *cobra.Command { cmd.Flags().IntVar(&ubatchSize, "ubatch-size", 0, "llama-server physical batch size (-ub); omitted when unset") cmd.Flags().IntVar(&threads, "threads", 0, "llama-server threads (-t); must be within observed CPU cores") cmd.Flags().IntVar(&mainGPU, "main-gpu", 0, "llama-server --main-gpu; omitted unless set; must match an observed nvidia-smi index") + cmd.Flags().StringVar(&ollamaModel, "ollama-model", "", "Ollama model name to load on 127.0.0.1:11434; replaces --weights") + cmd.Flags().StringVar(&ollamaKeepAlive, "ollama-keep-alive", "", "Ollama keep_alive for this load; omitted unless set") + cmd.Flags().IntVar(&ollamaNumCtx, "ollama-num-ctx", 0, "Ollama options.num_ctx; omitted unless set") cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().BoolVar(&live, "live", false, "Bypass daemon cache and perform live fleet discovery") cmd.Flags().StringVar(&format, "format", "text", "Start operation receipt format: text, json, or yaml") @@ -134,17 +137,26 @@ func modelStartCmd() *cobra.Command { } func modelStopCmd() *cobra.Command { - var node, cacheAddr, format string + var node, cacheAddr, format, ollamaModel string var port int cmd := &cobra.Command{ Use: "stop [generation-id]", - Short: "Stop an observed llama-server generation or use legacy node/port flags", + Short: "Stop an observed llama-server generation, unload an Ollama model, or use legacy node/port flags", Args: cobra.MaximumNArgs(1), SilenceUsage: true, PreRunE: validateOutputFormat(&format, "text", "json", "yaml"), RunE: func(cmd *cobra.Command, args []string) error { ctx, cancel := context.WithTimeout(cmd.Context(), 20*time.Second) defer cancel() + if strings.TrimSpace(ollamaModel) != "" { + if len(args) == 1 || port != 0 { + return fmt.Errorf("--ollama-model cannot be combined with a generation ID or --port") + } + if strings.TrimSpace(node) == "" { + return fmt.Errorf(`required flag(s) "node" not set`) + } + return runOllamaModelStop(ctx, cmd, node, ollamaModel, cacheAddr, format) + } if len(args) == 1 { if strings.TrimSpace(node) != "" || port != 0 { return fmt.Errorf("generation ID cannot be combined with --node or --port") @@ -154,8 +166,9 @@ func modelStopCmd() *cobra.Command { return runModelStop(ctx, cmd, node, port, defaultModelRunner) }, } - cmd.Flags().StringVar(&node, "node", "", "Legacy cluster node name") + cmd.Flags().StringVar(&node, "node", "", "Legacy cluster node name, or the node whose Ollama server unloads --ollama-model") cmd.Flags().IntVar(&port, "port", 0, "Legacy llama-server listen port") + cmd.Flags().StringVar(&ollamaModel, "ollama-model", "", "Unload this model from the Ollama server already listening on 127.0.0.1:11434") cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().StringVar(&format, "format", "text", "Generation-stop receipt format: text, json, or yaml") return cmd @@ -649,6 +662,9 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st if err != nil { return err } + if profile.Engine == models.EngineOllama { + return placeOllamaModel(ctx, cmd, nf, cfgNode, profile, source, snap, startedAt, format) + } for _, res := range nf.ResidentModels { if res.Port == profile.Port { @@ -822,6 +838,9 @@ func runModelStopGeneration(ctx context.Context, cmd *cobra.Command, generationI if instance.NodeStatus != models.StatusComplete { return fmt.Errorf("model generation %s is on node %s with status %s; refusing lifecycle mutation", generationID, instance.Node, instance.NodeStatus) } + if instance.Engine == models.EngineOllama { + return stopOllamaGeneration(ctx, cmd, snap, instance, format, startedAt) + } if instance.Engine != "llama.cpp" { return fmt.Errorf("model generation %s uses unsupported stop engine %q", generationID, instance.Engine) } @@ -1623,6 +1642,143 @@ func runOnNode(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeC return err } +// runNodeScript is the local-or-SSH curl seam for Ollama. Tests replace it. +// llama-server start and stop keep calling runOnNodeCapturing directly. +var runNodeScript = runOnNodeCapturing + +func placeOllamaModel(ctx context.Context, cmd *cobra.Command, node models.NodeFacts, cfgNode *config.NodeConfig, profile models.ModelRunProfile, source string, snap *models.ClusterSnapshot, startedAt time.Time, format string) error { + if err := profile.Validate(); err != nil { + return err + } + script, err := modellife.OllamaLoadScript(profile.OllamaModel, profile.OllamaKeepAlive, profile.OllamaNumCtx) + if err != nil { + return err + } + out, runErr := runNodeScript(ctx, node, cfgNode, script) + receipt := models.ModelOperationReceipt{ + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStart, + Status: models.ModelOperationCompleted, + Disposition: "placed", + Node: node.Name, + Engine: models.EngineOllama, + Model: profile.OllamaModel, + StartedAt: startedAt, + CompletedAt: time.Now().UTC(), + } + if snap != nil { + receipt.SnapshotSource = source + receipt.SnapshotAt = snap.Timestamp + if snap.Publication != nil { + receipt.PublicationID = snap.Publication.ID + } + } + if runErr != nil { + receipt.Status = models.ModelOperationFailed + receipt.Disposition = "failed" + receipt.Error = runErr.Error() + _ = writeModelStartReceipt(cmd, receipt, format) + return fmt.Errorf("ollama place failed: %w", runErr) + } + listed, parseErr := modellife.OllamaPSHasModel(out, profile.OllamaModel) + if parseErr != nil || !listed { + msg := "ollama /api/ps did not list the model" + if parseErr != nil { + msg = parseErr.Error() + } + receipt.Status = models.ModelOperationFailed + receipt.Disposition = "failed" + receipt.Error = msg + _ = writeModelStartReceipt(cmd, receipt, format) + return fmt.Errorf("%s", msg) + } + if writeErr := writeModelStartReceipt(cmd, receipt, format); writeErr != nil { + return writeErr + } + cacheAddr, _ := cmd.Flags().GetString("cache-addr") + warnModelDaemonRefresh(cmd, cacheAddr, "manual") + return nil +} + +func runOllamaModelStop(ctx context.Context, cmd *cobra.Command, nodeName, modelName, cacheAddr, format string) error { + live, _ := cmd.Flags().GetBool("live") + snap, _, err := loadModelCommandSnapshot(ctx, live, cacheAddr, "stop", false) + if err != nil { + return err + } + nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, nodeName) + if err != nil { + return err + } + if err := runOllamaUnload(ctx, nf, cfgNode, modelName); err != nil { + return err + } + receipt := models.ModelOperationReceipt{ + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStop, + Status: models.ModelOperationCompleted, + Disposition: "unloaded", + Node: nf.Name, + Engine: models.EngineOllama, + Model: modelName, + StartedAt: time.Now().UTC(), + CompletedAt: time.Now().UTC(), + } + return writeModelOperationReceipt(cmd, receipt, format) +} + +func stopOllamaGeneration(ctx context.Context, cmd *cobra.Command, snap *models.ClusterSnapshot, instance *models.ModelInstance, format string, startedAt time.Time) error { + nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, instance.Node) + if err != nil { + return err + } + unloadErr := runOllamaUnload(ctx, nf, cfgNode, instance.Model) + receipt := models.ModelOperationReceipt{ + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStop, + Status: models.ModelOperationCompleted, + Disposition: "unloaded", + InstanceID: instance.ID, + GenerationID: instance.GenerationID, + Node: instance.Node, + Engine: models.EngineOllama, + Model: instance.Model, + StartedAt: startedAt, + CompletedAt: time.Now().UTC(), + } + if unloadErr != nil { + receipt.Status = models.ModelOperationFailed + receipt.Disposition = "failed" + receipt.Error = unloadErr.Error() + } + if writeErr := writeModelOperationReceipt(cmd, receipt, format); writeErr != nil { + return writeErr + } + return unloadErr +} + +func runOllamaUnload(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig, modelName string) error { + script, err := modellife.OllamaUnloadScript(modelName) + if err != nil { + return err + } + out, err := runNodeScript(ctx, node, cfgNode, script) + if err != nil { + return fmt.Errorf("ollama unload failed: %w", err) + } + listed, err := modellife.OllamaPSHasModel(out, modelName) + if err != nil { + return err + } + if listed { + return fmt.Errorf("ollama /api/ps still lists %s", modelName) + } + return nil +} + // runOnNodeCapturing is runOnNode that also returns the command output, so // callers can read a result marker the script emitted. Both transports return // combined output, including on failure. diff --git a/cmd/axis/model_ollama_test.go b/cmd/axis/model_ollama_test.go new file mode 100644 index 00000000..0cd79446 --- /dev/null +++ b/cmd/axis/model_ollama_test.go @@ -0,0 +1,179 @@ +package main + +import ( + "bytes" + "context" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/config" + "github.com/toasterbook88/axis/internal/modelinventory" + "github.com/toasterbook88/axis/internal/models" +) + +func TestOllamaModelAndWeightsAreMutuallyExclusive(t *testing.T) { + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", "--ollama-model", "mistral"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "mutually exclusive") { + t.Fatalf("err=%v", err) + } +} + +func TestOllamaStartPlacesOnExistingLoopbackServer(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Running: true, Listening: true} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prevRunner := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prevRunner }) + + var scripts []string + prevScript := runNodeScript + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + scripts = append(scripts, script) + return `{"models":[{"name":"mistral","model":"mistral"}]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{ + "--node", "storage", "--ollama-model", "mistral", + "--ollama-keep-alive", "10m", "--ollama-num-ctx", "2048", + "--format", "json", + }) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if len(runner.started) != 0 || len(runner.probed) != 0 || len(runner.stopped) != 0 { + t.Fatalf("runner was used: started=%v probed=%v stopped=%v", runner.started, runner.probed, runner.stopped) + } + if len(scripts) != 1 { + t.Fatalf("scripts=%d", len(scripts)) + } + script := scripts[0] + if !strings.Contains(script, "http://127.0.0.1:11434/api/generate") || !strings.Contains(script, `"model":"mistral"`) || !strings.Contains(script, `"stream":false`) { + t.Fatalf("script=%s", script) + } + if !strings.Contains(script, `"keep_alive":"10m"`) || !strings.Contains(script, `"num_ctx":2048`) { + t.Fatalf("options missing: %s", script) + } + for _, banned := range []string{"num_gpu", "main_gpu", "num_batch", "num_thread", "ollama serve", "OLLAMA_HOST", "/v1/models", "prompt"} { + if strings.Contains(script, banned) { + t.Fatalf("script contains %s: %s", banned, script) + } + } + if !strings.Contains(buf.String(), `"engine": "ollama"`) || !strings.Contains(buf.String(), `"model": "mistral"`) { + t.Fatalf("receipt=%s", buf.String()) + } + + scripts = nil + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + scripts = append(scripts, script) + return "", context.DeadlineExceeded + } + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral", "--format", "text"}) + if err := cmd.Execute(); err == nil { + t.Fatal("curl failure must refuse even when OllamaInfo.Listening is true") + } + if len(runner.started) != 0 { + t.Fatal("curl failure started a process") + } + + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + return `{"models":[]}`, nil + } + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "api/ps") { + t.Fatalf("missing ps listing err=%v", err) + } +} + +func TestOllamaStopUnloadsWithoutKillingTheServer(t *testing.T) { + snap := testSnap() + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prevRunner := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prevRunner }) + var scripts []string + prevScript := runNodeScript + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + scripts = append(scripts, script) + return `{"models":[]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStopCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if len(runner.stopTargets) != 0 || len(runner.stopped) != 0 { + t.Fatalf("stop used the process killer: %#v", runner.stopTargets) + } + if len(scripts) != 1 || !strings.Contains(scripts[0], `"keep_alive":0`) { + t.Fatalf("scripts=%v", scripts) + } + for _, banned := range []string{"kill", "systemctl", "fuser"} { + if strings.Contains(scripts[0], banned) { + t.Fatalf("unload contains %s", banned) + } + } + + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + return `{"models":[{"name":"mistral","model":"mistral"}]}`, nil + } + cmd = modelStopCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "api/ps") { + t.Fatalf("still listed err=%v", err) + } +} + +func TestOllamaGenerationStopDoesNotReachProcessKill(t *testing.T) { + snap := generationStopSnapshot() + snap.Nodes[0].ResidentModels[0].Runtime = "ollama" + snap.Nodes[0].ResidentModels[0].Name = "mistral" + snap.Nodes[0].ResidentModels[0].Port = 11434 + want := modelinventory.FromSnapshot(snap, "daemon-cache").Instances[0] + if want.Engine != "ollama" { + t.Fatalf("engine=%s", want.Engine) + } + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prevRunner := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prevRunner }) + prevScript := runNodeScript + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + if strings.Contains(script, "kill") || strings.Contains(script, "systemctl") { + t.Fatalf("generation unload killed a process: %s", script) + } + return `{"models":[]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStopCmd() + cmd.SetOut(&bytes.Buffer{}) + if err := runModelStopGeneration(context.Background(), cmd, want.GenerationID, "test.sock", "text", runner); err != nil { + t.Fatal(err) + } + if len(runner.stopTargets) != 0 { + t.Fatalf("process kill targets=%#v", runner.stopTargets) + } +} diff --git a/cmd/axis/model_run_profile.go b/cmd/axis/model_run_profile.go index 7b3feb9b..6282cfec 100644 --- a/cmd/axis/model_run_profile.go +++ b/cmd/axis/model_run_profile.go @@ -15,6 +15,24 @@ func requireModelStartIdentity(cmd *cobra.Command) error { if cmd.Flags().Changed("from-plan") { return nil } + ollamaSet := cmd.Flags().Changed("ollama-model") + weightsSet := cmd.Flags().Changed("weights") + if ollamaSet && weightsSet { + return fmt.Errorf("--ollama-model and --weights are mutually exclusive") + } + if ollamaSet { + if !cmd.Flags().Changed("node") { + return fmt.Errorf(`required flag(s) "node" not set`) + } + name, err := cmd.Flags().GetString("ollama-model") + if err != nil { + return err + } + if strings.TrimSpace(name) == "" { + return fmt.Errorf("ollama model is required") + } + return nil + } var missing []string for _, name := range []string{"node", "weights", "port"} { if !cmd.Flags().Changed(name) { @@ -121,6 +139,35 @@ func applyChangedStartFlags(cmd *cobra.Command, profile *models.ModelRunProfile) profile.DeviceIndex = &value profile.IndexSource = models.IndexSourceNvidiaSMI } + if cmd.Flags().Changed("ollama-model") { + value, err := cmd.Flags().GetString("ollama-model") + if err != nil { + return err + } + profile.OllamaModel = strings.TrimSpace(value) + profile.Engine = models.EngineOllama + profile.ArtifactKind = models.ArtifactOllamaModelName + profile.WeightsPath = "" + profile.Volume = "" + profile.ToolName = "" + if profile.BindHost == "" { + profile.BindHost = "127.0.0.1" + } + } + if cmd.Flags().Changed("ollama-num-ctx") { + value, err := cmd.Flags().GetInt("ollama-num-ctx") + if err != nil { + return err + } + profile.OllamaNumCtx = &value + } + if cmd.Flags().Changed("ollama-keep-alive") { + value, err := cmd.Flags().GetString("ollama-keep-alive") + if err != nil { + return err + } + profile.OllamaKeepAlive = value + } return nil } diff --git a/internal/modellife/ollama.go b/internal/modellife/ollama.go new file mode 100644 index 00000000..52f7a1f7 --- /dev/null +++ b/internal/modellife/ollama.go @@ -0,0 +1,84 @@ +package modellife + +import ( + "encoding/json" + "fmt" + "strings" +) + +// OllamaLoadScript preloads a model on the Ollama server that is already +// listening on 127.0.0.1:11434. It does not exec ollama serve. The last +// command's stdout is GET /api/ps, which is the load fact. +func OllamaLoadScript(model, keepAlive string, numCtx *int) (string, error) { + model = strings.TrimSpace(model) + if model == "" { + return "", fmt.Errorf("ollama model is required") + } + if numCtx != nil && *numCtx < 1 { + return "", fmt.Errorf("ollama num_ctx must be >= 1") + } + payload := struct { + Model string `json:"model"` + Stream bool `json:"stream"` + KeepAlive string `json:"keep_alive,omitempty"` + Options map[string]int `json:"options,omitempty"` + }{ + Model: model, + Stream: false, + KeepAlive: keepAlive, + } + if numCtx != nil { + payload.Options = map[string]int{"num_ctx": *numCtx} + } + raw, err := json.Marshal(payload) + if err != nil { + return "", err + } + probe := "curl -fsS --max-time 5 http://127.0.0.1:11434/api/ps" + post := "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate -H 'Content-Type: application/json' -d " + shellSingleQuote(string(raw)) + return probe + " >/dev/null && " + post + " >/dev/null && " + probe, nil +} + +// OllamaUnloadScript asks the existing server to drop a model, then prints +// GET /api/ps. It does not kill comm=ollama or stop a supervisor unit. +func OllamaUnloadScript(model string) (string, error) { + model = strings.TrimSpace(model) + if model == "" { + return "", fmt.Errorf("ollama model is required") + } + payload := struct { + Model string `json:"model"` + KeepAlive int `json:"keep_alive"` + }{Model: model, KeepAlive: 0} + raw, err := json.Marshal(payload) + if err != nil { + return "", err + } + post := "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate -H 'Content-Type: application/json' -d " + shellSingleQuote(string(raw)) + probe := "curl -fsS --max-time 5 http://127.0.0.1:11434/api/ps" + return post + " >/dev/null && " + probe, nil +} + +// OllamaPSHasModel reports whether /api/ps lists the requested model name. +func OllamaPSHasModel(body, name string) (bool, error) { + var doc struct { + Models []struct { + Name string `json:"name"` + Model string `json:"model"` + } `json:"models"` + } + if err := json.Unmarshal([]byte(body), &doc); err != nil { + return false, fmt.Errorf("ollama /api/ps: %w", err) + } + name = strings.TrimSpace(name) + for _, model := range doc.Models { + if model.Name == name || model.Model == name { + return true, nil + } + } + return false, nil +} + +func shellSingleQuote(s string) string { + return "'" + strings.ReplaceAll(s, "'", `'"'"'`) + "'" +} diff --git a/internal/modellife/ollama_test.go b/internal/modellife/ollama_test.go new file mode 100644 index 00000000..e22007bf --- /dev/null +++ b/internal/modellife/ollama_test.go @@ -0,0 +1,115 @@ +package modellife + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestOllamaLoadScriptIsLoopbackGenerateWithoutPrompt(t *testing.T) { + script, err := OllamaLoadScript("mistral", "", nil) + if err != nil { + t.Fatal(err) + } + if strings.Count(script, "curl -fsS --max-time 5 http://127.0.0.1:11434/api/ps") < 2 { + t.Fatalf("load must probe /api/ps before and after POST:\n%s", script) + } + if !strings.Contains(script, "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate") { + t.Fatalf("post missing:\n%s", script) + } + if strings.Contains(script, "/v1/models") || strings.Contains(script, "ollama serve") || strings.Contains(script, "OLLAMA_HOST") || strings.Contains(script, "OLLAMA_KEEP_ALIVE") { + t.Fatalf("script widens or starts ollama:\n%s", script) + } + body := ollamaJSONBody(t, script) + if body["model"] != "mistral" || body["stream"] != false { + t.Fatalf("body=%v", body) + } + if _, ok := body["prompt"]; ok { + t.Fatalf("prompt must be omitted: %v", body) + } + if _, ok := body["keep_alive"]; ok { + t.Fatalf("unset keep_alive must be omitted: %v", body) + } + if _, ok := body["options"]; ok { + t.Fatalf("unset options must be omitted: %v", body) + } + for _, banned := range []string{"num_gpu", "main_gpu", "num_batch", "num_thread"} { + if strings.Contains(script, banned) { + t.Fatalf("script contains %s:\n%s", banned, script) + } + } +} + +func TestOllamaLoadScriptAddsOnlyKeepAliveAndNumCtx(t *testing.T) { + ctx := 2048 + script, err := OllamaLoadScript("mistral", "10m", &ctx) + if err != nil { + t.Fatal(err) + } + body := ollamaJSONBody(t, script) + if body["keep_alive"] != "10m" { + t.Fatalf("keep_alive=%v", body["keep_alive"]) + } + options, ok := body["options"].(map[string]any) + if !ok || len(options) != 1 || options["num_ctx"] != float64(2048) { + t.Fatalf("options=%v", body["options"]) + } +} + +func TestOllamaUnloadScriptKeepsAliveZeroAndDoesNotKill(t *testing.T) { + script, err := OllamaUnloadScript("mistral") + if err != nil { + t.Fatal(err) + } + if !strings.Contains(script, "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate") { + t.Fatalf("unload post missing:\n%s", script) + } + if !strings.Contains(script, "curl -fsS --max-time 5 http://127.0.0.1:11434/api/ps") { + t.Fatalf("unload probe missing:\n%s", script) + } + body := ollamaJSONBody(t, script) + if body["model"] != "mistral" || body["keep_alive"] != float64(0) { + t.Fatalf("body=%v", body) + } + if _, ok := body["stream"]; ok { + t.Fatalf("unload body=%v", body) + } + for _, banned := range []string{"kill", "systemctl", "fuser", "comm=ollama", "ollama serve", "num_gpu", "main_gpu"} { + if strings.Contains(script, banned) { + t.Fatalf("unload contains %s:\n%s", banned, script) + } + } +} + +func TestOllamaPSHasModelReadsAPIPs(t *testing.T) { + ok, err := OllamaPSHasModel(`{"models":[{"name":"mistral","model":"mistral"}]}`, "mistral") + if err != nil || !ok { + t.Fatalf("ok=%v err=%v", ok, err) + } + ok, err = OllamaPSHasModel(`{"models":[{"name":"other"}]}`, "mistral") + if err != nil || ok { + t.Fatalf("other ok=%v err=%v", ok, err) + } + if _, err := OllamaPSHasModel("not-json", "mistral"); err == nil { + t.Fatal("expected json error") + } +} + +func ollamaJSONBody(t *testing.T, script string) map[string]any { + t.Helper() + const marker = "-d '" + start := strings.Index(script, marker) + if start < 0 { + t.Fatalf("no json body in %s", script) + } + rest := script[start+len(marker):] + end := strings.Index(rest, "'") + if end < 0 { + t.Fatalf("unclosed body in %s", script) + } + var body map[string]any + if err := json.Unmarshal([]byte(rest[:end]), &body); err != nil { + t.Fatal(err) + } + return body +} diff --git a/internal/models/run_profile.go b/internal/models/run_profile.go index 6ff015b1..19da8016 100644 --- a/internal/models/run_profile.go +++ b/internal/models/run_profile.go @@ -23,12 +23,16 @@ const ( DeviceKindUnified = "unified" // DeviceKindCPU is a node with no discrete or unified accelerator. DeviceKindCPU = "cpu" - // EngineLlamaCpp is the only engine PR1 can launch. + // EngineLlamaCpp is the llama-server engine. EngineLlamaCpp = "llama.cpp" + // EngineOllama places a model on an Ollama server that is already listening. + EngineOllama = "ollama" // ToolLlamaServer is the observed tool name for llama.cpp. ToolLlamaServer = "llama-server" // ArtifactWeightsPath is a local weight file, not an Ollama model name. ArtifactWeightsPath = "weights-path" + // ArtifactOllamaModelName is an Ollama model name, not a GGUF path. + ArtifactOllamaModelName = "ollama-model-name" // IndexSourceNvidiaSMI is the only index source a llama-server pin may name. IndexSourceNvidiaSMI = "nvidia-smi" ) @@ -332,7 +336,11 @@ func (p ModelRunProfile) Validate() error { if p.Schema != ModelRunSchema { return fmt.Errorf("expected axis.model-run/v1") } - if p.Engine != EngineLlamaCpp { + switch p.Engine { + case EngineOllama: + return p.validateOllama() + case EngineLlamaCpp: + default: return fmt.Errorf("engine %q is not supported", p.Engine) } if p.BindHost != "127.0.0.1" { @@ -357,6 +365,37 @@ func (p ModelRunProfile) Validate() error { return nil } +func (p ModelRunProfile) validateOllama() error { + if strings.TrimSpace(p.OllamaModel) == "" { + return fmt.Errorf("ollama model is required") + } + if p.ArtifactKind != "" && p.ArtifactKind != ArtifactOllamaModelName { + return fmt.Errorf("ollama artifact must be %s", ArtifactOllamaModelName) + } + if strings.TrimSpace(p.WeightsPath) != "" && path.Clean(strings.TrimSpace(p.WeightsPath)) != "." { + return fmt.Errorf("--ollama-model and --weights are mutually exclusive") + } + if p.BindHost != "" && p.BindHost != "127.0.0.1" { + return fmt.Errorf("bind host must be 127.0.0.1") + } + if p.OllamaNumGPU != nil { + return fmt.Errorf("ollama num_gpu is not sent") + } + if p.DeviceIndex != nil || p.IndexSource != "" { + return fmt.Errorf("ollama main_gpu is not sent") + } + if p.MLXModel != "" || p.PrefillStepSize != nil || p.PromptCacheBytes != nil || p.KVBits != nil { + return fmt.Errorf("mlx fields are not supported") + } + if p.NGPULayers != nil || p.NGPULayersMode != "" || p.ContextTokens != nil || p.BatchSize != nil || p.UBatchSize != nil || p.Threads != nil { + return fmt.Errorf("only ollama launch fields are supported") + } + if p.OllamaNumCtx != nil && *p.OllamaNumCtx < 1 { + return fmt.Errorf("ollama num_ctx must be >= 1") + } + return nil +} + func (p ModelRunProfile) validateDevicePin() error { if p.DeviceIndex == nil && p.IndexSource == "" { return nil diff --git a/internal/models/run_profile_test.go b/internal/models/run_profile_test.go index 81dde732..b52fc5dc 100644 --- a/internal/models/run_profile_test.go +++ b/internal/models/run_profile_test.go @@ -158,3 +158,52 @@ func TestNewPlanProfileRecordsToolAndVolumeRefusals(t *testing.T) { t.Fatalf("unmeasured discrete device = %+v", profile) } } + +func TestValidateOllamaAllowsModelNameAndRefusesForeignFields(t *testing.T) { + profile := ModelRunProfile{ + Schema: ModelRunSchema, + Node: "storage", + Engine: EngineOllama, + ArtifactKind: ArtifactOllamaModelName, + OllamaModel: "mistral", + BindHost: "127.0.0.1", + } + if err := profile.Validate(); err != nil { + t.Fatal(err) + } + ctx := 1024 + profile.OllamaNumCtx = &ctx + profile.OllamaKeepAlive = "10m" + if err := profile.Validate(); err != nil { + t.Fatal(err) + } + gpus := 1 + profile.OllamaNumGPU = &gpus + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "num_gpu") { + t.Fatalf("num_gpu err=%v", err) + } + profile.OllamaNumGPU = nil + profile.MLXModel = "mlx" + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "mlx") { + t.Fatalf("mlx err=%v", err) + } + profile.MLXModel = "" + pin := 0 + profile.DeviceIndex = &pin + profile.IndexSource = IndexSourceNvidiaSMI + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "main_gpu") { + t.Fatalf("pin err=%v", err) + } + + llama := ModelRunProfile{ + Schema: ModelRunSchema, + Engine: EngineLlamaCpp, + BindHost: "127.0.0.1", + Port: 8081, + WeightsPath: "/mnt/models/a.gguf", + OllamaModel: "mistral", + } + if err := llama.Validate(); err == nil || !strings.Contains(err.Error(), "llama-server") { + t.Fatalf("llama with ollama field err=%v", err) + } +} From d169143429da7b757b588fa1ade6e184d9a1b585 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 16:38:39 -0400 Subject: [PATCH 4/9] fix(modellife): drop the unused PlanStart wrapper axis model start builds a profile and calls PlanStartProfile. PlanStart had no production caller, so the deadcode gate rejected the branch. Tests build that same default profile themselves. --- internal/modellife/plan.go | 18 --------------- internal/modellife/plan_profile_test.go | 2 +- internal/modellife/plan_test.go | 30 ++++++++++++++++++++----- 3 files changed, 25 insertions(+), 25 deletions(-) diff --git a/internal/modellife/plan.go b/internal/modellife/plan.go index 1504c204..758e4e7c 100644 --- a/internal/modellife/plan.go +++ b/internal/modellife/plan.go @@ -21,24 +21,6 @@ type StartPlan struct { Profile models.ModelRunProfile } -// PlanStart validates weights sit on a named local volume and that -// llama-server is an observed tool. Port must be explicit and valid. -// Optional launch fields stay unset, so the argv is the historical default. -func PlanStart(node models.NodeFacts, weights string, port int) (StartPlan, error) { - weights = path.Clean(strings.TrimSpace(weights)) - return PlanStartProfile(node, models.ModelRunProfile{ - Schema: models.ModelRunSchema, - Node: node.Name, - Engine: models.EngineLlamaCpp, - ToolName: models.ToolLlamaServer, - ArtifactKind: models.ArtifactWeightsPath, - WeightsPath: weights, - BindHost: "127.0.0.1", - Port: port, - PortSource: models.PortSourceExplicit, - }) -} - // PlanStartProfile validates profile against the observed node and derives argv. // A non-empty refusal list, a plan-default port, or an offload without measured // discrete VRAM returns an error and no argv. diff --git a/internal/modellife/plan_profile_test.go b/internal/modellife/plan_profile_test.go index ccf0875f..e2730564 100644 --- a/internal/modellife/plan_profile_test.go +++ b/internal/modellife/plan_profile_test.go @@ -99,7 +99,7 @@ func TestPlanStartProfileRefusesMainGPUThatWasNotObserved(t *testing.T) { } func TestPlanStartDefaultArgvIsExact(t *testing.T) { - plan, err := PlanStart(storageNode(), "/mnt/models/a.gguf", 8081) + plan, err := planStartDefault(storageNode(), "/mnt/models/a.gguf", 8081) if err != nil { t.Fatal(err) } diff --git a/internal/modellife/plan_test.go b/internal/modellife/plan_test.go index baedee6f..26091ea1 100644 --- a/internal/modellife/plan_test.go +++ b/internal/modellife/plan_test.go @@ -1,12 +1,30 @@ package modellife import ( + "path" "strings" "testing" "github.com/toasterbook88/axis/internal/models" ) +// planStartDefault is the no-flag llama-server profile. Production start +// builds that profile from flags and calls PlanStartProfile. +func planStartDefault(node models.NodeFacts, weights string, port int) (StartPlan, error) { + weights = path.Clean(strings.TrimSpace(weights)) + return PlanStartProfile(node, models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: node.Name, + Engine: models.EngineLlamaCpp, + ToolName: models.ToolLlamaServer, + ArtifactKind: models.ArtifactWeightsPath, + WeightsPath: weights, + BindHost: "127.0.0.1", + Port: port, + PortSource: models.PortSourceExplicit, + }) +} + func storageNode() models.NodeFacts { return models.NodeFacts{ Name: "storage", @@ -26,28 +44,28 @@ func storageNode() models.NodeFacts { func TestPlanStartRequiresValidPortAndWeights(t *testing.T) { n := storageNode() for _, port := range []int{-1, 0, 65536} { - if _, err := PlanStart(n, "/mnt/models/a.gguf", port); err == nil || !strings.Contains(err.Error(), "between 1 and 65535") { + if _, err := planStartDefault(n, "/mnt/models/a.gguf", port); err == nil || !strings.Contains(err.Error(), "between 1 and 65535") { t.Fatalf("port %d error = %v, want valid range error", port, err) } } - if _, err := PlanStart(n, "", 8081); err == nil { + if _, err := planStartDefault(n, "", 8081); err == nil { t.Fatal("expected error for missing weights") } } func TestPlanStartRequiresNamedLocalVolume(t *testing.T) { n := storageNode() - if _, err := PlanStart(n, "/mnt/nas/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "named local volume") { + if _, err := planStartDefault(n, "/mnt/nas/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "named local volume") { t.Fatalf("network path must be refused: %v", err) } - if _, err := PlanStart(n, "/not-a-volume/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "named local volume") { + if _, err := planStartDefault(n, "/not-a-volume/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "named local volume") { t.Fatalf("unknown path must be refused: %v", err) } } func TestPlanStartBuildsLlamaServerArgv(t *testing.T) { n := storageNode() - p, err := PlanStart(n, "/mnt/models/a.gguf", 8081) + p, err := planStartDefault(n, "/mnt/models/a.gguf", 8081) if err != nil { t.Fatal(err) } @@ -66,7 +84,7 @@ func TestPlanStartBuildsLlamaServerArgv(t *testing.T) { func TestPlanStartRequiresLlamaServerTool(t *testing.T) { n := storageNode() n.Tools = nil - if _, err := PlanStart(n, "/mnt/models/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "llama-server") { + if _, err := planStartDefault(n, "/mnt/models/a.gguf", 8081); err == nil || !strings.Contains(err.Error(), "llama-server") { t.Fatalf("expected missing tool error, got %v", err) } } From eee0e927863dcba40cf3dc0c44bc6794ba79b488 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 17:43:33 -0400 Subject: [PATCH 5/9] feat(model): start MLX on unified memory with a matching stop guard mlx_lm.server is the observed launch tool. The stop guard matches that basename or the argv sequence -m then mlx_lm.server. It does not kill comm=python or comm=mlx_lm. Port-only stops stay on the llama-server check. --- cmd/axis/model.go | 172 ++++++++++++++- cmd/axis/model_mlx_test.go | 323 ++++++++++++++++++++++++++++ cmd/axis/model_run_profile.go | 62 ++++++ internal/facts/remote_bundle.go | 4 +- internal/facts/tools.go | 1 + internal/modellife/mlx.go | 88 ++++++++ internal/modellife/mlx_test.go | 107 +++++++++ internal/modellife/stop.go | 3 + internal/models/run_profile.go | 42 ++++ internal/models/run_profile_test.go | 29 +++ 10 files changed, 820 insertions(+), 11 deletions(-) create mode 100644 cmd/axis/model_mlx_test.go create mode 100644 internal/modellife/mlx.go create mode 100644 internal/modellife/mlx_test.go diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 61ba77e2..7023b13e 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -98,12 +98,13 @@ func modelPlanCmd() *cobra.Command { } func modelStartCmd() *cobra.Command { - var node, weights, cacheAddr, format, fromPlan, nGPULayers, ollamaModel, ollamaKeepAlive string - var port, ctxSize, batchSize, ubatchSize, threads, mainGPU, ollamaNumCtx int + var node, weights, cacheAddr, format, fromPlan, nGPULayers, ollamaModel, ollamaKeepAlive, mlxModel string + var port, ctxSize, batchSize, ubatchSize, threads, mainGPU, ollamaNumCtx, prefillStepSize, kvBits int + var promptCacheBytes int64 var live bool cmd := &cobra.Command{ Use: "start", - Short: "Start llama-server, or place an Ollama model on the server already listening", + Short: "Start llama-server, place an Ollama model, or start mlx_lm.server (its HTTP API is not for production)", SilenceUsage: true, PreRunE: func(cmd *cobra.Command, args []string) error { if err := validateOutputFormat(&format, "text", "json", "yaml")(cmd, args); err != nil { @@ -130,6 +131,10 @@ func modelStartCmd() *cobra.Command { cmd.Flags().StringVar(&ollamaModel, "ollama-model", "", "Ollama model name to load on 127.0.0.1:11434; replaces --weights") cmd.Flags().StringVar(&ollamaKeepAlive, "ollama-keep-alive", "", "Ollama keep_alive for this load; omitted unless set") cmd.Flags().IntVar(&ollamaNumCtx, "ollama-num-ctx", 0, "Ollama options.num_ctx; omitted unless set") + cmd.Flags().StringVar(&mlxModel, "mlx-model", "", "Local MLX weight directory on a named volume; replaces --weights. The MLX HTTP API is not for production") + cmd.Flags().IntVar(&prefillStepSize, "prefill-step-size", 0, "mlx_lm.server --prefill-step-size; omitted unless set") + cmd.Flags().Int64Var(&promptCacheBytes, "prompt-cache-bytes", 0, "mlx_lm.server --prompt-cache-bytes; omitted unless set") + cmd.Flags().IntVar(&kvBits, "kv-bits", 0, "mlx_lm.server --kv-bits; omitted unless set") cmd.Flags().StringVar(&cacheAddr, "cache-addr", api.DefaultAddr(), "Address of the local AXIS daemon cache") cmd.Flags().BoolVar(&live, "live", false, "Bypass daemon cache and perform live fleet discovery") cmd.Flags().StringVar(&format, "format", "text", "Start operation receipt format: text, json, or yaml") @@ -141,7 +146,7 @@ func modelStopCmd() *cobra.Command { var port int cmd := &cobra.Command{ Use: "stop [generation-id]", - Short: "Stop an observed llama-server generation, unload an Ollama model, or use legacy node/port flags", + Short: "Stop an observed llama-server or MLX generation, unload an Ollama model, or use legacy node/port flags", Args: cobra.MaximumNArgs(1), SilenceUsage: true, PreRunE: validateOutputFormat(&format, "text", "json", "yaml"), @@ -694,6 +699,10 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st } } + if profile.Engine == models.EngineMLX { + return placeMLXModel(ctx, cmd, nf, cfgNode, profile, source, snap, startedAt, format) + } + plan, err := modellife.PlanStartProfile(nf, profile) if err != nil { return err @@ -841,7 +850,7 @@ func runModelStopGeneration(ctx context.Context, cmd *cobra.Command, generationI if instance.Engine == models.EngineOllama { return stopOllamaGeneration(ctx, cmd, snap, instance, format, startedAt) } - if instance.Engine != "llama.cpp" { + if instance.Engine != models.EngineLlamaCpp && instance.Engine != models.EngineMLX { return fmt.Errorf("model generation %s uses unsupported stop engine %q", generationID, instance.Engine) } target := modellife.StopTarget{ @@ -855,6 +864,9 @@ func runModelStopGeneration(ctx context.Context, cmd *cobra.Command, generationI SupervisorUnit: instance.SupervisorUnit, GPUIndices: append([]int(nil), instance.GPUIndices...), } + if instance.Engine == models.EngineMLX { + target.Engine = models.EngineMLX + } if err := target.Validate(); err != nil { return fmt.Errorf("model generation %s has incomplete stop evidence: %w", generationID, err) } @@ -1541,15 +1553,19 @@ func shellQuery(port int, req modellife.QueryRequest) (string, error) { } func shellStart(argv []string, port int) string { + return shellStartLabeled("llama-server", argv, port) +} + +func shellStartLabeled(label string, argv []string, port int) string { quoted := make([]string, len(argv)) for i, a := range argv { quoted[i] = shellQuote(a) } return shellListenerLookup(port) + fmt.Sprintf( "if test -n \"$_axis_pids\"; then "+ - "echo \"refusing to start llama-server: port %d already has listener pid(s) $_axis_pids\" >&2; exit 1; fi; "+ + "echo \"refusing to start %s: port %d already has listener pid(s) $_axis_pids\" >&2; exit 1; fi; "+ "nohup %s >/dev/null 2>&1 &", - port, strings.Join(quoted, " "), + label, port, strings.Join(quoted, " "), ) } @@ -1568,9 +1584,13 @@ func shellStopTarget(target modellife.StopTarget) string { supervisorCmd = fmt.Sprintf("systemctl stop %s 2>/dev/null || true; ", shellQuote(unit)) } } + ownerGuard := shellLlamaServerOwnerGuard(port) + if target.Engine == models.EngineMLX { + ownerGuard = shellMLXOwnerGuard(port) + } return shellListenerLookup(port) + "if test -z \"$_axis_pids\"; then echo '" + modelStopMarker + "not_running'; exit 0; fi; " + - shellLlamaServerOwnerGuard(port) + + ownerGuard + shellGenerationGuard(target) + supervisorCmd + killCmd + @@ -1607,6 +1627,31 @@ func shellProbe(port int) string { ) } +func shellMLXProbe(port int) string { + return shellListenerLookup(port) + fmt.Sprintf( + "if test -z \"$_axis_pids\"; then echo \"no listener on port %d\" >&2; exit 1; fi; ", + port, + ) + shellMLXOwnerGuard(port) + fmt.Sprintf( + "curl -fsS --max-time 5 http://127.0.0.1:%d/v1/models >/dev/null", + port, + ) +} + +func shellMLXOwnerGuard(port int) string { + return fmt.Sprintf( + "if ! command -v ps >/dev/null 2>&1; then echo 'axis model requires ps to verify process ownership' >&2; echo '"+modelStopMarker+"inspection_unavailable' >&2; exit 127; fi; "+ + "for _axis_pid in $_axis_pids; do "+ + "case \"$_axis_pid\" in ''|*[!0-9]*) echo \"refusing invalid listener pid $_axis_pid\" >&2; exit 1;; esac; "+ + "_axis_comm=$(ps -p \"$_axis_pid\" -o comm=) || exit $?; _axis_comm=${_axis_comm##*/}; "+ + "if test \"$_axis_comm\" = mlx_lm.server; then continue; fi; "+ + "_axis_args=$(ps -p \"$_axis_pid\" -o args=) || exit $?; "+ + "if ! printf '%%s\\n' \"$_axis_args\" | awk 'BEGIN{ok=0} {for(i=1;i&2; echo '"+modelStopMarker+"wrong_owner' >&2; exit 1; fi; "+ + "done; ", + port, + ) +} + func shellLlamaServerOwnerGuard(port int) string { return fmt.Sprintf( "if ! command -v ps >/dev/null 2>&1; then echo 'axis model requires ps to verify process ownership' >&2; echo '"+modelStopMarker+"inspection_unavailable' >&2; exit 127; fi; "+ @@ -1642,10 +1687,119 @@ func runOnNode(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeC return err } -// runNodeScript is the local-or-SSH curl seam for Ollama. Tests replace it. +// runNodeScript is the local-or-SSH seam for Ollama and MLX. Tests replace it. // llama-server start and stop keep calling runOnNodeCapturing directly. var runNodeScript = runOnNodeCapturing +func mlxImportObserved(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig) (bool, error) { + for _, tool := range node.Tools { + if strings.EqualFold(tool.Name, models.ToolMLXServer) { + return false, nil + } + } + python := "" + for _, tool := range node.Tools { + if strings.EqualFold(tool.Name, "python3") && strings.TrimSpace(tool.Path) != "" { + python = tool.Path + break + } + } + if python == "" { + return false, nil + } + script := shellQuote(python) + " -c " + shellQuote("import mlx_lm") + " >/dev/null 2>&1 && echo axis-mlx-import:ok || echo axis-mlx-import:missing" + out, err := runNodeScript(ctx, node, cfgNode, script) + if err != nil { + return false, err + } + return strings.Contains(out, "axis-mlx-import:ok"), nil +} + +func probeMLXServer(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig, port int) error { + script := shellMLXProbe(port) + var last error + for i := 0; i < 10; i++ { + if _, err := runNodeScript(ctx, node, cfgNode, script); err == nil { + return nil + } else { + last = err + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(500 * time.Millisecond): + } + } + if last == nil { + last = fmt.Errorf("probe failed") + } + return last +} + +func placeMLXModel(ctx context.Context, cmd *cobra.Command, node models.NodeFacts, cfgNode *config.NodeConfig, profile models.ModelRunProfile, source string, snap *models.ClusterSnapshot, startedAt time.Time, format string) error { + if err := profile.Validate(); err != nil { + return err + } + importOK, importErr := mlxImportObserved(ctx, node, cfgNode) + if importErr != nil { + return importErr + } + argv, err := modellife.MLXArgv(node, profile, importOK) + if err != nil { + return err + } + volume, _ := models.NamedLocalVolume(node, profile.MLXModel) + dev := models.ObserveLaunchDevice(node) + receipt := models.ModelOperationReceipt{ + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStart, + Status: models.ModelOperationCompleted, + Disposition: "started", + Node: node.Name, + Engine: models.EngineMLX, + Port: profile.Port, + Model: path.Base(profile.MLXModel), + Weights: profile.MLXModel, + Volume: volume, + Executable: argv[0], + SnapshotSource: source, + StartedAt: startedAt, + CompletedAt: time.Now().UTC(), + DeviceKind: dev.Kind, + PortSource: profile.PortSource, + } + if snap != nil { + receipt.SnapshotAt = snap.Timestamp + if snap.Publication != nil { + receipt.PublicationID = snap.Publication.ID + } + } + if _, startErr := runNodeScript(ctx, node, cfgNode, shellStartLabeled("mlx_lm.server", argv, profile.Port)); startErr != nil { + receipt.Status = models.ModelOperationFailed + receipt.Disposition = "failed" + receipt.Error = startErr.Error() + receipt.CompletedAt = time.Now().UTC() + _ = writeModelStartReceipt(cmd, receipt, format) + return fmt.Errorf("mlx start failed: %w", startErr) + } + if probeErr := probeMLXServer(ctx, node, cfgNode, profile.Port); probeErr != nil { + receipt.Status = models.ModelOperationFailed + receipt.Disposition = "failed" + receipt.Error = probeErr.Error() + receipt.CompletedAt = time.Now().UTC() + _ = writeModelStartReceipt(cmd, receipt, format) + return fmt.Errorf("started but probe failed: %w", probeErr) + } + receipt.CompletedAt = time.Now().UTC() + if writeErr := writeModelStartReceipt(cmd, receipt, format); writeErr != nil { + return writeErr + } + cacheAddr, _ := cmd.Flags().GetString("cache-addr") + warnModelDaemonRefresh(cmd, cacheAddr, "manual") + return nil +} + func placeOllamaModel(ctx context.Context, cmd *cobra.Command, node models.NodeFacts, cfgNode *config.NodeConfig, profile models.ModelRunProfile, source string, snap *models.ClusterSnapshot, startedAt time.Time, format string) error { if err := profile.Validate(); err != nil { return err diff --git a/cmd/axis/model_mlx_test.go b/cmd/axis/model_mlx_test.go new file mode 100644 index 00000000..8537f5ab --- /dev/null +++ b/cmd/axis/model_mlx_test.go @@ -0,0 +1,323 @@ +package main + +import ( + "bytes" + "context" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/config" + "github.com/toasterbook88/axis/internal/modelinventory" + "github.com/toasterbook88/axis/internal/modellife" + "github.com/toasterbook88/axis/internal/models" +) + +func TestMLXHelpSaysHTTPAPIIsNotForProduction(t *testing.T) { + cmd := modelStartCmd() + if !strings.Contains(cmd.Short, "not for production") { + t.Fatalf("short=%q", cmd.Short) + } + flag := cmd.Flags().Lookup("mlx-model") + if flag == nil || !strings.Contains(flag.Usage, "not for production") { + t.Fatalf("flag=%v", flag) + } +} + +func TestMLXModelIsMutuallyExclusiveWithWeightsAndOllama(t *testing.T) { + for _, args := range [][]string{ + {"--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8080", "--mlx-model", "/mnt/models/qwen"}, + {"--node", "storage", "--ollama-model", "mistral", "--mlx-model", "/mnt/models/qwen"}, + } { + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs(args) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "mutually exclusive") { + t.Fatalf("args=%v err=%v", args, err) + } + } +} + +func TestMLXStartUsesObservedServerAndDoesNotUseLlamaRunner(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified + snap.Nodes[0].Tools = append(snap.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{} + prevRunner := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prevRunner }) + + var scripts []string + prevScript := runNodeScript + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + scripts = append(scripts, script) + if strings.Contains(script, "import mlx_lm") { + t.Fatal("observed mlx_lm.server must not probe python import") + } + return "", nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{ + "--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080", + "--prefill-step-size", "2", "--prompt-cache-bytes", "4096", "--kv-bits", "4", + "--format", "json", + }) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if len(runner.started) != 0 || len(runner.probed) != 0 { + t.Fatalf("llama runner used: started=%v probed=%v", runner.started, runner.probed) + } + if len(scripts) != 2 { + t.Fatalf("scripts=%d", len(scripts)) + } + start := scripts[0] + for _, want := range []string{ + "nohup", "/usr/local/bin/mlx_lm.server", "--model", "/mnt/models/qwen", + "--port", "8080", "--host", "127.0.0.1", + "--prefill-step-size", "2", "--prompt-cache-bytes", "4096", "--kv-bits", "4", + } { + if !strings.Contains(start, want) { + t.Fatalf("start missing %s: %s", want, start) + } + } + for _, banned := range []string{"llama-server", "-ngl", "num_gpu", "iogpu", "--quant"} { + if strings.Contains(start, banned) { + t.Fatalf("start contains %s: %s", banned, start) + } + } + probe := scripts[1] + if !strings.Contains(probe, "/v1/models") || !strings.Contains(probe, "mlx_lm.server") || strings.Contains(probe, "not llama-server") { + t.Fatalf("probe=%s", probe) + } + if !strings.Contains(buf.String(), `"engine": "mlx"`) || !strings.Contains(buf.String(), `"device_kind": "unified"`) { + t.Fatalf("receipt=%s", buf.String()) + } + + scripts = nil + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + plain := scripts[0] + if strings.Contains(plain, "prefill") || strings.Contains(plain, "prompt-cache") || strings.Contains(plain, "kv-bits") { + t.Fatalf("unset flags leaked: %s", plain) + } +} + +func TestMLXStartUsesPythonModuleOnlyAfterObservedImport(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified + snap.Nodes[0].Tools = []models.ToolInfo{ + {Name: "mlx_lm", Path: "/usr/local/bin/mlx_lm"}, + {Name: "python3", Path: "/usr/bin/python3"}, + } + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + prevScript := runNodeScript + var scripts []string + runNodeScript = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + scripts = append(scripts, script) + if strings.Contains(script, "import mlx_lm") { + if !strings.Contains(script, "/usr/bin/python3") { + t.Fatalf("import did not use the observed python: %s", script) + } + return "axis-mlx-import:ok", nil + } + return "", nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080", "--format", "json"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + joined := strings.Join(scripts, "\n") + if !strings.Contains(joined, "/usr/bin/python3") || !strings.Contains(joined, "-m") || !strings.Contains(joined, "mlx_lm.server") { + t.Fatalf("scripts=%s", joined) + } + if strings.Contains(scripts[1], "/usr/local/bin/mlx_lm ") || strings.Contains(scripts[1], "mlx_lm server") { + t.Fatalf("console tool was launched: %s", scripts[1]) + } +} + +func TestMLXStartRefusesHubDiscreteFileAndOccupiedPort(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified + snap.Nodes[0].Tools = append(snap.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + t.Fatal("refused start executed a script") + return "", nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + refuses := []struct { + args []string + want string + }{ + {[]string{"--node", "storage", "--mlx-model", "mlx-community/Qwen", "--port", "8080"}, "hub"}, + {[]string{"--node", "storage", "--mlx-model", "/mnt/models/a.gguf", "--port", "8080"}, "directory"}, + } + for _, tc := range refuses { + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs(tc.args) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("args=%v err=%v", tc.args, err) + } + } + + discrete := testSnap() + discrete.Nodes[0].Tools = append(discrete.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + stubModelSnapshot(t, discrete) + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080"}) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "unified") { + t.Fatalf("discrete err=%v", err) + } + + occupied := testSnap() + occupied.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified + occupied.Nodes[0].Tools = append(occupied.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + occupied.Nodes[0].ResidentModels = []models.ResidentModel{{Name: "busy", Runtime: "llama.cpp", Port: 8080}} + stubModelSnapshot(t, occupied) + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080"}) + if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "occupied") { + t.Fatalf("occupied err=%v", err) + } +} + +func TestMLXGenerationStopSetsEngineGuard(t *testing.T) { + snap := generationStopSnapshot() + snap.Nodes[0].ResidentModels[0].Runtime = "mlx" + snap.Nodes[0].ResidentModels[0].Name = "qwen" + snap.Nodes[0].ResidentModels[0].Executable = "/usr/local/bin/mlx_lm.server" + want := modelinventory.FromSnapshot(snap, "daemon-cache").Instances[0] + if want.Engine != models.EngineMLX { + t.Fatalf("engine=%s", want.Engine) + } + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runner := &fakeModelRunner{stopDisposition: modelStopStopped} + cmd := modelStopCmd() + cmd.SetOut(&bytes.Buffer{}) + if err := runModelStopGeneration(context.Background(), cmd, want.GenerationID, "test.sock", "text", runner); err != nil { + t.Fatal(err) + } + if len(runner.stopTargets) != 1 || runner.stopTargets[0].Engine != models.EngineMLX { + t.Fatalf("targets=%#v", runner.stopTargets) + } + script := shellStopTarget(runner.stopTargets[0]) + if strings.Contains(script, "not llama-server") || !strings.Contains(script, "mlx_lm.server") { + t.Fatalf("script=%s", script) + } + if !strings.Contains(script, want.ProcessStartToken) || !strings.Contains(script, `kill -KILL "4242"`) { + t.Fatalf("generation evidence missing: %s", script) + } +} + +func TestShellStopMLXGuardMatchesServerNotPythonOrConsole(t *testing.T) { + cases := []struct { + name string + comm string + args string + allow bool + engine string + wantOut string + }{ + {name: "python", comm: "python", args: "python app.py", wantOut: "wrong_owner"}, + {name: "mlx console", comm: "mlx_lm", args: "mlx_lm server --port 8080", wantOut: "wrong_owner"}, + {name: "basename", comm: "/usr/local/bin/mlx_lm.server", args: "mlx_lm.server --model /mnt/models/qwen", allow: true}, + {name: "module", comm: "python", args: "python -m mlx_lm.server --model /mnt/models/qwen", allow: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + out, alive, err := runMLXStopScript(t, tc.comm, tc.args, models.EngineMLX) + if tc.allow { + if err != nil || alive { + t.Fatalf("err=%v alive=%v out=%s", err, alive, out) + } + if !strings.Contains(out, modelStopMarker+"stopped") { + t.Fatalf("out=%s", out) + } + return + } + var exitErr *exec.ExitError + if err == nil || !errors.As(err, &exitErr) || exitErr.ExitCode() != 1 || !alive { + t.Fatalf("err=%v alive=%v out=%s", err, alive, out) + } + if !strings.Contains(out, tc.wantOut) || strings.Contains(out, modelStopMarker+"stopped") { + t.Fatalf("out=%s", out) + } + }) + } + + out, alive, err := runMLXStopScript(t, "python", "python -m mlx_lm.server", "") + var exitErr *exec.ExitError + if err == nil || !errors.As(err, &exitErr) || !alive || !strings.Contains(out, "not llama-server") { + t.Fatalf("legacy python stop err=%v alive=%v out=%s", err, alive, out) + } +} + +func runMLXStopScript(t *testing.T, comm, args, engine string) (string, bool, error) { + t.Helper() + targetProc := exec.Command("sleep", "30") + if err := targetProc.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = targetProc.Process.Kill() + _, _ = targetProc.Process.Wait() + }) + pid := fmt.Sprintf("%d", targetProc.Process.Pid) + dir := t.TempDir() + writeModelTestExecutable(t, filepath.Join(dir, "fuser"), "#!/bin/sh\nprintf '%s\\n' "+pid+"\n") + writeModelTestExecutable(t, filepath.Join(dir, "lsof"), "#!/bin/sh\nprintf '%s\\n' "+pid+"\n") + writeModelTestExecutable(t, filepath.Join(dir, "ps"), fmt.Sprintf(`#!/bin/sh +case "$*" in + *"-o args="*) printf '%%s\n' %s ;; + *) printf '%%s\n' %s ;; +esac +`, shellQuote(args), shellQuote(comm))) + cmd := exec.Command("/bin/sh", "-c", shellStopTarget(modellife.StopTarget{Port: 8080, Engine: engine})) + cmd.Env = []string{"PATH=" + dir + ":/usr/bin:/bin"} + out, err := cmd.CombinedOutput() + return string(out), processStillRunning(targetProc.Process.Pid), err +} + +func processStillRunning(pid int) bool { + data, err := os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid)) + if err != nil { + return false + } + text := string(data) + i := strings.LastIndex(text, ")") + if i < 0 || i+2 >= len(text) { + return true + } + state := text[i+2] + return state != 'Z' && state != 'X' +} diff --git a/cmd/axis/model_run_profile.go b/cmd/axis/model_run_profile.go index 6282cfec..55461ff3 100644 --- a/cmd/axis/model_run_profile.go +++ b/cmd/axis/model_run_profile.go @@ -16,10 +16,17 @@ func requireModelStartIdentity(cmd *cobra.Command) error { return nil } ollamaSet := cmd.Flags().Changed("ollama-model") + mlxSet := cmd.Flags().Changed("mlx-model") weightsSet := cmd.Flags().Changed("weights") if ollamaSet && weightsSet { return fmt.Errorf("--ollama-model and --weights are mutually exclusive") } + if mlxSet && weightsSet { + return fmt.Errorf("--mlx-model and --weights are mutually exclusive") + } + if ollamaSet && mlxSet { + return fmt.Errorf("--mlx-model and --ollama-model are mutually exclusive") + } if ollamaSet { if !cmd.Flags().Changed("node") { return fmt.Errorf(`required flag(s) "node" not set`) @@ -33,6 +40,25 @@ func requireModelStartIdentity(cmd *cobra.Command) error { } return nil } + if mlxSet { + var missing []string + for _, name := range []string{"node", "port"} { + if !cmd.Flags().Changed(name) { + missing = append(missing, name) + } + } + if len(missing) > 0 { + return fmt.Errorf(`required flag(s) "%s" not set`, strings.Join(missing, `", "`)) + } + name, err := cmd.Flags().GetString("mlx-model") + if err != nil { + return err + } + if strings.TrimSpace(name) == "" { + return fmt.Errorf("mlx model is required") + } + return nil + } var missing []string for _, name := range []string{"node", "weights", "port"} { if !cmd.Flags().Changed(name) { @@ -168,6 +194,42 @@ func applyChangedStartFlags(cmd *cobra.Command, profile *models.ModelRunProfile) } profile.OllamaKeepAlive = value } + if cmd.Flags().Changed("mlx-model") { + value, err := cmd.Flags().GetString("mlx-model") + if err != nil { + return err + } + profile.MLXModel = strings.TrimSpace(value) + profile.Engine = models.EngineMLX + profile.ArtifactKind = models.ArtifactMLXModelDir + profile.ToolName = models.ToolMLXServer + profile.WeightsPath = "" + profile.Volume = "" + if profile.BindHost == "" { + profile.BindHost = "127.0.0.1" + } + } + if cmd.Flags().Changed("prefill-step-size") { + value, err := cmd.Flags().GetInt("prefill-step-size") + if err != nil { + return err + } + profile.PrefillStepSize = &value + } + if cmd.Flags().Changed("prompt-cache-bytes") { + value, err := cmd.Flags().GetInt64("prompt-cache-bytes") + if err != nil { + return err + } + profile.PromptCacheBytes = &value + } + if cmd.Flags().Changed("kv-bits") { + value, err := cmd.Flags().GetInt("kv-bits") + if err != nil { + return err + } + profile.KVBits = &value + } return nil } diff --git a/internal/facts/remote_bundle.go b/internal/facts/remote_bundle.go index b89a66de..d58ba7ad 100644 --- a/internal/facts/remote_bundle.go +++ b/internal/facts/remote_bundle.go @@ -85,7 +85,7 @@ if command -v lsblk >/dev/null 2>&1; then fi # Tools: path + version (same tool set as defaultToolDefs) -for t in go python3 git jq nix docker ollama mlx_lm llama-cli llama-server node swift cargo gcc; do +for t in go python3 git jq nix docker ollama mlx_lm mlx_lm.server llama-cli llama-server node swift cargo gcc; do p=$(command -v "$t" 2>/dev/null) if [ -n "$p" ]; then printf 'tool_%s=%s\n' "$t" "$p" @@ -97,7 +97,7 @@ for t in go python3 git jq nix docker ollama mlx_lm llama-cli llama-server node nix) v=$("$p" --version 2>/dev/null | head -1) ;; docker) v=$("$p" --version 2>/dev/null | head -1) ;; ollama) v=$("$p" --version 2>/dev/null | head -1) ;; - mlx_lm) v=$("$p" --help 2>/dev/null | head -1) ;; + mlx_lm|mlx_lm.server) v=$("$p" --help 2>/dev/null | head -1) ;; llama-cli|llama-server) v=$("$p" --version 2>/dev/null | head -1) ;; node) v=$("$p" --version 2>/dev/null | head -1) ;; swift) v=$("$p" --version 2>/dev/null | head -1) ;; diff --git a/internal/facts/tools.go b/internal/facts/tools.go index 6e708aad..8b30cf4c 100644 --- a/internal/facts/tools.go +++ b/internal/facts/tools.go @@ -342,6 +342,7 @@ func defaultToolDefs() []toolDef { {name: "docker", class: models.ToolClassContainer, versionCmd: "docker --version"}, {name: "ollama", class: models.ToolClassAICLI, versionCmd: "ollama --version"}, {name: "mlx_lm", class: models.ToolClassAICLI, versionCmd: "mlx_lm --help"}, + {name: "mlx_lm.server", class: models.ToolClassAICLI, versionCmd: "mlx_lm.server --help"}, {name: "llama-cli", class: models.ToolClassAICLI, versionCmd: "llama-cli --version"}, {name: "llama-server", class: models.ToolClassAICLI, versionCmd: "llama-server --version"}, {name: "node", class: models.ToolClassRuntime, versionCmd: "node --version"}, diff --git a/internal/modellife/mlx.go b/internal/modellife/mlx.go new file mode 100644 index 00000000..a52d85ed --- /dev/null +++ b/internal/modellife/mlx.go @@ -0,0 +1,88 @@ +package modellife + +import ( + "fmt" + "path" + "strconv" + "strings" + + "github.com/toasterbook88/axis/internal/models" +) + +// MLXArgv projects an mlx_lm.server command. importOK is the observed +// `python3 -c "import mlx_lm"` result and is used only when mlx_lm.server +// itself is not an observed tool. The mlx_lm console tool is not a server. +func MLXArgv(node models.NodeFacts, profile models.ModelRunProfile, importOK bool) ([]string, error) { + if profile.Engine != models.EngineMLX { + return nil, fmt.Errorf("engine %q is not mlx", profile.Engine) + } + if profile.BindHost != "127.0.0.1" { + return nil, fmt.Errorf("bind host must be 127.0.0.1") + } + if profile.Port < 1 || profile.Port > 65535 { + return nil, fmt.Errorf("port must be between 1 and 65535") + } + if node.Resources == nil || node.Resources.MemoryTopology != models.MemoryTopologyUnified { + return nil, fmt.Errorf("mlx requires unified memory") + } + model, err := mlxModelDir(node, profile.MLXModel) + if err != nil { + return nil, err + } + if profile.PrefillStepSize != nil && *profile.PrefillStepSize < 1 { + return nil, fmt.Errorf("prefill-step-size must be >= 1") + } + if profile.PromptCacheBytes != nil && *profile.PromptCacheBytes < 1 { + return nil, fmt.Errorf("prompt-cache-bytes must be >= 1") + } + if profile.KVBits != nil && *profile.KVBits < 1 { + return nil, fmt.Errorf("kv-bits must be >= 1") + } + argv, err := mlxArgvPrefix(node, importOK) + if err != nil { + return nil, err + } + argv = append(argv, "--model", model, "--port", strconv.Itoa(profile.Port), "--host", "127.0.0.1") + if profile.PrefillStepSize != nil { + argv = append(argv, "--prefill-step-size", strconv.Itoa(*profile.PrefillStepSize)) + } + if profile.PromptCacheBytes != nil { + argv = append(argv, "--prompt-cache-bytes", strconv.FormatInt(*profile.PromptCacheBytes, 10)) + } + if profile.KVBits != nil { + argv = append(argv, "--kv-bits", strconv.Itoa(*profile.KVBits)) + } + return argv, nil +} + +func mlxArgvPrefix(node models.NodeFacts, importOK bool) ([]string, error) { + if bin := toolPath(node, models.ToolMLXServer); bin != "" { + return []string{bin}, nil + } + if hasTool(node, models.ToolMLXServer) { + return []string{models.ToolMLXServer}, nil + } + python := toolPath(node, "python3") + if importOK && python != "" { + return []string{python, "-m", "mlx_lm.server"}, nil + } + return nil, fmt.Errorf("node %s has no observed mlx_lm.server tool", node.Name) +} + +func mlxModelDir(node models.NodeFacts, model string) (string, error) { + model = path.Clean(strings.TrimSpace(model)) + if model == "" || model == "." { + return "", fmt.Errorf("mlx model is required") + } + if !path.IsAbs(model) { + return "", fmt.Errorf("mlx hub repo id %q is not a local directory", model) + } + switch strings.ToLower(path.Ext(model)) { + case ".gguf", ".safetensors", ".bin", ".pt": + return "", fmt.Errorf("mlx model %s must be a local directory", model) + } + if _, ok := models.NamedLocalVolume(node, model); !ok { + return "", fmt.Errorf("mlx model %s is not on a named local volume", model) + } + return model, nil +} diff --git a/internal/modellife/mlx_test.go b/internal/modellife/mlx_test.go new file mode 100644 index 00000000..a6853470 --- /dev/null +++ b/internal/modellife/mlx_test.go @@ -0,0 +1,107 @@ +package modellife + +import ( + "reflect" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/models" +) + +func unifiedMLXNode() models.NodeFacts { + node := storageNode() + node.Resources.MemoryTopology = models.MemoryTopologyUnified + node.Tools = append(node.Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + return node +} + +func TestMLXArgvUsesObservedServerBinary(t *testing.T) { + node := unifiedMLXNode() + profile := models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: node.Name, + Engine: "mlx", + MLXModel: "/mnt/models/qwen", + BindHost: "127.0.0.1", + Port: 8080, + } + argv, err := MLXArgv(node, profile, false) + if err != nil { + t.Fatal(err) + } + want := []string{"/usr/local/bin/mlx_lm.server", "--model", "/mnt/models/qwen", "--port", "8080", "--host", "127.0.0.1"} + if !reflect.DeepEqual(argv, want) { + t.Fatalf("argv=%#v", argv) + } +} + +func TestMLXArgvUsesPythonModuleWhenImportWasObserved(t *testing.T) { + node := storageNode() + node.Resources.MemoryTopology = models.MemoryTopologyUnified + node.Tools = []models.ToolInfo{ + {Name: "mlx_lm", Path: "/usr/local/bin/mlx_lm"}, + {Name: "python3", Path: "/usr/bin/python3"}, + } + profile := models.ModelRunProfile{ + Schema: models.ModelRunSchema, Engine: models.EngineMLX, MLXModel: "/mnt/models/qwen", + BindHost: "127.0.0.1", Port: 8080, + } + argv, err := MLXArgv(node, profile, true) + if err != nil { + t.Fatal(err) + } + want := []string{"/usr/bin/python3", "-m", "mlx_lm.server", "--model", "/mnt/models/qwen", "--port", "8080", "--host", "127.0.0.1"} + if !reflect.DeepEqual(argv, want) { + t.Fatalf("argv=%#v", argv) + } + if _, err := MLXArgv(node, profile, false); err == nil || !strings.Contains(err.Error(), "mlx_lm.server") { + t.Fatalf("mlx_lm console tool err=%v", err) + } +} + +func TestMLXArgvRefusesHubDiscreteAndFileAndOmitsUnsetFlags(t *testing.T) { + node := unifiedMLXNode() + base := models.ModelRunProfile{ + Schema: models.ModelRunSchema, Engine: models.EngineMLX, BindHost: "127.0.0.1", Port: 8080, + } + hub := base + hub.MLXModel = "mlx-community/Qwen" + if _, err := MLXArgv(node, hub, false); err == nil || !strings.Contains(err.Error(), "hub") { + t.Fatalf("hub err=%v", err) + } + file := base + file.MLXModel = "/mnt/models/a.gguf" + if _, err := MLXArgv(node, file, false); err == nil || !strings.Contains(err.Error(), "directory") { + t.Fatalf("file err=%v", err) + } + discrete := node + copied := *node.Resources + discrete.Resources = &copied + discrete.Resources.MemoryTopology = "" + dir := base + dir.MLXModel = "/mnt/models/qwen" + if _, err := MLXArgv(discrete, dir, false); err == nil || !strings.Contains(err.Error(), "unified") { + t.Fatalf("discrete err=%v", err) + } + step, cache, bits := 2, int64(4096), 4 + flagged := dir + flagged.PrefillStepSize = &step + flagged.PromptCacheBytes = &cache + flagged.KVBits = &bits + argv, err := MLXArgv(node, flagged, false) + if err != nil { + t.Fatal(err) + } + got := strings.Join(argv, " ") + if !strings.Contains(got, "--prefill-step-size 2") || !strings.Contains(got, "--prompt-cache-bytes 4096") || !strings.Contains(got, "--kv-bits 4") { + t.Fatalf("argv=%v", argv) + } + plain, err := MLXArgv(node, dir, false) + if err != nil { + t.Fatal(err) + } + plainGot := strings.Join(plain, " ") + if strings.Contains(plainGot, "prefill") || strings.Contains(plainGot, "kv-bits") || strings.Contains(plainGot, "prompt-cache") { + t.Fatalf("unset flags leaked: %v", plain) + } +} diff --git a/internal/modellife/stop.go b/internal/modellife/stop.go index a44ca352..cc759388 100644 --- a/internal/modellife/stop.go +++ b/internal/modellife/stop.go @@ -18,6 +18,9 @@ type StopTarget struct { SupervisorType string SupervisorUnit string GPUIndices []int + // Engine selects the process-owner check. Empty keeps the llama-server + // comm check, including port-only legacy stops. + Engine string } func (t StopTarget) IsGenerationBound() bool { diff --git a/internal/models/run_profile.go b/internal/models/run_profile.go index 19da8016..428a1f9c 100644 --- a/internal/models/run_profile.go +++ b/internal/models/run_profile.go @@ -27,12 +27,18 @@ const ( EngineLlamaCpp = "llama.cpp" // EngineOllama places a model on an Ollama server that is already listening. EngineOllama = "ollama" + // EngineMLX starts mlx_lm.server on unified memory. + EngineMLX = "mlx" // ToolLlamaServer is the observed tool name for llama.cpp. ToolLlamaServer = "llama-server" + // ToolMLXServer is the observed mlx_lm.server binary, not the mlx_lm console tool. + ToolMLXServer = "mlx_lm.server" // ArtifactWeightsPath is a local weight file, not an Ollama model name. ArtifactWeightsPath = "weights-path" // ArtifactOllamaModelName is an Ollama model name, not a GGUF path. ArtifactOllamaModelName = "ollama-model-name" + // ArtifactMLXModelDir is a local MLX weight directory, not a Hub repo id. + ArtifactMLXModelDir = "mlx-model-dir" // IndexSourceNvidiaSMI is the only index source a llama-server pin may name. IndexSourceNvidiaSMI = "nvidia-smi" ) @@ -339,6 +345,8 @@ func (p ModelRunProfile) Validate() error { switch p.Engine { case EngineOllama: return p.validateOllama() + case EngineMLX: + return p.validateMLX() case EngineLlamaCpp: default: return fmt.Errorf("engine %q is not supported", p.Engine) @@ -396,6 +404,40 @@ func (p ModelRunProfile) validateOllama() error { return nil } +func (p ModelRunProfile) validateMLX() error { + if strings.TrimSpace(p.MLXModel) == "" { + return fmt.Errorf("mlx model is required") + } + if p.ArtifactKind != "" && p.ArtifactKind != ArtifactMLXModelDir { + return fmt.Errorf("mlx artifact must be %s", ArtifactMLXModelDir) + } + if strings.TrimSpace(p.WeightsPath) != "" && path.Clean(strings.TrimSpace(p.WeightsPath)) != "." { + return fmt.Errorf("--mlx-model and --weights are mutually exclusive") + } + if p.BindHost != "" && p.BindHost != "127.0.0.1" { + return fmt.Errorf("bind host must be 127.0.0.1") + } + if p.Port < 1 || p.Port > 65535 { + return fmt.Errorf("port must be between 1 and 65535") + } + if p.OllamaModel != "" || p.OllamaNumCtx != nil || p.OllamaKeepAlive != "" || p.OllamaNumGPU != nil { + return fmt.Errorf("only mlx launch fields are supported") + } + if p.NGPULayers != nil || p.NGPULayersMode != "" || p.ContextTokens != nil || p.BatchSize != nil || p.UBatchSize != nil || p.Threads != nil || p.DeviceIndex != nil || p.IndexSource != "" { + return fmt.Errorf("only mlx launch fields are supported") + } + if p.PrefillStepSize != nil && *p.PrefillStepSize < 1 { + return fmt.Errorf("prefill-step-size must be >= 1") + } + if p.PromptCacheBytes != nil && *p.PromptCacheBytes < 1 { + return fmt.Errorf("prompt-cache-bytes must be >= 1") + } + if p.KVBits != nil && *p.KVBits < 1 { + return fmt.Errorf("kv-bits must be >= 1") + } + return nil +} + func (p ModelRunProfile) validateDevicePin() error { if p.DeviceIndex == nil && p.IndexSource == "" { return nil diff --git a/internal/models/run_profile_test.go b/internal/models/run_profile_test.go index b52fc5dc..289fc4e8 100644 --- a/internal/models/run_profile_test.go +++ b/internal/models/run_profile_test.go @@ -207,3 +207,32 @@ func TestValidateOllamaAllowsModelNameAndRefusesForeignFields(t *testing.T) { t.Fatalf("llama with ollama field err=%v", err) } } + +func TestValidateMLXRequiresDirectoryAndRefusesForeignFields(t *testing.T) { + profile := ModelRunProfile{ + Schema: ModelRunSchema, + Engine: EngineMLX, + ArtifactKind: ArtifactMLXModelDir, + MLXModel: "/mnt/models/qwen", + BindHost: "127.0.0.1", + Port: 8080, + } + if err := profile.Validate(); err != nil { + t.Fatal(err) + } + profile.WeightsPath = "/mnt/models/a.gguf" + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "mutually exclusive") { + t.Fatalf("weights err=%v", err) + } + profile.WeightsPath = "" + profile.OllamaModel = "mistral" + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "only mlx") { + t.Fatalf("ollama err=%v", err) + } + profile.OllamaModel = "" + bits := 0 + profile.KVBits = &bits + if err := profile.Validate(); err == nil || !strings.Contains(err.Error(), "kv-bits") { + t.Fatalf("kv err=%v", err) + } +} From acf97dd6e2d840f8984a87854c9f855bfaaab379 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 17:58:11 -0400 Subject: [PATCH 6/9] feat(reservation): keep device holds apart from entry VRAM A device hold stores one GPU index and that device's MiB. Index 0 round-trips. Load drops an expired hold without adding it to the VRAM sums, and entry writes keep the hold slice. --- cmd/axis/reservations.go | 74 +++- cmd/axis/reservations_test.go | 137 ++++++++ internal/reservation/device_hold.go | 161 +++++++++ internal/reservation/device_hold_test.go | 412 +++++++++++++++++++++++ internal/reservation/ledger.go | 24 +- internal/reservation/persist.go | 35 +- 6 files changed, 814 insertions(+), 29 deletions(-) create mode 100644 internal/reservation/device_hold.go create mode 100644 internal/reservation/device_hold_test.go diff --git a/cmd/axis/reservations.go b/cmd/axis/reservations.go index 5ebb8c4b..54726f40 100644 --- a/cmd/axis/reservations.go +++ b/cmd/axis/reservations.go @@ -223,23 +223,47 @@ func reservationsListCmd() *cobra.Command { } return nil default: - if len(entries) == 0 { + holds := ledger.DeviceHolds() + if len(entries) == 0 && len(holds) == 0 { _, err := fmt.Fprintln(cmd.OutOrStdout(), "No active reservations") return err } - tbl := ui.NewTable("ID", "NODE", "RAM MB", "OWNER", "CREATED AT", "LAST HEARTBEAT") - for _, e := range entries { - tbl.AddRow( - truncateID(e.ID, 20), - e.Node, - fmt.Sprintf("%d", e.RAMMB), - truncateID(e.OwnerSurface, 15), - e.CreatedAt.Format(time.RFC3339), - e.LastHeartbeat.Format(time.RFC3339), - ) - } var b strings.Builder - tbl.Render(&b) + if len(entries) == 0 { + fmt.Fprintln(&b, "No active reservations") + } else { + tbl := ui.NewTable("ID", "NODE", "RAM MB", "OWNER", "CREATED AT", "LAST HEARTBEAT") + for _, e := range entries { + tbl.AddRow( + truncateID(e.ID, 20), + e.Node, + fmt.Sprintf("%d", e.RAMMB), + truncateID(e.OwnerSurface, 15), + e.CreatedAt.Format(time.RFC3339), + e.LastHeartbeat.Format(time.RFC3339), + ) + } + tbl.Render(&b) + } + if len(holds) > 0 { + fmt.Fprintln(&b, "DEVICE HOLDS") + ht := ui.NewTable("ID", "NODE", "GPU INDEX", "MIB", "OWNER", "EXPIRES AT") + for _, h := range holds { + gpu := "" + if h.GPUIndex != nil { + gpu = fmt.Sprintf("%d", *h.GPUIndex) + } + ht.AddRow( + truncateID(h.ID, 20), + h.Node, + gpu, + fmt.Sprintf("%d", h.MiB), + truncateID(h.Owner, 15), + h.ExpiresAt.Format(time.RFC3339), + ) + } + ht.Render(&b) + } _, err := io.WriteString(cmd.OutOrStdout(), b.String()) return err } @@ -274,6 +298,11 @@ func reservationsInspectCmd() *cobra.Command { } if found == nil { + for _, hold := range ledger.DeviceHolds() { + if hold.ID == id { + return writeDeviceHold(cmd.OutOrStdout(), format, hold) + } + } return ExitCodeError{Code: ExitErrGeneric, Message: fmt.Sprintf("reservation %q not found", id)} } @@ -310,6 +339,25 @@ func reservationsInspectCmd() *cobra.Command { return cmd } +func writeDeviceHold(w io.Writer, format string, hold reservation.DeviceHold) error { + switch format { + case "json": + return json.NewEncoder(w).Encode(hold) + default: + var b strings.Builder + fmt.Fprintf(&b, "ID: %s\n", hold.ID) + fmt.Fprintf(&b, "Node: %s\n", hold.Node) + if hold.GPUIndex != nil { + fmt.Fprintf(&b, "GPU index: %d\n", *hold.GPUIndex) + } + fmt.Fprintf(&b, "MiB: %d\n", hold.MiB) + fmt.Fprintf(&b, "Owner: %s\n", hold.Owner) + fmt.Fprintf(&b, "Expires At: %s\n", hold.ExpiresAt.Format(time.RFC3339)) + _, err := io.WriteString(w, b.String()) + return err + } +} + func reservationsReleaseCmd() *cobra.Command { var force bool var format string diff --git a/cmd/axis/reservations_test.go b/cmd/axis/reservations_test.go index e3ccf9e3..c6b161f0 100644 --- a/cmd/axis/reservations_test.go +++ b/cmd/axis/reservations_test.go @@ -505,6 +505,143 @@ func TestReservationsListText(t *testing.T) { _ = stderr } +func TestReservationsListAndInspectDeviceHold(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + ledger := reservation.NewLedger(reservation.DefaultLimits(), nil) + ledger.SetNodeCapacity("node-a", 16384) + if _, err := ledger.Reserve(reservation.Entry{ + ID: "exec-1", + Node: "node-a", + OwnerSurface: "guarded-exec", + RAMMB: 1024, + VRAMMB: 2048, + }); err != nil { + t.Fatal(err) + } + gpu := 0 + if _, err := ledger.HoldDevice(reservation.DeviceHold{ + ID: "gpu-hold", + Node: "node-a", + GPUIndex: &gpu, + MiB: 1536, + Owner: "operator", + ExpiresAt: time.Now().Add(time.Hour).UTC(), + }); err != nil { + t.Fatal(err) + } + + list := reservationsListCmd() + stdout, stderr, err := captureProcessOutput(t, func() error { + list.SetArgs(nil) + return list.Execute() + }) + if err != nil { + t.Fatalf("list text: %v\nstdout: %s\nstderr: %s", err, stdout, stderr) + } + if !strings.Contains(stdout, "DEVICE HOLDS") { + t.Fatalf("missing device hold table:\n%s", stdout) + } + var holdCells []string + for _, line := range strings.Split(stdout, "\n") { + if strings.Contains(line, "gpu-hold") { + for _, cell := range strings.Split(line, "│") { + cell = strings.TrimSpace(cell) + if cell != "" { + holdCells = append(holdCells, cell) + } + } + } + } + if len(holdCells) < 4 || holdCells[2] != "0" || holdCells[3] != "1536" { + t.Fatalf("hold row = %#v, want GPU index 0 and MiB 1536\n%s", holdCells, stdout) + } + + listJSON := reservationsListCmd() + jsonOut, _, err := captureProcessOutput(t, func() error { + listJSON.SetArgs([]string{"--format", "json"}) + return listJSON.Execute() + }) + if err != nil { + t.Fatal(err) + } + var entries []reservation.Entry + if err := json.Unmarshal([]byte(jsonOut), &entries); err != nil { + t.Fatalf("json list changed shape: %v\n%s", err, jsonOut) + } + if len(entries) != 1 || entries[0].ID != "exec-1" || strings.Contains(jsonOut, "gpu-hold") { + t.Fatalf("json list = %s", jsonOut) + } + + inspectEntry := reservationsInspectCmd() + entryOut, _, err := captureProcessOutput(t, func() error { + inspectEntry.SetArgs([]string{"exec-1"}) + return inspectEntry.Execute() + }) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(entryOut, "VRAM MB:") || strings.Contains(entryOut, "GPU index") { + t.Fatalf("entry inspect relabeled VRAM:\n%s", entryOut) + } + + inspectHold := reservationsInspectCmd() + holdOut, _, err := captureProcessOutput(t, func() error { + inspectHold.SetArgs([]string{"gpu-hold"}) + return inspectHold.Execute() + }) + if err != nil { + t.Fatalf("inspect hold: %v\n%s", err, holdOut) + } + if !strings.Contains(holdOut, "GPU index:") || !strings.Contains(holdOut, "1536") || strings.Contains(holdOut, "VRAM MB") { + t.Fatalf("hold inspect =\n%s", holdOut) + } + + releaseHold := reservationsReleaseCmd() + _, _, err = captureProcessOutput(t, func() error { + releaseHold.SetArgs([]string{"gpu-hold"}) + return releaseHold.Execute() + }) + if err == nil || !strings.Contains(err.Error(), `reservation "gpu-hold" not found`) { + t.Fatalf("release hold error = %v", err) + } + reloaded := reservation.NewLedger(reservation.DefaultLimits(), nil) + if err := reloaded.Load(); err != nil { + t.Fatal(err) + } + holds := reloaded.DeviceHolds() + if len(holds) != 1 || holds[0].ID != "gpu-hold" || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("release of hold id changed holds: %+v", holds) + } + if len(reloaded.Entries()) != 1 { + t.Fatalf("hold release removed the entry: %+v", reloaded.Entries()) + } +} + +func TestReservationsListHoldOnlyPrintsSecondTable(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + ledger := reservation.NewLedger(reservation.DefaultLimits(), nil) + gpu := 0 + if _, err := ledger.HoldDevice(reservation.DeviceHold{ + ID: "gpu-hold", Node: "node-a", GPUIndex: &gpu, MiB: 64, Owner: "operator", + ExpiresAt: time.Now().Add(time.Hour).UTC(), + }); err != nil { + t.Fatal(err) + } + cmd := reservationsListCmd() + stdout, stderr, err := captureProcessOutput(t, func() error { + cmd.SetArgs(nil) + return cmd.Execute() + }) + if err != nil { + t.Fatalf("list hold-only: %v\nstdout: %s\nstderr: %s", err, stdout, stderr) + } + if !strings.Contains(stdout, "No active reservations") || !strings.Contains(stdout, "DEVICE HOLDS") || !strings.Contains(stdout, "gpu-hold") { + t.Fatalf("hold-only list =\n%s", stdout) + } +} + func TestFormatDuration(t *testing.T) { tests := []struct { d time.Duration diff --git a/internal/reservation/device_hold.go b/internal/reservation/device_hold.go new file mode 100644 index 00000000..1f3c1d55 --- /dev/null +++ b/internal/reservation/device_hold.go @@ -0,0 +1,161 @@ +package reservation + +import ( + "context" + "fmt" + "sort" + "strings" + "time" +) + +// DeviceHold is one device's reserved MiB. It is not Entry.VRAMMB and it is +// not added to node or cluster VRAM sums. GPUIndex has no omitempty: a nil +// index is a refusal, and 0 is a legal observed index that must round-trip. +// A past expiry is stored; Load drops it. Expiry does not stop a process. +type DeviceHold struct { + ID string `json:"id"` + Node string `json:"node"` + GPUIndex *int `json:"gpu_index"` + MiB int64 `json:"mib"` + Owner string `json:"owner"` + ExpiresAt time.Time `json:"expires_at"` +} + +func cloneDeviceHold(hold *DeviceHold) *DeviceHold { + if hold == nil { + return nil + } + cp := *hold + if hold.GPUIndex != nil { + idx := *hold.GPUIndex + cp.GPUIndex = &idx + } + return &cp +} + +// HoldDevice records one device hold. It does not look up a GPU fact and it +// does not spend the hold. The caller owns neither the stored index pointer +// nor the returned one. +func (l *Ledger) HoldDevice(hold DeviceHold) (*DeviceHold, error) { + l.fileMu.Lock() + defer l.fileMu.Unlock() + + wasLocked := l.lockFile != nil + if !wasLocked { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := l.lockFileLocked(ctx); err != nil { + return nil, err + } + defer l.unlockFileLocked() + } + + var snap []*Entry + var holds []*DeviceHold + var stored *DeviceHold + err := func() error { + l.mu.Lock() + defer l.mu.Unlock() + + hold.ID = strings.TrimSpace(hold.ID) + hold.Node = strings.TrimSpace(hold.Node) + hold.Owner = strings.TrimSpace(hold.Owner) + if hold.ID == "" { + return fmt.Errorf("reservation: device hold ID required") + } + if hold.Node == "" { + return fmt.Errorf("reservation: device hold node required") + } + if hold.GPUIndex == nil { + return fmt.Errorf("reservation: device hold GPU index required") + } + if hold.MiB <= 0 { + return fmt.Errorf("reservation: device hold MiB must be > 0") + } + if hold.Owner == "" { + return fmt.Errorf("reservation: device hold owner required") + } + if hold.ExpiresAt.IsZero() { + return fmt.Errorf("reservation: device hold expiry required") + } + if l.deviceHolds == nil { + l.deviceHolds = make(map[string]*DeviceHold) + } + if _, exists := l.entries[hold.ID]; exists { + return fmt.Errorf("reservation: duplicate ID %q", hold.ID) + } + if _, exists := l.deviceHolds[hold.ID]; exists { + return fmt.Errorf("reservation: duplicate ID %q", hold.ID) + } + stored = cloneDeviceHold(&hold) + l.deviceHolds[stored.ID] = stored + snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() + return nil + }() + if err != nil { + return nil, err + } + if err := l.writeSnapshot(snap, holds); err != nil { + l.logger.Error("failed to persist device hold", "error", err) + } + return cloneDeviceHold(stored), nil +} + +// DeviceHolds returns a copy of every device hold, sorted by ID. +func (l *Ledger) DeviceHolds() []DeviceHold { + l.mu.RLock() + defer l.mu.RUnlock() + snap := l.snapshotDeviceHoldsLocked() + out := make([]DeviceHold, 0, len(snap)) + for _, hold := range snap { + out = append(out, *hold) + } + return out +} + +// snapshotDeviceHoldsLocked returns independent copies. The caller must hold l.mu. +func (l *Ledger) snapshotDeviceHoldsLocked() []*DeviceHold { + out := make([]*DeviceHold, 0, len(l.deviceHolds)) + for _, hold := range l.deviceHolds { + if cp := cloneDeviceHold(hold); cp != nil { + out = append(out, cp) + } + } + sort.Slice(out, func(i, j int) bool { return out[i].ID < out[j].ID }) + return out +} + +func (l *Ledger) replaceDeviceHolds(holds []*DeviceHold) { + l.mu.Lock() + defer l.mu.Unlock() + l.replaceDeviceHoldsLocked(holds) +} + +func (l *Ledger) replaceDeviceHoldsLocked(holds []*DeviceHold) { + l.deviceHolds = make(map[string]*DeviceHold, len(holds)) + for _, hold := range holds { + cp := cloneDeviceHold(hold) + if cp == nil || cp.ID == "" { + continue + } + l.deviceHolds[cp.ID] = cp + } +} + +// dropExpiredDeviceHoldsLocked removes holds whose expiry is before now. +// A zero expiry is kept. Equal-to-now is kept. The caller must hold l.mu. +func (l *Ledger) dropExpiredDeviceHoldsLocked() int { + now := l.now().UTC() + dropped := 0 + for id, hold := range l.deviceHolds { + if hold == nil || hold.ExpiresAt.IsZero() { + continue + } + if now.After(hold.ExpiresAt) { + delete(l.deviceHolds, id) + dropped++ + } + } + return dropped +} diff --git a/internal/reservation/device_hold_test.go b/internal/reservation/device_hold_test.go new file mode 100644 index 00000000..098ebb7c --- /dev/null +++ b/internal/reservation/device_hold_test.go @@ -0,0 +1,412 @@ +package reservation + +import ( + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func intPtr(v int) *int { return &v } + +func futureHold(id string, gpu int) DeviceHold { + return DeviceHold{ + ID: id, + Node: "node-a", + GPUIndex: intPtr(gpu), + MiB: 512, + Owner: "operator", + ExpiresAt: time.Now().Add(time.Hour), + } +} + +func writeLedgerFixture(t *testing.T, df diskFormat) []byte { + t.Helper() + data, err := json.MarshalIndent(df, "", " ") + if err != nil { + t.Fatalf("marshal ledger fixture: %v", err) + } + if err := os.MkdirAll(filepath.Dir(Path()), 0o700); err != nil { + t.Fatalf("create ledger directory: %v", err) + } + if err := os.WriteFile(Path(), data, 0o600); err != nil { + t.Fatalf("write ledger fixture: %v", err) + } + return data +} + +func TestHoldDeviceRefusesIncomplete(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + expiry := time.Now().Add(time.Hour) + base := DeviceHold{ID: "gpu-0", Node: "node-a", GPUIndex: intPtr(1), MiB: 256, Owner: "operator", ExpiresAt: expiry} + cases := []struct { + name string + edit func(*DeviceHold) + want string + }{ + {name: "nil index", edit: func(h *DeviceHold) { h.GPUIndex = nil }, want: "GPU index"}, + {name: "zero mib", edit: func(h *DeviceHold) { h.MiB = 0 }, want: "MiB"}, + {name: "negative mib", edit: func(h *DeviceHold) { h.MiB = -5 }, want: "MiB"}, + {name: "empty id", edit: func(h *DeviceHold) { h.ID = "" }, want: "ID"}, + {name: "blank id", edit: func(h *DeviceHold) { h.ID = " " }, want: "ID"}, + {name: "empty node", edit: func(h *DeviceHold) { h.Node = "" }, want: "node"}, + {name: "blank node", edit: func(h *DeviceHold) { h.Node = " " }, want: "node"}, + {name: "empty owner", edit: func(h *DeviceHold) { h.Owner = "" }, want: "owner"}, + {name: "blank owner", edit: func(h *DeviceHold) { h.Owner = "\t" }, want: "owner"}, + {name: "zero expiry", edit: func(h *DeviceHold) { h.ExpiresAt = time.Time{} }, want: "expiry"}, + } + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + hold := base + tt.edit(&hold) + if _, err := l.HoldDevice(hold); err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("HoldDevice error = %v, want substring %q", err, tt.want) + } + }) + } + if got := len(l.DeviceHolds()); got != 0 { + t.Fatalf("refusals stored %d holds", got) + } +} + +func TestHoldDeviceIndexZeroRoundTrips(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + gpu := 0 + hold := futureHold("gpu-0", 1) + hold.GPUIndex = &gpu + got, err := l.HoldDevice(hold) + if err != nil { + t.Fatal(err) + } + gpu = 4 + *got.GPUIndex = 9 + holds := l.DeviceHolds() + if len(holds) != 1 || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("in-memory hold = %+v, want GPU index 0", holds) + } + + raw, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(raw), `"gpu_index": 0`) { + t.Fatalf("ledger JSON omitted GPU index 0:\n%s", raw) + } + + reloaded := NewLedger(DefaultLimits(), nil) + if err := reloaded.Load(); err != nil { + t.Fatal(err) + } + again := reloaded.DeviceHolds() + if len(again) != 1 || again[0].GPUIndex == nil || *again[0].GPUIndex != 0 { + t.Fatalf("reloaded hold = %+v, want GPU index 0", again) + } +} + +func TestHoldDeviceRejectsDuplicateEntryOrHold(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + l.SetNodeCapacity("node-a", 16384) + if _, err := l.Reserve(Entry{ID: "exec-1", Node: "node-a", RAMMB: 1024}); err != nil { + t.Fatal(err) + } + if _, err := l.HoldDevice(futureHold("exec-1", 0)); err == nil || !strings.Contains(err.Error(), "duplicate") { + t.Fatalf("hold matching an entry: %v", err) + } + if _, err := l.HoldDevice(futureHold("gpu-0", 0)); err != nil { + t.Fatal(err) + } + if _, err := l.HoldDevice(futureHold("gpu-0", 1)); err == nil || !strings.Contains(err.Error(), "duplicate") { + t.Fatalf("duplicate hold: %v", err) + } + if _, err := l.Reserve(Entry{ID: "gpu-0", Node: "node-a", RAMMB: 128}); err == nil || !strings.Contains(err.Error(), "duplicate") { + t.Fatalf("entry matching a hold: %v", err) + } +} + +func TestLoadOldEntriesFileHasNoHolds(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + now := time.Now().UTC() + writeLedgerFixture(t, diskFormat{Entries: []*Entry{{ + ID: "exec-1", + Node: "node-a", + RAMMB: 128, + CreatedAt: now, + LastHeartbeat: now, + }}}) + if err := l.Load(); err != nil { + t.Fatal(err) + } + if len(l.DeviceHolds()) != 0 { + t.Fatalf("old entries file invented holds: %+v", l.DeviceHolds()) + } + entries := l.Entries() + if len(entries) != 1 || entries[0].ID != "exec-1" { + t.Fatalf("entries = %+v", entries) + } +} + +func TestLoadDropsExpiredHoldWithoutReclaimingEntries(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + now := time.Now().UTC() + writeLedgerFixture(t, diskFormat{ + Entries: []*Entry{{ + ID: "exec-fresh", + Node: "node-a", + RAMMB: 128, + CreatedAt: now, + LastHeartbeat: now, + }}, + DeviceHolds: []*DeviceHold{ + {ID: "hold-old", Node: "node-a", GPUIndex: intPtr(1), MiB: 100, Owner: "op", ExpiresAt: now.Add(-time.Minute)}, + {ID: "hold-new", Node: "node-a", GPUIndex: intPtr(2), MiB: 200, Owner: "op", ExpiresAt: now.Add(time.Hour)}, + }, + }) + if err := l.Load(); err != nil { + t.Fatal(err) + } + if entries := l.Entries(); len(entries) != 1 || entries[0].ID != "exec-fresh" { + t.Fatalf("entries after hold expiry = %+v", entries) + } + holds := l.DeviceHolds() + if len(holds) != 1 || holds[0].ID != "hold-new" { + t.Fatalf("holds after expiry = %+v", holds) + } + raw, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(raw), "hold-old") { + t.Fatalf("expired hold was left on disk:\n%s", raw) + } + if !strings.Contains(string(raw), "hold-new") || !strings.Contains(string(raw), "exec-fresh") { + t.Fatalf("rewrite dropped a live record:\n%s", raw) + } + + again := NewLedger(DefaultLimits(), nil) + if err := again.Load(); err != nil { + t.Fatal(err) + } + if len(again.DeviceHolds()) != 1 || again.DeviceHolds()[0].ID != "hold-new" { + t.Fatalf("second load resurrected or dropped holds: %+v", again.DeviceHolds()) + } +} + +func TestLoadKeepsHoldExpiringAtNow(t *testing.T) { + now := time.Date(2026, 10, 3, 15, 0, 0, 0, time.UTC) + l := setupTestLedger(t, DefaultLimits()) + l.now = func() time.Time { return now } + writeLedgerFixture(t, diskFormat{ + Entries: []*Entry{{ + ID: "exec-fresh", Node: "node-a", RAMMB: 64, + CreatedAt: now, LastHeartbeat: now, + }}, + DeviceHolds: []*DeviceHold{{ + ID: "hold-eq", Node: "node-a", GPUIndex: intPtr(0), MiB: 32, Owner: "op", ExpiresAt: now, + }}, + }) + before, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if err := l.Load(); err != nil { + t.Fatal(err) + } + holds := l.DeviceHolds() + if len(holds) != 1 || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("equal expiry dropped the hold: %+v", holds) + } + after, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if string(before) != string(after) { + t.Fatalf("equal-expiry load rewrote the ledger\n got: %s\nwant: %s", after, before) + } +} + +func TestLoadEntryReclaimKeepsFutureHold(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + now := time.Now().UTC() + stale := now.Add(-10 * time.Minute) + writeLedgerFixture(t, diskFormat{ + Entries: []*Entry{{ + ID: "exec-stale", Node: "node-a", RAMMB: 256, + CreatedAt: stale, LastHeartbeat: stale, + }}, + DeviceHolds: []*DeviceHold{{ + ID: "hold-live", Node: "node-a", GPUIndex: intPtr(0), MiB: 768, Owner: "op", + ExpiresAt: now.Add(time.Hour), + }}, + }) + if err := l.Load(); err != nil { + t.Fatal(err) + } + if len(l.Entries()) != 0 { + t.Fatalf("stale entry survived: %+v", l.Entries()) + } + holds := l.DeviceHolds() + if len(holds) != 1 || holds[0].ID != "hold-live" || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("entry reclaim deleted the hold: %+v", holds) + } + raw, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + text := string(raw) + if strings.Contains(text, "exec-stale") || !strings.Contains(text, "hold-live") || !strings.Contains(text, `"gpu_index": 0`) { + t.Fatalf("reclaim snapshot = %s", text) + } +} + +func TestLoadKeepsZeroExpiryHold(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + now := time.Now().UTC() + before := writeLedgerFixture(t, diskFormat{ + Entries: []*Entry{{ + ID: "exec-fresh", Node: "node-a", RAMMB: 64, + CreatedAt: now, LastHeartbeat: now, + }}, + DeviceHolds: []*DeviceHold{{ + ID: "hold-open", Node: "node-a", GPUIndex: intPtr(3), MiB: 16, Owner: "op", + }}, + }) + if err := l.Load(); err != nil { + t.Fatal(err) + } + if len(l.DeviceHolds()) != 1 || l.DeviceHolds()[0].ID != "hold-open" { + t.Fatalf("zero expiry was dropped: %+v", l.DeviceHolds()) + } + after, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if string(before) != string(after) { + t.Fatalf("zero-expiry load rewrote the ledger\n got: %s\nwant: %s", after, before) + } +} + +func TestLoadReadOnlyKeepsExpiredHoldAndBytes(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + now := time.Now().UTC() + before := writeLedgerFixture(t, diskFormat{DeviceHolds: []*DeviceHold{{ + ID: "hold-old", Node: "node-a", GPUIndex: intPtr(1), MiB: 100, Owner: "op", + ExpiresAt: now.Add(-time.Hour), + }}}) + if err := l.LoadReadOnly(); err != nil { + t.Fatal(err) + } + holds := l.DeviceHolds() + if len(holds) != 1 || holds[0].ID != "hold-old" { + t.Fatalf("read-only load dropped the expired hold: %+v", holds) + } + after, err := os.ReadFile(Path()) + if err != nil { + t.Fatal(err) + } + if string(before) != string(after) { + t.Fatalf("read-only load rewrote ledger.json\n got: %s\nwant: %s", after, before) + } +} + +func TestLoadReadOnlyMissingFileClearsHolds(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + if _, err := l.HoldDevice(futureHold("gpu-0", 0)); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(Path()); err != nil { + t.Fatalf("HoldDevice did not persist: %v", err) + } + if err := os.Remove(Path()); err != nil { + t.Fatal(err) + } + if err := l.LoadReadOnly(); err != nil { + t.Fatal(err) + } + if len(l.DeviceHolds()) != 0 { + t.Fatalf("missing file left holds in memory: %+v", l.DeviceHolds()) + } +} + +func TestReleaseAndHeartbeatKeepDeviceHold(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + l.SetNodeCapacity("node-a", 16384) + if _, err := l.Reserve(Entry{ID: "exec-1", Node: "node-a", RAMMB: 1024, VRAMMB: 2048}); err != nil { + t.Fatal(err) + } + if _, err := l.HoldDevice(futureHold("gpu-0", 0)); err != nil { + t.Fatal(err) + } + if err := l.Heartbeat("exec-1"); err != nil { + t.Fatal(err) + } + if holds := l.DeviceHolds(); len(holds) != 1 || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("heartbeat dropped the hold: %+v", holds) + } + if err := l.Release("exec-1"); err != nil { + t.Fatal(err) + } + if err := l.Release("gpu-0"); err == nil { + t.Fatal("release of a hold id deleted or accepted the hold") + } + if holds := l.DeviceHolds(); len(holds) != 1 { + t.Fatalf("release removed the hold from memory: %+v", holds) + } + + reloaded := NewLedger(DefaultLimits(), nil) + if err := reloaded.Load(); err != nil { + t.Fatal(err) + } + if len(reloaded.Entries()) != 0 { + t.Fatalf("released entry returned: %+v", reloaded.Entries()) + } + holds := reloaded.DeviceHolds() + if len(holds) != 1 || holds[0].ID != "gpu-0" || holds[0].GPUIndex == nil || *holds[0].GPUIndex != 0 { + t.Fatalf("reload lost the hold: %+v", holds) + } +} + +func TestSummaryIgnoresHoldMiB(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + l.SetNodeCapacity("node-a", 16384) + if _, err := l.Reserve(Entry{ID: "exec-1", Node: "node-a", RAMMB: 1024, VRAMMB: 2048}); err != nil { + t.Fatal(err) + } + hold := futureHold("gpu-0", 0) + hold.MiB = 9000 + if _, err := l.HoldDevice(hold); err != nil { + t.Fatal(err) + } + summary := l.Summary() + if summary.TotalVRAMMB != 2048 { + t.Fatalf("TotalVRAMMB = %d, want entry VRAM 2048", summary.TotalVRAMMB) + } + var reserved int64 + for _, node := range summary.Nodes { + if node.Node == "node-a" { + reserved = node.ReservedVRAMMB + } + } + if reserved != 2048 { + t.Fatalf("ReservedVRAMMB = %d, want 2048", reserved) + } +} + +func TestPastExpiryIsStoredUntilLoad(t *testing.T) { + l := setupTestLedger(t, DefaultLimits()) + hold := futureHold("hold-old", 1) + hold.ExpiresAt = time.Now().Add(-time.Second) + if _, err := l.HoldDevice(hold); err != nil { + t.Fatal(err) + } + if len(l.DeviceHolds()) != 1 { + t.Fatal("past expiry was refused") + } + if err := l.Load(); err != nil { + t.Fatal(err) + } + if len(l.DeviceHolds()) != 0 { + t.Fatalf("Load kept the expired hold: %+v", l.DeviceHolds()) + } +} diff --git a/internal/reservation/ledger.go b/internal/reservation/ledger.go index 502bfd9c..eac44252 100644 --- a/internal/reservation/ledger.go +++ b/internal/reservation/ledger.go @@ -153,6 +153,10 @@ type Ledger struct { // When zero or missing, falls back to limits.SystemReserveMB. nodeReserve map[string]int64 + // deviceHolds maps hold ID → one device's MiB. It is not part of entries + // and it is not added to VRAM totals. + deviceHolds map[string]*DeviceHold + // fileMu serializes disk persistence (file-lock bookkeeping + marshal + // atomic write) across goroutines. It is always acquired before mu when // both are needed. Holding it during I/O leaves mu free so concurrent @@ -178,6 +182,7 @@ func NewLedger(limits Limits, logger *slog.Logger) *Ledger { logger: logger.With("component", "reservation-ledger"), nodeRAM: make(map[string]int64), nodeReserve: make(map[string]int64), + deviceHolds: make(map[string]*DeviceHold), now: time.Now, } } @@ -227,6 +232,7 @@ func (l *Ledger) Reserve(req Entry) (*Entry, error) { } var snap []*Entry + var holds []*DeviceHold result, err := func() (*Entry, error) { l.mu.Lock() defer l.mu.Unlock() @@ -250,6 +256,9 @@ func (l *Ledger) Reserve(req Entry) (*Entry, error) { if _, exists := l.entries[req.ID]; exists { return nil, fmt.Errorf("reservation: duplicate ID %q", req.ID) } + if _, exists := l.deviceHolds[req.ID]; exists { + return nil, fmt.Errorf("reservation: duplicate ID %q", req.ID) + } // Check per-node cap nodeCount := 0 @@ -296,6 +305,7 @@ func (l *Ledger) Reserve(req Entry) (*Entry, error) { l.entries[req.ID] = &req l.totalReserved += req.RAMMB snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() return &req, nil }() if err != nil { @@ -317,7 +327,7 @@ func (l *Ledger) Reserve(req Entry) (*Entry, error) { "owner": req.OwnerSurface, }) - if err := l.writeSnapshot(snap); err != nil { + if err := l.writeSnapshot(snap, holds); err != nil { l.logger.Error("failed to persist ledger", "error", err) } return result, nil @@ -339,6 +349,7 @@ func (l *Ledger) Release(id string) error { } var snap []*Entry + var holds []*DeviceHold var released Entry err := func() error { l.mu.Lock() @@ -351,6 +362,7 @@ func (l *Ledger) Release(id string) error { released = *e delete(l.entries, id) snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() return nil }() if err != nil { @@ -366,7 +378,7 @@ func (l *Ledger) Release(id string) error { "ram_mb": released.RAMMB, }) - return l.writeSnapshot(snap) + return l.writeSnapshot(snap, holds) } // Heartbeat updates the liveness timestamp for a reservation. @@ -385,6 +397,7 @@ func (l *Ledger) Heartbeat(id string) error { } var snap []*Entry + var holds []*DeviceHold err := func() error { l.mu.Lock() defer l.mu.Unlock() @@ -394,12 +407,13 @@ func (l *Ledger) Heartbeat(id string) error { } e.LastHeartbeat = l.now() snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() return nil }() if err != nil { return err } - return l.writeSnapshot(snap) + return l.writeSnapshot(snap, holds) } // Reclaim removes all stale and expired reservations. Returns count reclaimed. @@ -422,13 +436,15 @@ func (l *Ledger) Reclaim() int { l.mu.Lock() reclaimed, receipts := l.reclaimInMemoryLocked() var snap []*Entry + var holds []*DeviceHold if reclaimed > 0 { snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() } l.mu.Unlock() if reclaimed > 0 { - if err := l.writeSnapshot(snap); err != nil { + if err := l.writeSnapshot(snap, holds); err != nil { l.logger.Error("failed to persist ledger during reclaim", "error", err) } else { repairs.EmitAll(l.logger, receipts) diff --git a/internal/reservation/persist.go b/internal/reservation/persist.go index 91a89543..7453ff30 100644 --- a/internal/reservation/persist.go +++ b/internal/reservation/persist.go @@ -20,7 +20,8 @@ func Path() string { // diskFormat represents the serialized ledger. type diskFormat struct { - Entries []*Entry `json:"entries"` + Entries []*Entry `json:"entries"` + DeviceHolds []*DeviceHold `json:"device_holds,omitempty"` } // LockFile acquires an exclusive lock on the ledger lockfile with a 500ms timeout. @@ -123,19 +124,25 @@ func (l *Ledger) Load() error { l.entries[e.ID] = e l.totalReserved += e.RAMMB } + l.replaceDeviceHoldsLocked(df.DeviceHolds) + droppedHolds := l.dropExpiredDeviceHoldsLocked() // Startup reconciliation pass (in-memory only; persist after unlocking mu). reclaimed, receipts := l.reclaimInMemoryLocked() var snap []*Entry - if reclaimed > 0 { + var holds []*DeviceHold + if reclaimed > 0 || droppedHolds > 0 { snap = l.snapshotEntriesLocked() + holds = l.snapshotDeviceHoldsLocked() } l.mu.Unlock() - if reclaimed > 0 { - l.logger.Info("startup reconciliation complete", "reclaimed", reclaimed) - if err := l.writeSnapshot(snap); err != nil { + if reclaimed > 0 || droppedHolds > 0 { + if reclaimed > 0 { + l.logger.Info("startup reconciliation complete", "reclaimed", reclaimed) + } + if err := l.writeSnapshot(snap, holds); err != nil { l.logger.Error("failed to persist ledger during startup reconciliation", "error", err) - } else { + } else if reclaimed > 0 { repairs.EmitAll(l.logger, receipts) } } @@ -165,6 +172,7 @@ func (l *Ledger) LoadReadOnly() error { if err != nil { if os.IsNotExist(err) { l.replaceEntries(nil) + l.replaceDeviceHolds(nil) return nil } return err @@ -175,6 +183,7 @@ func (l *Ledger) LoadReadOnly() error { return fmt.Errorf("decode reservation ledger without recovery: %w", err) } l.replaceEntries(df.Entries) + l.replaceDeviceHolds(df.DeviceHolds) return nil } @@ -212,9 +221,10 @@ func (l *Ledger) Save() error { l.mu.RLock() snap := l.snapshotEntriesLocked() + holds := l.snapshotDeviceHoldsLocked() l.mu.RUnlock() - return l.writeSnapshot(snap) + return l.writeSnapshot(snap, holds) } // snapshotEntriesLocked returns independent copies of every entry. The caller @@ -229,16 +239,17 @@ func (l *Ledger) snapshotEntriesLocked() []*Entry { return out } -// writeSnapshot marshals entries and atomically writes them to disk. It must be -// called with the file lock held (fileMu) but WITHOUT l.mu, so the marshal and -// write never block in-memory readers. -func (l *Ledger) writeSnapshot(entries []*Entry) error { +// writeSnapshot marshals entries and device holds and atomically writes them +// to disk. It must be called with the file lock held (fileMu) but WITHOUT l.mu, +// so the marshal and write never block in-memory readers. Every caller passes +// the remaining holds so an entry write does not delete device_holds. +func (l *Ledger) writeSnapshot(entries []*Entry, holds []*DeviceHold) error { path := Path() if err := persist.EnsurePrivateDir(filepath.Dir(path)); err != nil { return err } - df := diskFormat{Entries: entries} + df := diskFormat{Entries: entries, DeviceHolds: holds} data, err := json.MarshalIndent(df, "", " ") if err != nil { return err From 1768ad2d846cae9da1976a801d523e54bd75d7d1 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 18:59:50 -0400 Subject: [PATCH 7/9] feat(model): record a probed llama-server peak and exclude it at plan time After the llama-server probe succeeds, sample that node's listener RSS and store one ExecutionObservation. A failed sample warns and leaves the process up. Model plan drops a candidate whose fresh peak exceeds allocatable RAM. --- cmd/axis/model.go | 202 +++++++- cmd/axis/model_observation_test.go | 499 +++++++++++++++++++ cmd/axis/model_test.go | 12 + internal/models/model_operation.go | 1 + internal/models/types.go | 6 + internal/placement/llama_server_peak.go | 20 + internal/placement/llama_server_peak_test.go | 93 ++++ internal/state/observations.go | 16 + internal/state/observations_test.go | 114 +++++ 9 files changed, 958 insertions(+), 5 deletions(-) create mode 100644 cmd/axis/model_observation_test.go create mode 100644 internal/placement/llama_server_peak.go create mode 100644 internal/placement/llama_server_peak_test.go diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 7023b13e..99566b68 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -7,6 +7,7 @@ import ( "os" "path" "path/filepath" + "sort" "strconv" "strings" "time" @@ -21,7 +22,9 @@ import ( "github.com/toasterbook88/axis/internal/modellife" "github.com/toasterbook88/axis/internal/modelplan" "github.com/toasterbook88/axis/internal/models" + "github.com/toasterbook88/axis/internal/placement" "github.com/toasterbook88/axis/internal/runtimectx" + "github.com/toasterbook88/axis/internal/state" "github.com/toasterbook88/axis/internal/transport" ) @@ -531,6 +534,11 @@ func runModelPlan(ctx context.Context, cmd *cobra.Command, specOrWeights string, if err != nil { return err } + st, err := state.Load() + if err != nil { + return err + } + applyLlamaServerPeakExclusions(&plan, snap, st) plan.SnapshotSource = source if cmd.Flags().Changed("port") && plan.Selected != nil { plan.Selected.PortSource = models.PortSourceExplicit @@ -749,6 +757,8 @@ func runModelStart(ctx context.Context, cmd *cobra.Command, nodeName, weights st receipt.Status = models.ModelOperationFailed receipt.Disposition = "failed" receipt.Error = startErr.Error() + } else if warning := recordLlamaServerObservation(ctx, cmd, nf, cfgNode, plan, startedAt); warning != "" { + receipt.Warnings = append(receipt.Warnings, warning) } if writeErr := writeModelStartReceipt(cmd, receipt, format); writeErr != nil { @@ -772,13 +782,21 @@ func writeModelStartReceipt(cmd *cobra.Command, receipt models.ModelOperationRec return printOutput(cmd.OutOrStdout(), receipt, format) } if receipt.Status == models.ModelOperationCompleted { - _, err := fmt.Fprintf(cmd.OutOrStdout(), "started %s on %s:%d volume %s operation %s\n", - receipt.Executable, receipt.Node, receipt.Port, receipt.Volume, receipt.ID) - if err != nil || receipt.DeviceNote == "" { + if _, err := fmt.Fprintf(cmd.OutOrStdout(), "started %s on %s:%d volume %s operation %s\n", + receipt.Executable, receipt.Node, receipt.Port, receipt.Volume, receipt.ID); err != nil { return err } - _, err = fmt.Fprintf(cmd.OutOrStdout(), "%s\n", receipt.DeviceNote) - return err + if receipt.DeviceNote != "" { + if _, err := fmt.Fprintf(cmd.OutOrStdout(), "%s\n", receipt.DeviceNote); err != nil { + return err + } + } + for _, warning := range receipt.Warnings { + if _, err := fmt.Fprintf(cmd.OutOrStdout(), "warning: %s\n", warning); err != nil { + return err + } + } + return nil } _, err := fmt.Fprintf(cmd.OutOrStdout(), "%s %s:%d: %s operation %s\n", receipt.Disposition, receipt.Node, receipt.Port, receipt.Error, receipt.ID) @@ -1652,6 +1670,52 @@ func shellMLXOwnerGuard(port int) string { ) } +const ( + llamaSampleRAMPrefix = "axis-llama-sample:ram " + llamaSampleVRAMPrefix = "axis-llama-sample:vram " + llamaSampleRAMFailed = "axis-llama-sample:ram-failed" +) + +func shellLlamaServerSample(port int, deviceIndex *int) string { + script := shellListenerLookup(port) + shellLlamaServerOwnerGuard(port) + + `if test -z "$_axis_pids"; then echo '` + llamaSampleRAMFailed + `' >&2; exit 1; fi; ` + + `_axis_peak=0; _axis_any=0; ` + + `for _axis_pid in $_axis_pids; do ` + + `_axis_rss=$(ps -p "$_axis_pid" -o rss=) || { echo '` + llamaSampleRAMFailed + `' >&2; exit 1; }; ` + + `_axis_mib=$((_axis_rss / 1024)); ` + + `if test "$_axis_mib" -gt "$_axis_peak"; then _axis_peak=$_axis_mib; fi; ` + + `_axis_any=1; ` + + `done; ` + + `if test "$_axis_any" != 1; then echo '` + llamaSampleRAMFailed + `' >&2; exit 1; fi; ` + + `echo "` + llamaSampleRAMPrefix + `$_axis_peak"; ` + if deviceIndex == nil { + return script + } + return script + fmt.Sprintf(`_axis_want=%d; `, *deviceIndex) + + `_axis_smi=$(nvidia-smi --query-gpu=index,memory.used --format=csv,noheader,nounits 2>/dev/null) || _axis_smi=""; ` + + `_axis_used=$(printf '%s\n' "$_axis_smi" | awk -F, -v want="$_axis_want" '{ idx=$1; gsub(/ /, "", idx); used=$2; gsub(/ /, "", used); if (idx == want && used ~ /^[0-9]+$/) { print used; exit } }'); ` + + `if test -n "$_axis_used"; then echo "` + llamaSampleVRAMPrefix + `$_axis_used"; fi; ` +} + +func parseLlamaServerSample(out string) (ram int64, ramOK bool, vram int64, vramOK bool) { + for _, line := range strings.Split(out, "\n") { + line = strings.TrimSpace(line) + switch { + case strings.HasPrefix(line, llamaSampleRAMPrefix): + n, err := strconv.ParseInt(strings.TrimSpace(strings.TrimPrefix(line, llamaSampleRAMPrefix)), 10, 64) + if err == nil && n >= 0 { + ram, ramOK = n, true + } + case strings.HasPrefix(line, llamaSampleVRAMPrefix): + n, err := strconv.ParseInt(strings.TrimSpace(strings.TrimPrefix(line, llamaSampleVRAMPrefix)), 10, 64) + if err == nil && n >= 0 { + vram, vramOK = n, true + } + } + } + return ram, ramOK, vram, vramOK +} + func shellLlamaServerOwnerGuard(port int) string { return fmt.Sprintf( "if ! command -v ps >/dev/null 2>&1; then echo 'axis model requires ps to verify process ownership' >&2; echo '"+modelStopMarker+"inspection_unavailable' >&2; exit 127; fi; "+ @@ -1691,6 +1755,134 @@ func runOnNode(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeC // llama-server start and stop keep calling runOnNodeCapturing directly. var runNodeScript = runOnNodeCapturing +// runLlamaServerSample is the post-probe RSS seam. Production calls +// runOnNodeCapturing, so a remote node is sampled over its SSH session and a +// local node is sampled with the local executor. Tests replace it. +var runLlamaServerSample = runOnNodeCapturing + +func recordLlamaServerObservation(ctx context.Context, cmd *cobra.Command, node models.NodeFacts, cfgNode *config.NodeConfig, plan modellife.StartPlan, startedAt time.Time) string { + modelName := llamaServerObservationModelName(plan.Weights) + obs := models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: plan.Node, + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + ModelName: modelName, + }, + ObservedAt: time.Now().UTC(), + SampleCount: 1, + LastSuccess: true, + WallTimeMS: observationWallMS(time.Since(startedAt)), + ModelName: modelName, + ContextTokens: copyOptionalInt(plan.Profile.ContextTokens), + DeviceIndex: copyOptionalInt(plan.Profile.DeviceIndex), + } + warning := "" + out, err := runLlamaServerSample(ctx, node, cfgNode, shellLlamaServerSample(plan.Port, plan.Profile.DeviceIndex)) + if err != nil { + warning = "llama-server RSS sample failed" + } else { + ram, ramOK, vram, vramOK := parseLlamaServerSample(out) + if !ramOK { + warning = "llama-server RSS sample failed" + } else { + obs.PeakRAMMB = ram + if plan.Profile.DeviceIndex != nil && vramOK { + obs.PeakVRAMMB = vram + } + } + } + if err := state.Update(func(latest *state.ClusterState) error { + latest.RecordObservation(obs) + return nil + }); err != nil && cmd != nil { + fmt.Fprintf(cmd.ErrOrStderr(), "warning: execution observation persistence failed: %v\n", err) + } + return warning +} + +func observationWallMS(elapsed time.Duration) int64 { + if elapsed <= 0 { + return 1 + } + if ms := elapsed.Milliseconds(); ms > 0 { + return ms + } + return 1 +} + +func copyOptionalInt(v *int) *int { + if v == nil { + return nil + } + n := *v + return &n +} + +func llamaServerObservationModelName(weights string) string { + weights = strings.TrimSpace(weights) + if weights == "" || weights == "." { + return "" + } + base := path.Base(weights) + if base == "." || base == "/" { + return "" + } + return base +} + +func applyLlamaServerPeakExclusions(plan *modelplan.ModelPlacementPlan, snap *models.ClusterSnapshot, st *state.ClusterState) { + if plan == nil || snap == nil || st == nil { + return + } + modelName := llamaServerObservationModelName(plan.Spec.WeightsPath) + kept := make([]modelplan.ModelCandidateScore, 0, len(plan.Candidates)) + for _, cand := range plan.Candidates { + node, ok := modelSnapshotNode(snap, cand.Node) + if ok { + if reason, blocked := placement.LlamaServerPeakExclusion(node, modelName, st); blocked { + plan.Excluded = append(plan.Excluded, modelplan.ModelExcludedCandidate{ + Node: cand.Node, + Reasons: []string{reason}, + }) + continue + } + } + kept = append(kept, cand) + } + plan.Candidates = kept + sort.Slice(plan.Excluded, func(i, j int) bool { + return plan.Excluded[i].Node < plan.Excluded[j].Node + }) + if len(plan.Candidates) == 0 { + plan.BestCandidate = "" + plan.Selected = nil + return + } + plan.BestCandidate = plan.Candidates[0].Node + plan.Selected = nil + for i := range snap.Nodes { + if snap.Nodes[i].Name == plan.BestCandidate { + selected := models.NewPlanProfile(snap.Nodes[i], plan.Spec, plan.TargetPort, plan.PublicationID) + plan.Selected = &selected + return + } + } +} + +func modelSnapshotNode(snap *models.ClusterSnapshot, name string) (models.NodeFacts, bool) { + if snap == nil { + return models.NodeFacts{}, false + } + for i := range snap.Nodes { + if snap.Nodes[i].Name == name { + return snap.Nodes[i], true + } + } + return models.NodeFacts{}, false +} + func mlxImportObserved(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig) (bool, error) { for _, tool := range node.Tools { if strings.EqualFold(tool.Name, models.ToolMLXServer) { diff --git a/cmd/axis/model_observation_test.go b/cmd/axis/model_observation_test.go new file mode 100644 index 00000000..e850ec0e --- /dev/null +++ b/cmd/axis/model_observation_test.go @@ -0,0 +1,499 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/toasterbook88/axis/internal/config" + "github.com/toasterbook88/axis/internal/modelplan" + "github.com/toasterbook88/axis/internal/models" + "github.com/toasterbook88/axis/internal/state" +) + +func TestLlamaServerStartRecordsPeakAndContext(t *testing.T) { + snap := testSnap() + zero := 0 + snap.Nodes[0].Resources.CPUCores = 8 + snap.Nodes[0].Resources.GPUs = []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", Index: &zero, IndexSource: models.IndexSourceNvidiaSMI, + VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + useFakeModelRunner(t) + var scripts []string + runLlamaServerSample = func(_ context.Context, node models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + if node.Name != "storage" { + t.Fatalf("sampled node %q", node.Name) + } + scripts = append(scripts, script) + return "axis-llama-sample:ram 40\naxis-llama-sample:vram 7\n", nil + } + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{ + "--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", + "--ctx-size", "2048", "--main-gpu", "0", "--format", "json", + }) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if len(scripts) != 1 { + t.Fatalf("sample calls = %d", len(scripts)) + } + if !strings.Contains(scripts[0], "nvidia-smi --query-gpu=index,memory.used --format=csv,noheader,nounits") || !strings.Contains(scripts[0], "_axis_want=0") { + t.Fatalf("script = %s", scripts[0]) + } + if strings.Contains(scripts[0], "s+=") { + t.Fatalf("sample sums VRAM: %s", scripts[0]) + } + var receipt models.ModelOperationReceipt + if err := json.Unmarshal(buf.Bytes(), &receipt); err != nil { + t.Fatalf("receipt: %v\n%s", err, buf.String()) + } + if len(receipt.Warnings) != 0 { + t.Fatalf("warnings = %v", receipt.Warnings) + } + obs := loadedLlamaObservation(t, "storage", "a.gguf") + if obs.PeakRAMMB != 40 || obs.PeakVRAMMB != 7 || !obs.LastSuccess || obs.WallTimeMS < 1 { + t.Fatalf("observation = %+v", obs) + } + if obs.Scope.Backend != "llama.cpp" || obs.Scope.Tool != "llama-server" || obs.Scope.Workload != models.ClassLlamaServer { + t.Fatalf("scope = %+v", obs.Scope) + } + if obs.ContextTokens == nil || *obs.ContextTokens != 2048 || obs.DeviceIndex == nil || *obs.DeviceIndex != 0 { + t.Fatalf("context/device = %+v %+v", obs.ContextTokens, obs.DeviceIndex) + } +} + +func TestLlamaServerStartWithoutDeviceSkipsVRAM(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + useFakeModelRunner(t) + var script string + runLlamaServerSample = func(_ context.Context, _ models.NodeFacts, _ *config.NodeConfig, got string) (string, error) { + script = got + return "axis-llama-sample:ram 11\naxis-llama-sample:vram 99\n", nil + } + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if strings.Contains(script, "nvidia-smi") { + t.Fatalf("unset device queried VRAM: %s", script) + } + obs := loadedLlamaObservation(t, "storage", "a.gguf") + if obs.PeakRAMMB != 11 || obs.PeakVRAMMB != 0 || obs.DeviceIndex != nil { + t.Fatalf("observation = %+v", obs) + } +} + +func TestLlamaServerRSSFailureWarnsAndStillStarts(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + useFakeModelRunner(t) + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return "", errors.New("ps failed") + } + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081", "--format", "text"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if !strings.Contains(buf.String(), "started /usr/local/bin/llama-server on storage:8081") || !strings.Contains(buf.String(), "warning: llama-server RSS sample failed") { + t.Fatalf("output = %q", buf.String()) + } + obs := loadedLlamaObservation(t, "storage", "a.gguf") + if obs.PeakRAMMB != 0 || obs.PeakVRAMMB != 0 || !obs.LastSuccess { + t.Fatalf("observation = %+v", obs) + } +} + +func TestProbeFailureRecordsNoObservation(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + calls := 0 + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return "axis-llama-sample:ram 5\n", nil + } + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + err := runModelStart(context.Background(), cmd, "storage", "/mnt/models/a.gguf", 8081, &probeFailRunner{}) + if err == nil || !strings.Contains(err.Error(), "probe failed") { + t.Fatalf("err = %v", err) + } + if calls != 0 { + t.Fatalf("probe failure sampled %d times", calls) + } + loaded, loadErr := state.Load() + if loadErr != nil { + t.Fatal(loadErr) + } + if len(loaded.Observations) != 0 { + t.Fatalf("observations = %+v", loaded.Observations) + } +} + +func TestWriterFailureStillRecordsLlamaObservation(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return "axis-llama-sample:ram 12\n", nil + } + want := errors.New("writer unavailable") + cmd := modelStartCmd() + cmd.SetOut(rejectingOutputWriter{err: want}) + err := runModelStart(context.Background(), cmd, "storage", "/mnt/models/a.gguf", 8081, &fakeModelRunner{}) + if !errors.Is(err, want) { + t.Fatalf("err = %v", err) + } + obs := loadedLlamaObservation(t, "storage", "a.gguf") + if obs.PeakRAMMB != 12 { + t.Fatalf("observation = %+v", obs) + } +} + +func TestObservationPersistFailureDoesNotRollBackStart(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + notDir := filepath.Join(t.TempDir(), "not-a-directory") + if err := os.WriteFile(notDir, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + t.Setenv("AXIS_HOME", notDir) + useFakeModelRunner(t) + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return "axis-llama-sample:ram 9\n", nil + } + cmd := modelStartCmd() + var out, errBuf bytes.Buffer + cmd.SetOut(&out) + cmd.SetErr(&errBuf) + cmd.SetArgs([]string{"--node", "storage", "--weights", "/mnt/models/a.gguf", "--port", "8081"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if !strings.Contains(out.String(), "started ") || !strings.Contains(errBuf.String(), "execution observation persistence failed") { + t.Fatalf("stdout=%q stderr=%q", out.String(), errBuf.String()) + } +} + +func TestOllamaAndMLXStartsDoNotRecordLlamaObservation(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Running: true, Listening: true} + snap.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified + snap.Nodes[0].Tools = append(snap.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + useFakeModelRunner(t) + calls := 0 + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return "", errors.New("llama sample must not run") + } + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return `{"models":[{"name":"mistral","model":"mistral"}]}`, nil + } + t.Cleanup(func() { runNodeScript = runOnNodeCapturing }) + + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral", "--format", "json"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + cmd = modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--mlx-model", "/mnt/models/qwen", "--port", "8080", "--format", "json"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if calls != 0 { + t.Fatalf("llama sample calls = %d", calls) + } + loaded, err := state.Load() + if err != nil { + t.Fatal(err) + } + if len(loaded.Observations) != 0 { + t.Fatalf("observations = %+v", loaded.Observations) + } +} + +func TestRunModelPlanExcludesFreshLlamaPeak(t *testing.T) { + stubModelSnapshot(t, llamaPlanSnapshot()) + if err := state.Update(func(st *state.ClusterState) error { + st.RecordObservation(models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: "tight", + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + ModelName: "qwen2.5-7b.gguf", + }, + ObservedAt: time.Now().UTC(), + LastSuccess: true, + WallTimeMS: 30, + PeakRAMMB: 8000, + }) + return nil + }); err != nil { + t.Fatal(err) + } + plan := executeModelPlan(t) + if plan.BestCandidate != "wide" { + t.Fatalf("best = %q", plan.BestCandidate) + } + if plan.Selected == nil || plan.Selected.Node != "wide" { + t.Fatalf("selected = %+v", plan.Selected) + } + if len(plan.Candidates) != 1 || plan.Candidates[0].Node != "wide" { + t.Fatalf("candidates = %+v", plan.Candidates) + } + var found bool + for _, ex := range plan.Excluded { + if ex.Node == "tight" { + found = true + if len(ex.Reasons) != 1 || ex.Reasons[0] != "empirical peak RAM 8000MB exceeds allocatable 512MB" { + t.Fatalf("reasons = %#v", ex.Reasons) + } + } + } + if !found { + t.Fatalf("excluded = %+v", plan.Excluded) + } +} + +func TestRunModelPlanWithoutObservationStaysPut(t *testing.T) { + stubModelSnapshot(t, llamaPlanSnapshot()) + plan := executeModelPlan(t) + if plan.BestCandidate != "tight" || len(plan.Candidates) != 2 { + t.Fatalf("best=%s candidates=%+v", plan.BestCandidate, plan.Candidates) + } + if plan.Selected == nil || plan.Selected.Node != "tight" { + t.Fatalf("selected = %+v", plan.Selected) + } +} + +func TestRunModelPlanReturnsStateLoadError(t *testing.T) { + stubModelSnapshot(t, llamaPlanSnapshot()) + notDir := filepath.Join(t.TempDir(), "not-a-directory") + if err := os.WriteFile(notDir, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + t.Setenv("AXIS_HOME", notDir) + cmd := modelPlanCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetArgs([]string{"qwen2.5-7b", "--format", "json"}) + err := cmd.Execute() + if err == nil { + t.Fatal("expected state load error") + } + if strings.Contains(err.Error(), "no eligible") { + t.Fatalf("load error treated as no peak: %v", err) + } +} + +func TestRunModelPlanAllExcludedByPeakFails(t *testing.T) { + snap := llamaPlanSnapshot() + snap.Nodes = snap.Nodes[:1] + stubModelSnapshot(t, snap) + if err := state.Update(func(st *state.ClusterState) error { + st.RecordObservation(models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: "tight", Workload: models.ClassLlamaServer, Backend: "llama.cpp", + Tool: "llama-server", ModelName: "qwen2.5-7b.gguf", + }, + ObservedAt: time.Now().UTC(), LastSuccess: true, WallTimeMS: 10, PeakRAMMB: 8000, + }) + return nil + }); err != nil { + t.Fatal(err) + } + cmd := modelPlanCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"qwen2.5-7b", "--format", "json"}) + err := cmd.Execute() + if err == nil || ExitCode(err) != ExitErrCommandFail || !strings.Contains(err.Error(), "no eligible") { + t.Fatalf("err = %v", err) + } + var plan modelplan.ModelPlacementPlan + if jsonErr := json.Unmarshal(buf.Bytes(), &plan); jsonErr != nil { + t.Fatalf("plan json: %v\n%s", jsonErr, buf.String()) + } + if plan.BestCandidate != "" || plan.Selected != nil || len(plan.Candidates) != 0 { + t.Fatalf("plan still names the excluded node: best=%q selected=%+v candidates=%+v", plan.BestCandidate, plan.Selected, plan.Candidates) + } +} + +func TestLlamaServerSampleScriptMaxRSSAndDeviceRow(t *testing.T) { + dir := t.TempDir() + writeSampleStub(t, dir, "fuser", "#!/bin/sh\necho '111 222'\n") + writeSampleStub(t, dir, "ps", `#!/bin/sh +pid= +mode= +while [ $# -gt 0 ]; do + case "$1" in + -p) pid=$2; shift 2 ;; + -o) mode=$2; shift 2 ;; + *) shift ;; + esac +done +case "$mode" in + comm=) printf '%s\n' llama-server ;; + rss=) + case "$pid" in + 111) printf '%s\n' ' 1500' ;; + 222) printf '%s\n' ' 2048' ;; + *) exit 1 ;; + esac + ;; + *) exit 1 ;; +esac +`) + writeSampleStub(t, dir, "nvidia-smi", "#!/bin/sh\nprintf '%s\n' '1, 9000' '0, 100'\n") + zero := 0 + out, err := runSampleShell(dir, shellLlamaServerSample(8081, &zero)) + if err != nil { + t.Fatalf("sample script: %v\n%s", err, out) + } + if !strings.Contains(out, "axis-llama-sample:ram 2") { + t.Fatalf("max RSS missing: %q", out) + } + if strings.Contains(out, "axis-llama-sample:ram 3") || strings.Contains(out, "axis-llama-sample:vram 9100") || strings.Contains(out, "axis-llama-sample:vram 9000") { + t.Fatalf("sample summed or picked the wrong row: %q", out) + } + if !strings.Contains(out, "axis-llama-sample:vram 100") { + t.Fatalf("index 0 row missing: %q", out) + } + + nilScript := shellLlamaServerSample(8081, nil) + if strings.Contains(nilScript, "nvidia-smi") { + t.Fatal("nil device index emits nvidia-smi") + } + marker := filepath.Join(dir, "smi-ran") + writeSampleStub(t, dir, "nvidia-smi", "#!/bin/sh\ntouch "+marker+"\nexit 1\n") + out, err = runSampleShell(dir, nilScript) + if err != nil { + t.Fatalf("nil-device script: %v\n%s", err, out) + } + if _, statErr := os.Stat(marker); !os.IsNotExist(statErr) { + t.Fatal("nil device index ran nvidia-smi") + } + + writeSampleStub(t, dir, "ps", "#!/bin/sh\nexit 1\n") + out, err = runSampleShell(dir, shellLlamaServerSample(8081, nil)) + if err == nil || strings.Contains(out, "axis-llama-sample:ram ") { + t.Fatalf("ps failure err=%v out=%q", err, out) + } + + writeSampleStub(t, dir, "fuser", "#!/bin/sh\nexit 1\n") + writeSampleStub(t, dir, "ps", "#!/bin/sh\nprintf '%s\n' llama-server\n") + out, err = runSampleShell(dir, shellLlamaServerSample(8081, nil)) + if err == nil { + t.Fatalf("no listener was a successful sample: %q", out) + } +} + +func loadedLlamaObservation(t *testing.T, node, model string) models.ExecutionObservation { + t.Helper() + loaded, err := state.Load() + if err != nil { + t.Fatal(err) + } + obs, ok := loaded.Observation(models.ObservationScope{ + Node: node, + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + ModelName: model, + }) + if !ok || obs == nil { + t.Fatalf("missing observation for %s/%s in %+v", node, model, loaded.Observations) + } + return *obs +} + +func llamaPlanSnapshot() *models.ClusterSnapshot { + return &models.ClusterSnapshot{ + Timestamp: time.Now().UTC(), + Nodes: []models.NodeFacts{ + { + Name: "tight", + Status: models.StatusComplete, + RAMAllocatableMB: 512, + Resources: &models.Resources{ + RAMFreeMB: 100000, RAMTotalMB: 128000, + Volumes: []models.Volume{{Mount: "/data/models", Kind: "local"}}, + GPUs: []models.GPUInfo{{ + Vendor: "nvidia", Model: "RTX 4090", VRAMMB: 24576, + VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }}, + }, + DiskWeights: []models.DiskWeight{{ + Name: "qwen2.5-7b", Path: "/data/models/qwen2.5-7b.gguf", + Bytes: 4 * 1024 * 1024 * 1024, Format: "gguf", + }}, + }, + { + Name: "wide", + Status: models.StatusComplete, + RAMAllocatableMB: 64000, + Resources: &models.Resources{RAMFreeMB: 100000, RAMTotalMB: 128000}, + }, + }, + } +} + +func executeModelPlan(t *testing.T) modelplan.ModelPlacementPlan { + t.Helper() + cmd := modelPlanCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"qwen2.5-7b", "--format", "json"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + var plan modelplan.ModelPlacementPlan + if err := json.Unmarshal(buf.Bytes(), &plan); err != nil { + t.Fatalf("plan: %v\n%s", err, buf.String()) + } + return plan +} + +func useFakeModelRunner(t *testing.T) { + t.Helper() + prev := defaultModelRunner + defaultModelRunner = &fakeModelRunner{} + t.Cleanup(func() { defaultModelRunner = prev }) +} + +func writeSampleStub(t *testing.T, dir, name, body string) { + t.Helper() + if err := os.WriteFile(filepath.Join(dir, name), []byte(body), 0o755); err != nil { + t.Fatal(err) + } +} + +func runSampleShell(dir, script string) (string, error) { + cmd := exec.Command("sh", "-c", script) + cmd.Env = []string{"PATH=" + dir + ":/usr/bin:/bin", "HOME=" + dir} + out, err := cmd.CombinedOutput() + return string(out), err +} diff --git a/cmd/axis/model_test.go b/cmd/axis/model_test.go index 65ab78bf..e2b3650d 100644 --- a/cmd/axis/model_test.go +++ b/cmd/axis/model_test.go @@ -115,8 +115,19 @@ func (f *fakeModelRunner) Query(_ context.Context, _ models.NodeFacts, _ *config }, nil } +func isolateModelState(t *testing.T) { + t.Helper() + t.Setenv("AXIS_HOME", t.TempDir()) + prev := runLlamaServerSample + runLlamaServerSample = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return "", errors.New("llama-server sample not stubbed") + } + t.Cleanup(func() { runLlamaServerSample = prev }) +} + func stubModelSnapshot(t *testing.T, snap *models.ClusterSnapshot) { t.Helper() + isolateModelState(t) prevLive := loadModelSnapshot loadModelSnapshot = func(context.Context) (*models.ClusterSnapshot, error) { return snap, nil } prevFetch := fetchModelInventorySnapshot @@ -131,6 +142,7 @@ func stubModelSnapshot(t *testing.T, snap *models.ClusterSnapshot) { func stubModelConfig(t *testing.T, cfg *config.Config) { t.Helper() + isolateModelState(t) prev := loadModelConfig loadModelConfig = func() (*config.Config, error) { return cfg, nil } t.Cleanup(func() { loadModelConfig = prev }) diff --git a/internal/models/model_operation.go b/internal/models/model_operation.go index c7f78ba2..24ad113a 100644 --- a/internal/models/model_operation.go +++ b/internal/models/model_operation.go @@ -58,5 +58,6 @@ type ModelOperationReceipt struct { VRAMFreeMeasured bool `json:"vram_free_measured,omitempty" yaml:"vram_free_measured,omitempty"` PortSource string `json:"port_source,omitempty" yaml:"port_source,omitempty"` DeviceNote string `json:"device_note,omitempty" yaml:"device_note,omitempty"` + Warnings []string `json:"warnings,omitempty" yaml:"warnings,omitempty"` Error string `json:"error,omitempty" yaml:"error,omitempty"` } diff --git a/internal/models/types.go b/internal/models/types.go index 57086fd2..a97ccfd0 100644 --- a/internal/models/types.go +++ b/internal/models/types.go @@ -727,6 +727,12 @@ type ExecutionObservation struct { WallTimeMS int64 `json:"wall_time_ms" yaml:"wall_time_ms"` PeakRAMMB int64 `json:"peak_ram_mb,omitempty" yaml:"peak_ram_mb,omitempty"` PeakVRAMMB int64 `json:"peak_vram_mb,omitempty" yaml:"peak_vram_mb,omitempty"` + // ContextTokens and DeviceIndex are optional facts from a llama-server + // start. They are not part of ObservationKey. A nil pointer on a later + // sample keeps the previous value; a non-nil pointer, including zero, + // replaces it. + ContextTokens *int `json:"context_tokens,omitempty" yaml:"context_tokens,omitempty"` + DeviceIndex *int `json:"device_index,omitempty" yaml:"device_index,omitempty"` // ModelName is the inference model name observed during execution // (e.g. "llama3.2:latest", "qwen2.5-coder:7b"). Populated when a model // name is extractable from the task command or description. Used by diff --git a/internal/placement/llama_server_peak.go b/internal/placement/llama_server_peak.go new file mode 100644 index 00000000..d00b68ec --- /dev/null +++ b/internal/placement/llama_server_peak.go @@ -0,0 +1,20 @@ +package placement + +import ( + "strings" + + "github.com/toasterbook88/axis/internal/models" + "github.com/toasterbook88/axis/internal/state" +) + +// LlamaServerPeakExclusion reports whether a fresh llama-server RAM peak for +// modelName exceeds this node's allocatable RAM. modelName is the weights +// file base stored on the start observation. It is not parsed from a task +// description. The caller passes the snapshot node it already has. +func LlamaServerPeakExclusion(n models.NodeFacts, modelName string, st *state.ClusterState) (string, bool) { + reqs := models.TaskRequirements{ + Workload: models.WorkloadProfileMatch{Class: models.ClassLlamaServer}, + RequiredTools: []string{"llama-server"}, + } + return empiricalPeakRAMExclusionReason(n, reqs, st, allocatableRAM(n), strings.TrimSpace(modelName)) +} diff --git a/internal/placement/llama_server_peak_test.go b/internal/placement/llama_server_peak_test.go new file mode 100644 index 00000000..4466f086 --- /dev/null +++ b/internal/placement/llama_server_peak_test.go @@ -0,0 +1,93 @@ +package placement + +import ( + "testing" + "time" + + "github.com/toasterbook88/axis/internal/models" + "github.com/toasterbook88/axis/internal/state" +) + +func llamaPeakNode(name string, allocatable int64) models.NodeFacts { + return models.NodeFacts{ + Name: name, + RAMAllocatableMB: allocatable, + Resources: &models.Resources{RAMTotalMB: allocatable + 2048, RAMFreeMB: allocatable}, + } +} + +func recordLlamaPeak(st *state.ClusterState, node, model string, peak int64, observedAt time.Time) { + st.RecordObservation(models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: node, + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + ModelName: model, + }, + ObservedAt: observedAt, + LastSuccess: true, + WallTimeMS: 20, + PeakRAMMB: peak, + }) +} + +func TestLlamaServerPeakExclusionUsesRecordedPeak(t *testing.T) { + now := time.Now().UTC() + st := &state.ClusterState{} + recordLlamaPeak(st, "tight", "a.gguf", 5000, now) + recordLlamaPeak(st, "wide", "a.gguf", 5000, now) + st.RecordObservation(models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: "tight", + Workload: models.ClassLocalLLMInference, + Backend: "ollama", + Tool: "ollama", + ModelName: "a.gguf", + }, + ObservedAt: now, + LastSuccess: true, + WallTimeMS: 20, + PeakRAMMB: 99999, + }) + + reason, blocked := LlamaServerPeakExclusion(llamaPeakNode("tight", 1000), "a.gguf", st) + if !blocked || reason != "empirical peak RAM 5000MB exceeds allocatable 1000MB" { + t.Fatalf("exclusion = %q blocked=%v", reason, blocked) + } + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("wide", 8000), "a.gguf", st); blocked { + t.Fatal("peak within allocatable RAM excluded the node") + } + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("tight", 1000), "other.gguf", st); blocked { + t.Fatal("a different model name reused the llama-server peak") + } + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("missing", 1000), "a.gguf", st); blocked { + t.Fatal("missing observation excluded the node") + } + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("tight", 1000), "a.gguf", nil); blocked { + t.Fatal("nil state excluded the node") + } + ollamaOnly := &state.ClusterState{} + ollamaOnly.RecordObservation(models.ExecutionObservation{ + Scope: models.ObservationScope{ + Node: "tight", + Workload: models.ClassLocalLLMInference, + Backend: "ollama", + Tool: "ollama", + ModelName: "a.gguf", + }, + ObservedAt: now, + LastSuccess: true, + WallTimeMS: 20, + PeakRAMMB: 99999, + }) + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("tight", 1000), "a.gguf", ollamaOnly); blocked { + t.Fatal("an ollama peak excluded a llama-server plan") + } + + stale := &state.ClusterState{} + recordLlamaPeak(stale, "tight", "a.gguf", 5000, now.Add(-(state.ObservationStaleAfter + time.Hour))) + if _, blocked := LlamaServerPeakExclusion(llamaPeakNode("tight", 1000), "a.gguf", stale); blocked { + t.Fatal("stale peak excluded the node") + } +} diff --git a/internal/state/observations.go b/internal/state/observations.go index 96cafb54..8d5e6929 100644 --- a/internal/state/observations.go +++ b/internal/state/observations.go @@ -59,9 +59,19 @@ func normalizeObservation(obs models.ExecutionObservation) models.ExecutionObser if obs.PeakVRAMMB < 0 { obs.PeakVRAMMB = 0 } + obs.ContextTokens = copyOptionalInt(obs.ContextTokens) + obs.DeviceIndex = copyOptionalInt(obs.DeviceIndex) return obs } +func copyOptionalInt(v *int) *int { + if v == nil { + return nil + } + n := *v + return &n +} + func weightedAverage(current int64, currentSamples int, next int64, nextSamples int) int64 { if current <= 0 { return next @@ -100,6 +110,12 @@ func mergeObservation(existing, next models.ExecutionObservation) models.Executi if next.ModelName != "" { merged.ModelName = next.ModelName } + if next.ContextTokens != nil { + merged.ContextTokens = copyOptionalInt(next.ContextTokens) + } + if next.DeviceIndex != nil { + merged.DeviceIndex = copyOptionalInt(next.DeviceIndex) + } return merged } diff --git a/internal/state/observations_test.go b/internal/state/observations_test.go index e49596bb..ee378c73 100644 --- a/internal/state/observations_test.go +++ b/internal/state/observations_test.go @@ -1,6 +1,8 @@ package state import ( + "encoding/json" + "strings" "testing" "time" @@ -208,6 +210,118 @@ func TestNormalizeObservationSyncsScopeModelName(t *testing.T) { } } +func TestMergeKeepsUnsetContextAndDeviceAndReplacesSetValues(t *testing.T) { + s := &ClusterState{} + scope := models.ObservationScope{ + Node: "storage", + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + ModelName: "a.gguf", + } + ctx := 2048 + zero := 0 + s.RecordObservation(models.ExecutionObservation{ + Scope: scope, + ObservedAt: time.Now().UTC(), + LastSuccess: true, + WallTimeMS: 10, + PeakRAMMB: 100, + ContextTokens: &ctx, + DeviceIndex: &zero, + }) + zero = 7 + ctx = 1 + s.RecordObservation(models.ExecutionObservation{ + Scope: scope, + ObservedAt: time.Now().UTC(), + LastSuccess: true, + WallTimeMS: 10, + PeakRAMMB: 50, + }) + obs, ok := s.Observation(scope) + if !ok || obs == nil || obs.ContextTokens == nil || *obs.ContextTokens != 2048 { + t.Fatalf("nil context sample cleared the previous value: %+v", obs) + } + if obs.DeviceIndex == nil || *obs.DeviceIndex != 0 { + t.Fatalf("nil device sample cleared index 0: %+v", obs) + } + if obs.PeakRAMMB != 100 { + t.Fatalf("peak = %d, want the previous max 100", obs.PeakRAMMB) + } + + nextCtx := 4096 + one := 1 + s.RecordObservation(models.ExecutionObservation{ + Scope: scope, + ObservedAt: time.Now().UTC(), + LastSuccess: true, + WallTimeMS: 10, + ContextTokens: &nextCtx, + DeviceIndex: &one, + }) + one = 9 + obs, ok = s.Observation(scope) + if !ok || obs.ContextTokens == nil || *obs.ContextTokens != 4096 || obs.DeviceIndex == nil || *obs.DeviceIndex != 1 { + t.Fatalf("set sample did not replace context and device: %+v", obs) + } + if len(s.Observations) != 1 { + t.Fatalf("optional fields changed the observation key: %d entries", len(s.Observations)) + } + + other := scope + if ObservationKey(other) != ObservationKey(obs.Scope) { + t.Fatal("context and device index changed ObservationKey") + } +} + +func TestObservationOptionalFieldsRoundTripIndexZero(t *testing.T) { + t.Setenv("AXIS_HOME", t.TempDir()) + s := &ClusterState{} + scope := models.ObservationScope{ + Node: "storage", + Workload: models.ClassLlamaServer, + Backend: "llama.cpp", + Tool: "llama-server", + } + zero := 0 + ctx := 1024 + s.RecordObservation(models.ExecutionObservation{ + Scope: scope, + ObservedAt: time.Now().UTC(), + LastSuccess: true, + WallTimeMS: 4, + ContextTokens: &ctx, + DeviceIndex: &zero, + }) + raw, err := json.Marshal(s.Observations[ObservationKey(scope)]) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(raw), `"device_index":0`) { + t.Fatalf("index 0 was omitted: %s", raw) + } + if err := s.Save(); err != nil { + t.Fatal(err) + } + loaded, err := Load() + if err != nil { + t.Fatal(err) + } + obs, ok := loaded.Observation(scope) + if !ok || obs.DeviceIndex == nil || *obs.DeviceIndex != 0 || obs.ContextTokens == nil || *obs.ContextTokens != 1024 { + t.Fatalf("round trip = %+v", obs) + } + unset := models.ExecutionObservation{Scope: scope, ObservedAt: time.Now().UTC(), WallTimeMS: 1, LastSuccess: true} + raw, err = json.Marshal(normalizeObservation(unset)) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(raw), "device_index") || strings.Contains(string(raw), "context_tokens") { + t.Fatalf("nil pointers were written as zero: %s", raw) + } +} + func TestObservationIsFresh(t *testing.T) { now := time.Now().UTC() fresh := models.ExecutionObservation{ObservedAt: now.Add(-time.Hour)} From fec781b22fb1cd8e996ed7a2755a74475223f1cc Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sat, 3 Oct 2026 20:06:01 -0400 Subject: [PATCH 8/9] fix(model): close the model-run review findings Encode every remote nvidia-smi GPU row before base64. Match an untagged Ollama name only to the :latest tag. Refuse an Ollama profile with refusals before the load POST. Keep ModelRunProfile YAML keys in snake_case. Print the placed Ollama model on a successful text receipt. Refresh the daemon cache after a successful Ollama unload. Deep-copy GPUInfo.Index when cloning a snapshot. Gate n-gpu-layers on the pinned GPU's measured free VRAM. Name Ollama and MLX in the model help and two stale sentences. Wait briefly for SIGKILL before the MLX stop script asserts death. --- AGENTS.md | 2 +- cmd/axis/model.go | 26 ++++- cmd/axis/model_mlx_test.go | 13 ++- cmd/axis/model_ollama_test.go | 86 +++++++++++++++++ docs/current-state.md | 2 +- internal/facts/remote_bundle.go | 2 +- internal/facts/remote_bundle_gpu_test.go | 35 +++++++ internal/modellife/ollama.go | 22 ++++- internal/modellife/ollama_test.go | 12 +++ internal/modellife/plan.go | 5 + internal/modellife/plan_profile_test.go | 60 ++++++++++++ internal/models/run_profile.go | 105 +++++++++++++-------- internal/models/run_profile_test.go | 37 ++++++++ internal/snapshotview/overlay.go | 4 + internal/snapshotview/snapshotview_test.go | 28 ++++++ 15 files changed, 390 insertions(+), 49 deletions(-) create mode 100644 internal/facts/remote_bundle_gpu_test.go diff --git a/AGENTS.md b/AGENTS.md index dc573cbc..ae78f31c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -249,7 +249,7 @@ heavy inference. | `axis agent [--auto-approve] [--autonomy MODE] [--plain] [--console] [--live]` | Agentic tool-calling assistant; REPL slash commands `/plan /todo /diff /undo /compact /autonomy /export /fleet`; default cluster context is daemon/disk cache (`LoadCached`), `--live` is an explicit discovery sweep. On an interactive TTY the transcript console is the default; `--plain` selects the legacy line reader (takes precedence over `--console`); `--console` forces the console. Tool approvals use an overlay (`y` yes, `n` no; Enter does not approve; timeout and cancel deny) | | `axis llm` | Removed; prints `use axis ai route` | | `axis ai` | Inference backends, roles, dry-run route resolve | -| `axis model` | List/inspect resident instances, dry-run placement planning, start/stop llama-server, await readiness, or query models | +| `axis model` | List/inspect resident instances, dry-run placement planning, start and stop llama-server or MLX, place or unload an Ollama model, await readiness, or query models | | `axis cluster` | Fleet snapshot: `status`, `summary` | | `axis node` | This machine: `facts` | | `axis cortex` | Distributed vector memory / event bus (resolves node via AXIS_CORTEX_NODE, role: cortex, or name cortex/foundry) | diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 99566b68..910dc027 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -62,7 +62,7 @@ var defaultModelRunner modelProcessRunner = liveModelRunner{} func modelCmd() *cobra.Command { cmd := &cobra.Command{ Use: "model", - Short: "Inspect resident models or manage llama-server on a named node", + Short: "Inspect resident models, plan a placement, or manage llama-server, Ollama, and MLX", } cmd.AddCommand(modelListCmd()) cmd.AddCommand(modelInspectCmd()) @@ -782,6 +782,11 @@ func writeModelStartReceipt(cmd *cobra.Command, receipt models.ModelOperationRec return printOutput(cmd.OutOrStdout(), receipt, format) } if receipt.Status == models.ModelOperationCompleted { + if receipt.Engine == models.EngineOllama { + _, err := fmt.Fprintf(cmd.OutOrStdout(), "placed ollama model %s on %s operation %s\n", + receipt.Model, receipt.Node, receipt.ID) + return err + } if _, err := fmt.Fprintf(cmd.OutOrStdout(), "started %s on %s:%d volume %s operation %s\n", receipt.Executable, receipt.Node, receipt.Port, receipt.Volume, receipt.ID); err != nil { return err @@ -866,7 +871,7 @@ func runModelStopGeneration(ctx context.Context, cmd *cobra.Command, generationI return fmt.Errorf("model generation %s is on node %s with status %s; refusing lifecycle mutation", generationID, instance.Node, instance.NodeStatus) } if instance.Engine == models.EngineOllama { - return stopOllamaGeneration(ctx, cmd, snap, instance, format, startedAt) + return stopOllamaGeneration(ctx, cmd, snap, instance, cacheAddr, format, startedAt) } if instance.Engine != models.EngineLlamaCpp && instance.Engine != models.EngineMLX { return fmt.Errorf("model generation %s uses unsupported stop engine %q", generationID, instance.Engine) @@ -1996,6 +2001,9 @@ func placeOllamaModel(ctx context.Context, cmd *cobra.Command, node models.NodeF if err := profile.Validate(); err != nil { return err } + if len(profile.Refusals) > 0 { + return fmt.Errorf("%s", strings.Join(profile.Refusals, "; ")) + } script, err := modellife.OllamaLoadScript(profile.OllamaModel, profile.OllamaKeepAlive, profile.OllamaNumCtx) if err != nil { return err @@ -2072,10 +2080,14 @@ func runOllamaModelStop(ctx context.Context, cmd *cobra.Command, nodeName, model StartedAt: time.Now().UTC(), CompletedAt: time.Now().UTC(), } - return writeModelOperationReceipt(cmd, receipt, format) + if err := writeModelOperationReceipt(cmd, receipt, format); err != nil { + return err + } + warnModelDaemonRefresh(cmd, cacheAddr, "manual") + return nil } -func stopOllamaGeneration(ctx context.Context, cmd *cobra.Command, snap *models.ClusterSnapshot, instance *models.ModelInstance, format string, startedAt time.Time) error { +func stopOllamaGeneration(ctx context.Context, cmd *cobra.Command, snap *models.ClusterSnapshot, instance *models.ModelInstance, cacheAddr, format string, startedAt time.Time) error { nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, instance.Node) if err != nil { return err @@ -2103,7 +2115,11 @@ func stopOllamaGeneration(ctx context.Context, cmd *cobra.Command, snap *models. if writeErr := writeModelOperationReceipt(cmd, receipt, format); writeErr != nil { return writeErr } - return unloadErr + if unloadErr != nil { + return unloadErr + } + warnModelDaemonRefresh(cmd, cacheAddr, "manual") + return nil } func runOllamaUnload(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig, modelName string) error { diff --git a/cmd/axis/model_mlx_test.go b/cmd/axis/model_mlx_test.go index 8537f5ab..9f11cc79 100644 --- a/cmd/axis/model_mlx_test.go +++ b/cmd/axis/model_mlx_test.go @@ -10,6 +10,7 @@ import ( "path/filepath" "strings" "testing" + "time" "github.com/toasterbook88/axis/internal/config" "github.com/toasterbook88/axis/internal/modelinventory" @@ -305,7 +306,17 @@ esac cmd := exec.Command("/bin/sh", "-c", shellStopTarget(modellife.StopTarget{Port: 8080, Engine: engine})) cmd.Env = []string{"PATH=" + dir + ":/usr/bin:/bin"} out, err := cmd.CombinedOutput() - return string(out), processStillRunning(targetProc.Process.Pid), err + text := string(out) + alive := processStillRunning(targetProc.Process.Pid) + // kill(2) can return before /proc shows the victim as dead. + if strings.Contains(text, modelStopMarker+"stopped") { + deadline := time.Now().Add(500 * time.Millisecond) + for alive && time.Now().Before(deadline) { + time.Sleep(10 * time.Millisecond) + alive = processStillRunning(targetProc.Process.Pid) + } + } + return text, alive, err } func processStillRunning(pid int) bool { diff --git a/cmd/axis/model_ollama_test.go b/cmd/axis/model_ollama_test.go index 0cd79446..d6b490bb 100644 --- a/cmd/axis/model_ollama_test.go +++ b/cmd/axis/model_ollama_test.go @@ -115,12 +115,23 @@ func TestOllamaStopUnloadsWithoutKillingTheServer(t *testing.T) { } t.Cleanup(func() { runNodeScript = prevScript }) + var refreshes int + prevRefresh := signalModelDaemonRefresh + signalModelDaemonRefresh = func(context.Context, string, string) error { + refreshes++ + return nil + } + t.Cleanup(func() { signalModelDaemonRefresh = prevRefresh }) + cmd := modelStopCmd() cmd.SetOut(&bytes.Buffer{}) cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) if err := cmd.Execute(); err != nil { t.Fatal(err) } + if refreshes != 1 { + t.Fatalf("daemon refreshes=%d", refreshes) + } if len(runner.stopTargets) != 0 || len(runner.stopped) != 0 { t.Fatalf("stop used the process killer: %#v", runner.stopTargets) } @@ -142,6 +153,9 @@ func TestOllamaStopUnloadsWithoutKillingTheServer(t *testing.T) { if err := cmd.Execute(); err == nil || !strings.Contains(err.Error(), "api/ps") { t.Fatalf("still listed err=%v", err) } + if refreshes != 1 { + t.Fatalf("failed unload refreshed the daemon: %d", refreshes) + } } func TestOllamaGenerationStopDoesNotReachProcessKill(t *testing.T) { @@ -167,13 +181,85 @@ func TestOllamaGenerationStopDoesNotReachProcessKill(t *testing.T) { return `{"models":[]}`, nil } t.Cleanup(func() { runNodeScript = prevScript }) + var refreshes int + prevRefresh := signalModelDaemonRefresh + signalModelDaemonRefresh = func(context.Context, string, string) error { + refreshes++ + return nil + } + t.Cleanup(func() { signalModelDaemonRefresh = prevRefresh }) cmd := modelStopCmd() cmd.SetOut(&bytes.Buffer{}) if err := runModelStopGeneration(context.Background(), cmd, want.GenerationID, "test.sock", "text", runner); err != nil { t.Fatal(err) } + if refreshes != 1 { + t.Fatalf("daemon refreshes=%d", refreshes) + } if len(runner.stopTargets) != 0 { t.Fatalf("process kill targets=%#v", runner.stopTargets) } } + +func TestOllamaStartTextNamesTheModel(t *testing.T) { + snap := testSnap() + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + return `{"models":[{"name":"mistral:latest","model":"mistral:latest"}]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + prevRefresh := signalModelDaemonRefresh + signalModelDaemonRefresh = func(context.Context, string, string) error { return nil } + t.Cleanup(func() { signalModelDaemonRefresh = prevRefresh }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral", "--format", "text"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + got := buf.String() + if !strings.Contains(got, "placed ollama model mistral on storage operation ") { + t.Fatalf("receipt=%q", got) + } + if strings.Contains(got, "started ") || strings.Contains(got, ":0") { + t.Fatalf("receipt used the llama-server line: %q", got) + } +} + +func TestOllamaStartRefusesProfileRefusals(t *testing.T) { + stubModelSnapshot(t, testSnap()) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + path := writeProfile(t, models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: "storage", + Engine: models.EngineOllama, + ArtifactKind: models.ArtifactOllamaModelName, + OllamaModel: "mistral", + BindHost: "127.0.0.1", + Refusals: []string{"weights are not on a named local volume"}, + }) + var calls int + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return `{"models":[{"name":"mistral"}]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--from-plan", path, "--format", "text"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "named local volume") { + t.Fatalf("err=%v", err) + } + if calls != 0 { + t.Fatalf("refusals still reached ollama: %d", calls) + } +} diff --git a/docs/current-state.md b/docs/current-state.md index 80bd5627..c11d38ed 100644 --- a/docs/current-state.md +++ b/docs/current-state.md @@ -108,7 +108,7 @@ Top-level commands currently registered in the binary: | `axis agent` | Agentic tool-calling assistant | Cluster tools + Layer-4 guarded `run_shell` / `run_on_node` / `axis_run_task`; injects nearest `AGENTS.md` into the system prompt when present; `--auto-approve` for safe commands; `--system` appends to system prompt | | `axis llm` | Removed | Prints `use: axis ai route` | | `axis model list\|inspect` | Inspect resident model instances | Daemon cache by default; `--live` explicitly performs a fresh cluster collection; text, JSON, and YAML output | -| `axis model start\|stop` | Manage llama-server | Requires an explicit node and port; start also requires a weight path on an observed local volume | +| `axis model start\|stop` | Manage llama-server, Ollama, and MLX | llama-server start requires a node, a port, and a weight path on an observed local volume; MLX start requires a node, a port, and a local model directory; Ollama load and unload require a node and a model name | | `axis cluster` | Fleet snapshot | `status` (cache-first for 5 minutes; `--live` sweeps), `summary` | | `axis node` | This machine | `facts` (localhost). Root `axis facts` still works | diff --git a/internal/facts/remote_bundle.go b/internal/facts/remote_bundle.go index d58ba7ad..0287e3b0 100644 --- a/internal/facts/remote_bundle.go +++ b/internal/facts/remote_bundle.go @@ -46,7 +46,7 @@ case "$(printf '%s' "$OS" | tr '[:upper:]' '[:lower:]')" in printf 'meminfo_b64=%s\n' "$(grep -E 'MemTotal|MemAvailable|MemFree' /proc/meminfo 2>/dev/null | base64 | tr -d '\n')" printf 'loadavg=%s\n' "$(cat /proc/loadavg 2>/dev/null)" printf 'pressure_b64=%s\n' "$(cat /proc/pressure/memory 2>/dev/null | base64 | tr -d '\n')" - printf 'gpu_b64=%s\n' "$(nvidia-smi --query-gpu=index,name,memory.total,memory.free --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //' | base64 | tr -d '\n')" + printf 'gpu_b64=%s\n' "$( { nvidia-smi --query-gpu=index,name,memory.total,memory.free --format=csv,noheader,nounits 2>/dev/null || lspci 2>/dev/null | grep -iE 'vga|3d' | sed 's/.*: //'; } | base64 | tr -d '\n')" printf 'identity=%s\n' "$(cat /etc/machine-id 2>/dev/null || cat /var/lib/dbus/machine-id 2>/dev/null)" printf 'battery=%s\n' "$(cat /sys/class/power_supply/BAT0/capacity /sys/class/power_supply/BAT1/capacity /sys/class/power_supply/BATT/capacity 2>/dev/null | head -1)" printf 'power=%s\n' "$(for n in AC ADP0 ACAD Mains; do s=$(cat /sys/class/power_supply/$n/status 2>/dev/null); [ -n "$s" ] && echo "$s" && break; done)" diff --git a/internal/facts/remote_bundle_gpu_test.go b/internal/facts/remote_bundle_gpu_test.go new file mode 100644 index 00000000..8eb4c3e2 --- /dev/null +++ b/internal/facts/remote_bundle_gpu_test.go @@ -0,0 +1,35 @@ +package facts + +import ( + "os" + "os/exec" + "path/filepath" + "strings" + "testing" +) + +func TestRemoteBundleKeepsEveryNvidiaSMIRow(t *testing.T) { + if !strings.Contains(remoteFactBundleScript, "{ nvidia-smi "+nvidiaSMIMemoryQuery) || + !strings.Contains(remoteFactBundleScript, "} | base64 | tr -d '\\n'") { + t.Fatal("nvidia-smi success path is not grouped into base64") + } + dir := t.TempDir() + stub := "#!/bin/sh\nprintf '%s\\n' '0, RTX 4090, 24576, 20000' '1, RTX 4090, 24576, 100'\n" + if err := os.WriteFile(filepath.Join(dir, "nvidia-smi"), []byte(stub), 0o755); err != nil { + t.Fatal(err) + } + cmd := exec.Command("bash", "-c", remoteFactBundleScript) + cmd.Env = append(os.Environ(), "PATH="+dir+":/usr/bin:/bin") + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("bundle: %v\n%s", err, out) + } + kv, err := parseRemoteFactBundle(string(out)) + if err != nil { + t.Fatal(err) + } + gpus := parseNvidiaSMIOutput(strings.TrimSpace(b64field(kv, "gpu_b64"))) + if len(gpus) != 2 || gpus[0].Index == nil || *gpus[0].Index != 0 || gpus[1].Index == nil || *gpus[1].Index != 1 { + t.Fatalf("gpus=%+v bundle=%s", gpus, out) + } +} diff --git a/internal/modellife/ollama.go b/internal/modellife/ollama.go index 52f7a1f7..386fbd43 100644 --- a/internal/modellife/ollama.go +++ b/internal/modellife/ollama.go @@ -72,13 +72,33 @@ func OllamaPSHasModel(body, name string) (bool, error) { } name = strings.TrimSpace(name) for _, model := range doc.Models { - if model.Name == name || model.Model == name { + if ollamaNameMatches(model.Name, name) || ollamaNameMatches(model.Model, name) { return true, nil } } return false, nil } +// ollamaNameMatches accepts an exact name, and the default tag when one side +// is untagged. A different tag, such as q4, does not match. +func ollamaNameMatches(listed, requested string) bool { + listed = strings.TrimSpace(listed) + requested = strings.TrimSpace(requested) + if listed == "" || requested == "" { + return false + } + if listed == requested { + return true + } + if !strings.Contains(requested, ":") && listed == requested+":latest" { + return true + } + if !strings.Contains(listed, ":") && requested == listed+":latest" { + return true + } + return false +} + func shellSingleQuote(s string) string { return "'" + strings.ReplaceAll(s, "'", `'"'"'`) + "'" } diff --git a/internal/modellife/ollama_test.go b/internal/modellife/ollama_test.go index e22007bf..5d7149c5 100644 --- a/internal/modellife/ollama_test.go +++ b/internal/modellife/ollama_test.go @@ -93,6 +93,18 @@ func TestOllamaPSHasModelReadsAPIPs(t *testing.T) { if _, err := OllamaPSHasModel("not-json", "mistral"); err == nil { t.Fatal("expected json error") } + ok, err = OllamaPSHasModel(`{"models":[{"name":"mistral:latest","model":"mistral:latest"}]}`, "mistral") + if err != nil || !ok { + t.Fatalf("untagged request latest listing ok=%v err=%v", ok, err) + } + ok, err = OllamaPSHasModel(`{"models":[{"name":"mistral","model":"mistral"}]}`, "mistral:latest") + if err != nil || !ok { + t.Fatalf("latest request untagged listing ok=%v err=%v", ok, err) + } + ok, err = OllamaPSHasModel(`{"models":[{"name":"mistral:q4","model":"mistral:q4"}]}`, "mistral") + if err != nil || ok { + t.Fatalf("other tag ok=%v err=%v", ok, err) + } } func ollamaJSONBody(t *testing.T, script string) map[string]any { diff --git a/internal/modellife/plan.go b/internal/modellife/plan.go index 758e4e7c..0f5aa012 100644 --- a/internal/modellife/plan.go +++ b/internal/modellife/plan.go @@ -98,6 +98,11 @@ func normalizeStartProfile(node models.NodeFacts, profile models.ModelRunProfile } } dev := models.ObserveLaunchDevice(node) + if profile.DeviceIndex != nil && dev.Kind == models.DeviceKindDiscrete { + if pinned, ok := models.LaunchDeviceForNvidiaIndex(node, *profile.DeviceIndex); ok { + dev = pinned + } + } profile.DeviceKind = dev.Kind profile.DeviceModel = dev.Model profile.Accelerator = dev.Accelerator diff --git a/internal/modellife/plan_profile_test.go b/internal/modellife/plan_profile_test.go index e2730564..5f38cf5d 100644 --- a/internal/modellife/plan_profile_test.go +++ b/internal/modellife/plan_profile_test.go @@ -199,6 +199,66 @@ func TestPlanStartProfileNGPULayersModeAndRefusals(t *testing.T) { } } +func TestPlanStartProfilePinUsesThatGPUsMeasuredVRAM(t *testing.T) { + node := storageNode() + zero, one := 0, 1 + node.Resources.GPUs = []models.GPUInfo{ + { + Vendor: "nvidia", Model: "GPU0", Index: &zero, IndexSource: models.IndexSourceNvidiaSMI, + VRAMMB: 24576, VRAMFreeMB: 20000, VRAMFreeMeasured: true, Capabilities: []string{"cuda"}, + }, + { + Vendor: "nvidia", Model: "GPU1", Index: &one, IndexSource: models.IndexSourceNvidiaSMI, + VRAMMB: 8192, Capabilities: []string{"cuda"}, + }, + } + pinned := readyProfile(node) + pinned.NGPULayersMode = "auto" + pinned.DeviceIndex = &one + pinned.IndexSource = models.IndexSourceNvidiaSMI + if _, err := PlanStartProfile(node, pinned); err == nil || !strings.Contains(err.Error(), "measured free VRAM") { + t.Fatalf("unmeasured pin err=%v", err) + } + + node.Resources.GPUs[1].VRAMFreeMB = 100 + node.Resources.GPUs[1].VRAMFreeMeasured = true + pinned = readyProfile(node) + pinned.NGPULayersMode = "auto" + pinned.DeviceIndex = &one + pinned.IndexSource = models.IndexSourceNvidiaSMI + plan, err := PlanStartProfile(node, pinned) + if err != nil { + t.Fatal(err) + } + if plan.Profile.DeviceModel != "GPU1" || plan.Profile.VRAMFreeMB != 100 || !plan.Profile.VRAMFreeMeasured { + t.Fatalf("pinned device=%s free=%d measured=%v", plan.Profile.DeviceModel, plan.Profile.VRAMFreeMB, plan.Profile.VRAMFreeMeasured) + } + if !strings.Contains(strings.Join(plan.Argv, " "), "--main-gpu 1") { + t.Fatalf("argv=%v", plan.Argv) + } + + unpinned := readyProfile(node) + unpinned.NGPULayersMode = "auto" + plan, err = PlanStartProfile(node, unpinned) + if err != nil { + t.Fatal(err) + } + if plan.Profile.DeviceModel != "GPU0" || plan.Profile.VRAMFreeMB != 20000 { + t.Fatalf("unpinned device=%s free=%d", plan.Profile.DeviceModel, plan.Profile.VRAMFreeMB) + } + + unified := storageNode() + unified.Resources.MemoryTopology = models.MemoryTopologyUnified + unified.Resources.GPUs = append([]models.GPUInfo(nil), node.Resources.GPUs...) + pinned = readyProfile(unified) + pinned.NGPULayersMode = "auto" + pinned.DeviceIndex = &zero + pinned.IndexSource = models.IndexSourceNvidiaSMI + if _, err := PlanStartProfile(unified, pinned); err == nil || !strings.Contains(err.Error(), "measured free VRAM") { + t.Fatalf("unified pin err=%v", err) + } +} + func TestPlanStartProfileCtxSizeDoesNotCheckVRAM(t *testing.T) { node := storageNode() ctx := 100000 diff --git a/internal/models/run_profile.go b/internal/models/run_profile.go index 428a1f9c..20130a74 100644 --- a/internal/models/run_profile.go +++ b/internal/models/run_profile.go @@ -55,45 +55,45 @@ func MainGPUPinNote(index *int) string { // ModelRunProfile is the launch description shared by model plan and model start. // It does not exec. Argv is derived from it. type ModelRunProfile struct { - Schema string `json:"schema"` - Node string `json:"node"` - Engine string `json:"engine"` - EngineBinary string `json:"engine_binary,omitempty"` - ToolName string `json:"tool_name,omitempty"` - SpecID string `json:"spec_id,omitempty"` - ArtifactKind string `json:"artifact_kind,omitempty"` - WeightsPath string `json:"weights_path,omitempty"` - OllamaModel string `json:"ollama_model,omitempty"` - Format ModelFormat `json:"format,omitempty"` - Quantization string `json:"quantization,omitempty"` - Volume string `json:"volume,omitempty"` - SpecSource string `json:"spec_source,omitempty"` - DeviceKind string `json:"device_kind,omitempty"` - DeviceIndex *int `json:"device_index,omitempty"` - IndexSource string `json:"index_source,omitempty"` - DeviceModel string `json:"device_model,omitempty"` - MemoryTopology MemoryTopology `json:"memory_topology,omitempty"` - VRAMFreeMB int64 `json:"vram_free_mb,omitempty"` - VRAMFreeMeasured bool `json:"vram_free_measured,omitempty"` - Accelerator string `json:"accelerator,omitempty"` - BindHost string `json:"bind_host"` - Port int `json:"port"` - ContextTokens *int `json:"context_tokens,omitempty"` - NGPULayers *int `json:"n_gpu_layers,omitempty"` - NGPULayersMode string `json:"n_gpu_layers_mode,omitempty"` - BatchSize *int `json:"batch_size,omitempty"` - UBatchSize *int `json:"ubatch_size,omitempty"` - Threads *int `json:"threads,omitempty"` - OllamaNumCtx *int `json:"ollama_num_ctx,omitempty"` - OllamaKeepAlive string `json:"ollama_keep_alive,omitempty"` - OllamaNumGPU *int `json:"ollama_num_gpu,omitempty"` - MLXModel string `json:"mlx_model,omitempty"` - PrefillStepSize *int `json:"prefill_step_size,omitempty"` - PromptCacheBytes *int64 `json:"prompt_cache_bytes,omitempty"` - KVBits *int `json:"kv_bits,omitempty"` - PortSource string `json:"port_source,omitempty"` - SnapshotPublicationID string `json:"snapshot_publication_id,omitempty"` - Refusals []string `json:"refusals,omitempty"` + Schema string `json:"schema" yaml:"schema"` + Node string `json:"node" yaml:"node"` + Engine string `json:"engine" yaml:"engine"` + EngineBinary string `json:"engine_binary,omitempty" yaml:"engine_binary,omitempty"` + ToolName string `json:"tool_name,omitempty" yaml:"tool_name,omitempty"` + SpecID string `json:"spec_id,omitempty" yaml:"spec_id,omitempty"` + ArtifactKind string `json:"artifact_kind,omitempty" yaml:"artifact_kind,omitempty"` + WeightsPath string `json:"weights_path,omitempty" yaml:"weights_path,omitempty"` + OllamaModel string `json:"ollama_model,omitempty" yaml:"ollama_model,omitempty"` + Format ModelFormat `json:"format,omitempty" yaml:"format,omitempty"` + Quantization string `json:"quantization,omitempty" yaml:"quantization,omitempty"` + Volume string `json:"volume,omitempty" yaml:"volume,omitempty"` + SpecSource string `json:"spec_source,omitempty" yaml:"spec_source,omitempty"` + DeviceKind string `json:"device_kind,omitempty" yaml:"device_kind,omitempty"` + DeviceIndex *int `json:"device_index,omitempty" yaml:"device_index,omitempty"` + IndexSource string `json:"index_source,omitempty" yaml:"index_source,omitempty"` + DeviceModel string `json:"device_model,omitempty" yaml:"device_model,omitempty"` + MemoryTopology MemoryTopology `json:"memory_topology,omitempty" yaml:"memory_topology,omitempty"` + VRAMFreeMB int64 `json:"vram_free_mb,omitempty" yaml:"vram_free_mb,omitempty"` + VRAMFreeMeasured bool `json:"vram_free_measured,omitempty" yaml:"vram_free_measured,omitempty"` + Accelerator string `json:"accelerator,omitempty" yaml:"accelerator,omitempty"` + BindHost string `json:"bind_host" yaml:"bind_host"` + Port int `json:"port" yaml:"port"` + ContextTokens *int `json:"context_tokens,omitempty" yaml:"context_tokens,omitempty"` + NGPULayers *int `json:"n_gpu_layers,omitempty" yaml:"n_gpu_layers,omitempty"` + NGPULayersMode string `json:"n_gpu_layers_mode,omitempty" yaml:"n_gpu_layers_mode,omitempty"` + BatchSize *int `json:"batch_size,omitempty" yaml:"batch_size,omitempty"` + UBatchSize *int `json:"ubatch_size,omitempty" yaml:"ubatch_size,omitempty"` + Threads *int `json:"threads,omitempty" yaml:"threads,omitempty"` + OllamaNumCtx *int `json:"ollama_num_ctx,omitempty" yaml:"ollama_num_ctx,omitempty"` + OllamaKeepAlive string `json:"ollama_keep_alive,omitempty" yaml:"ollama_keep_alive,omitempty"` + OllamaNumGPU *int `json:"ollama_num_gpu,omitempty" yaml:"ollama_num_gpu,omitempty"` + MLXModel string `json:"mlx_model,omitempty" yaml:"mlx_model,omitempty"` + PrefillStepSize *int `json:"prefill_step_size,omitempty" yaml:"prefill_step_size,omitempty"` + PromptCacheBytes *int64 `json:"prompt_cache_bytes,omitempty" yaml:"prompt_cache_bytes,omitempty"` + KVBits *int `json:"kv_bits,omitempty" yaml:"kv_bits,omitempty"` + PortSource string `json:"port_source,omitempty" yaml:"port_source,omitempty"` + SnapshotPublicationID string `json:"snapshot_publication_id,omitempty" yaml:"snapshot_publication_id,omitempty"` + Refusals []string `json:"refusals,omitempty" yaml:"refusals,omitempty"` } // LaunchDevice is the one device a llama-server launch is allowed to name. @@ -183,6 +183,33 @@ func ObserveLaunchDevice(node NodeFacts) LaunchDevice { } } +// LaunchDeviceForNvidiaIndex returns the discrete GPU with this nvidia-smi +// index. It does not fall back to a different GPU. +func LaunchDeviceForNvidiaIndex(node NodeFacts, index int) (LaunchDevice, bool) { + if node.Resources == nil { + return LaunchDevice{}, false + } + for _, gpu := range node.Resources.GPUs { + if gpu.Index == nil || *gpu.Index != index || gpu.IndexSource != IndexSourceNvidiaSMI { + continue + } + acc, ok := discreteAccelerator(gpu) + if !ok { + return LaunchDevice{}, false + } + free, measured := MeasuredFreeVRAM(gpu) + return LaunchDevice{ + Kind: DeviceKindDiscrete, + Model: gpu.Model, + Accelerator: acc, + MemoryTopology: node.Resources.MemoryTopology, + VRAMFreeMB: free, + VRAMFreeMeasured: measured, + }, true + } + return LaunchDevice{}, false +} + func unifiedLaunchDevice(res *Resources) LaunchDevice { dev := LaunchDevice{ Kind: DeviceKindUnified, diff --git a/internal/models/run_profile_test.go b/internal/models/run_profile_test.go index 289fc4e8..d66af3d3 100644 --- a/internal/models/run_profile_test.go +++ b/internal/models/run_profile_test.go @@ -3,6 +3,8 @@ package models import ( "strings" "testing" + + "gopkg.in/yaml.v3" ) func TestNamedLocalVolumeSkipsNetworkAndKeepsLongestMount(t *testing.T) { @@ -236,3 +238,38 @@ func TestValidateMLXRequiresDirectoryAndRefusesForeignFields(t *testing.T) { t.Fatalf("kv err=%v", err) } } + +func TestModelRunProfileYAMLKeepsSnakeCase(t *testing.T) { + zero := 0 + profile := ModelRunProfile{ + Schema: ModelRunSchema, + Node: "storage", + Engine: EngineLlamaCpp, + EngineBinary: "/usr/local/bin/llama-server", + DeviceIndex: &zero, + BindHost: "127.0.0.1", + Port: 8080, + } + data, err := yaml.Marshal(profile) + if err != nil { + t.Fatal(err) + } + text := string(data) + for _, bad := range []string{"enginebinary:", "deviceindex:", "bindhost:", "spec_id:", "batch_size:"} { + if strings.Contains(text, bad) { + t.Fatalf("yaml contains %s\n%s", bad, text) + } + } + for _, want := range []string{"engine_binary:", "device_index: 0", "bind_host:", "port: 8080"} { + if !strings.Contains(text, want) { + t.Fatalf("yaml missing %s\n%s", want, text) + } + } + var decoded ModelRunProfile + if err := yaml.Unmarshal(data, &decoded); err != nil { + t.Fatal(err) + } + if decoded.EngineBinary != profile.EngineBinary || decoded.DeviceIndex == nil || *decoded.DeviceIndex != 0 { + t.Fatalf("decoded=%+v", decoded) + } +} diff --git a/internal/snapshotview/overlay.go b/internal/snapshotview/overlay.go index 41e30bee..218a7d37 100644 --- a/internal/snapshotview/overlay.go +++ b/internal/snapshotview/overlay.go @@ -16,6 +16,10 @@ func cloneGPUInfos(gpus []models.GPUInfo) []models.GPUInfo { for i, gpu := range gpus { gpuCopy := gpu gpuCopy.Capabilities = append([]string(nil), gpu.Capabilities...) + if gpu.Index != nil { + index := *gpu.Index + gpuCopy.Index = &index + } cloned[i] = gpuCopy } return cloned diff --git a/internal/snapshotview/snapshotview_test.go b/internal/snapshotview/snapshotview_test.go index 0dbbe387..b60254f2 100644 --- a/internal/snapshotview/snapshotview_test.go +++ b/internal/snapshotview/snapshotview_test.go @@ -87,6 +87,34 @@ func TestCloneDeepCopiesResources(t *testing.T) { } } +func TestCloneDeepCopiesGPUIndex(t *testing.T) { + index := 1 + orig := &models.ClusterSnapshot{ + Nodes: []models.NodeFacts{{ + Name: "gpu-node", + Resources: &models.Resources{ + GPUs: []models.GPUInfo{{ + Model: "RTX 4090", + Index: &index, + }}, + }, + }}, + } + clone := snapshotview.Clone(orig) + *clone.Nodes[0].Resources.GPUs[0].Index = 9 + if *orig.Nodes[0].Resources.GPUs[0].Index != 1 { + t.Fatalf("clone index write changed original to %d", *orig.Nodes[0].Resources.GPUs[0].Index) + } + if clone.Nodes[0].Resources.GPUs[0].Index == orig.Nodes[0].Resources.GPUs[0].Index { + t.Fatal("clone shares the GPU index pointer") + } + + // Resources pointer itself must be different. + if clone.Nodes[0].Resources == orig.Nodes[0].Resources { + t.Error("expected cloned Resources to be a new pointer") + } +} + func TestCloneDeepCopiesAddressesAndTools(t *testing.T) { orig := &models.ClusterSnapshot{ Nodes: []models.NodeFacts{ From 71ad6e609df2606b5849d1a27cee2ce5c2e6abc7 Mon Sep 17 00:00:00 2001 From: AXIS Contributor Date: Sun, 4 Oct 2026 17:44:44 -0400 Subject: [PATCH 9/9] feat(model): start on the node that already has the runtime axis model start reads the snapshot and names one complete node whose runtime already has that model. A resident model is reported and left running. An Ollama server that is listening and lists the model is loaded on loopback, and that load is accepted only when the server returns done_reason load. Ollama SSH is refused unless the node is complete and the server is listening. --- AGENTS.md | 2 +- cmd/axis/main.go | 4 +- cmd/axis/model.go | 22 ++- cmd/axis/model_daily.go | 93 +++++++++ cmd/axis/model_daily_test.go | 187 ++++++++++++++++++ cmd/axis/model_mlx_test.go | 4 +- cmd/axis/model_observation_test.go | 1 + cmd/axis/model_ollama_test.go | 6 + cmd/axis/model_run_profile.go | 18 ++ cmd/axis/task.go | 2 +- cmd/axis/task_context_enrichment_test.go | 2 +- .../testdata/task_context_turboquant.golden | 2 +- docs/current-state.md | 2 +- internal/chat/system.go | 2 +- internal/modellife/ollama.go | 23 ++- internal/modellife/ollama_test.go | 19 ++ internal/modellife/serve_pick.go | 124 ++++++++++++ internal/modellife/serve_pick_test.go | 104 ++++++++++ 18 files changed, 602 insertions(+), 15 deletions(-) create mode 100644 cmd/axis/model_daily.go create mode 100644 cmd/axis/model_daily_test.go create mode 100644 internal/modellife/serve_pick.go create mode 100644 internal/modellife/serve_pick_test.go diff --git a/AGENTS.md b/AGENTS.md index ae78f31c..76108a26 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -249,7 +249,7 @@ heavy inference. | `axis agent [--auto-approve] [--autonomy MODE] [--plain] [--console] [--live]` | Agentic tool-calling assistant; REPL slash commands `/plan /todo /diff /undo /compact /autonomy /export /fleet`; default cluster context is daemon/disk cache (`LoadCached`), `--live` is an explicit discovery sweep. On an interactive TTY the transcript console is the default; `--plain` selects the legacy line reader (takes precedence over `--console`); `--console` forces the console. Tool approvals use an overlay (`y` yes, `n` no; Enter does not approve; timeout and cancel deny) | | `axis llm` | Removed; prints `use axis ai route` | | `axis ai` | Inference backends, roles, dry-run route resolve | -| `axis model` | List/inspect resident instances, dry-run placement planning, start and stop llama-server or MLX, place or unload an Ollama model, await readiness, or query models | +| `axis model` | List/inspect resident instances, dry-run placement planning, `start ` on the node that already has the runtime, or pin llama-server, Ollama, and MLX start/stop, await readiness, or query models | | `axis cluster` | Fleet snapshot: `status`, `summary` | | `axis node` | This machine: `facts` | | `axis cortex` | Distributed vector memory / event bus (resolves node via AXIS_CORTEX_NODE, role: cortex, or name cortex/foundry) | diff --git a/cmd/axis/main.go b/cmd/axis/main.go index 0430fea1..8678f9c9 100644 --- a/cmd/axis/main.go +++ b/cmd/axis/main.go @@ -47,12 +47,12 @@ func newRootCmd() *cobra.Command { axis cluster status every node (cache-first for 5 minutes; --live to sweep) axis node facts this machine axis agent ask questions (advisory) - axis model start llama-server on a named node (--node --weights --port) + axis model start one model name; pick the node that already has the runtime axis daemon status local cache axis status, axis facts, axis summary, and axis doctor still work. axis chat and axis llm were removed; use axis agent and axis ai route.`, - Example: " axis cluster status\n axis node facts\n axis agent\n axis model start --node storage --weights /mnt/models/a.gguf --port 8081", + Example: " axis cluster status\n axis node facts\n axis agent\n axis model start mistral", PersistentPreRunE: func(cmd *cobra.Command, args []string) error { ui.Init(noColor) diff --git a/cmd/axis/model.go b/cmd/axis/model.go index 910dc027..b53a9329 100644 --- a/cmd/axis/model.go +++ b/cmd/axis/model.go @@ -106,18 +106,30 @@ func modelStartCmd() *cobra.Command { var promptCacheBytes int64 var live bool cmd := &cobra.Command{ - Use: "start", - Short: "Start llama-server, place an Ollama model, or start mlx_lm.server (its HTTP API is not for production)", + Use: "start [model]", + Short: "Place a model on the node that already has its runtime", + Long: `Daily path: axis model start + +Read the snapshot, pick one complete node that already has the runtime, and print that choice. A resident model is reported and left running. An Ollama server that is already listening and lists the model is loaded on 127.0.0.1:11434. + +Explicit pin: --node with --weights and --port, --ollama-model, --mlx-model, or --from-plan. mlx_lm.server's HTTP API is not for production.`, + Args: cobra.MaximumNArgs(1), SilenceUsage: true, PreRunE: func(cmd *cobra.Command, args []string) error { if err := validateOutputFormat(&format, "text", "json", "yaml")(cmd, args); err != nil { return err } + if len(args) == 1 { + return requireDailyModelStart(cmd) + } return requireModelStartIdentity(cmd) }, RunE: func(cmd *cobra.Command, args []string) error { ctx, cancel := context.WithTimeout(cmd.Context(), 45*time.Second) defer cancel() + if len(args) == 1 { + return runDailyModelStart(ctx, cmd, args[0]) + } return runModelStart(ctx, cmd, node, weights, port, defaultModelRunner) }, } @@ -2004,6 +2016,9 @@ func placeOllamaModel(ctx context.Context, cmd *cobra.Command, node models.NodeF if len(profile.Refusals) > 0 { return fmt.Errorf("%s", strings.Join(profile.Refusals, "; ")) } + if err := modellife.OllamaServerReady(node); err != nil { + return err + } script, err := modellife.OllamaLoadScript(profile.OllamaModel, profile.OllamaKeepAlive, profile.OllamaNumCtx) if err != nil { return err @@ -2123,6 +2138,9 @@ func stopOllamaGeneration(ctx context.Context, cmd *cobra.Command, snap *models. } func runOllamaUnload(ctx context.Context, node models.NodeFacts, cfgNode *config.NodeConfig, modelName string) error { + if err := modellife.OllamaServerReady(node); err != nil { + return err + } script, err := modellife.OllamaUnloadScript(modelName) if err != nil { return err diff --git a/cmd/axis/model_daily.go b/cmd/axis/model_daily.go new file mode 100644 index 00000000..e6e4af75 --- /dev/null +++ b/cmd/axis/model_daily.go @@ -0,0 +1,93 @@ +package main + +import ( + "context" + "fmt" + "strings" + "time" + + "github.com/spf13/cobra" + + "github.com/toasterbook88/axis/internal/api" + "github.com/toasterbook88/axis/internal/modellife" + "github.com/toasterbook88/axis/internal/models" +) + +func runDailyModelStart(ctx context.Context, cmd *cobra.Command, modelName string) error { + modelName = strings.TrimSpace(modelName) + if modelName == "" { + return fmt.Errorf("model name is required") + } + live, _ := cmd.Flags().GetBool("live") + cacheAddr, _ := cmd.Flags().GetString("cache-addr") + if cacheAddr == "" { + cacheAddr = api.DefaultAddr() + } + format, _ := cmd.Flags().GetString("format") + if format == "" { + format = "text" + } + snap, source, err := loadModelCommandSnapshot(ctx, live, cacheAddr, "start", false) + if err != nil { + return err + } + pick, err := modellife.PickServingNode(snap.Nodes, modelName) + if err != nil { + return err + } + if pick.Already { + return writeDailyPick(cmd, pick, snap, source, format) + } + if pick.Runtime != models.EngineOllama { + return fmt.Errorf("node %s runtime %s is not a load target for %s", pick.Node, pick.Runtime, modelName) + } + if format == "text" { + if _, err := fmt.Fprintln(cmd.OutOrStdout(), pick.Sentence()); err != nil { + return err + } + } + nf, cfgNode, err := resolveModelNodeFromSnapshot(snap, pick.Node) + if err != nil { + return err + } + profile := models.ModelRunProfile{ + Schema: models.ModelRunSchema, + Node: pick.Node, + Engine: models.EngineOllama, + ArtifactKind: models.ArtifactOllamaModelName, + OllamaModel: modelName, + BindHost: "127.0.0.1", + } + if snap.Publication != nil { + profile.SnapshotPublicationID = snap.Publication.ID + } + return placeOllamaModel(ctx, cmd, nf, cfgNode, profile, source, snap, time.Now().UTC(), format) +} + +func writeDailyPick(cmd *cobra.Command, pick modellife.ServingPick, snap *models.ClusterSnapshot, source, format string) error { + if format == "text" || format == "" { + _, err := fmt.Fprintln(cmd.OutOrStdout(), pick.Sentence()) + return err + } + now := time.Now().UTC() + receipt := models.ModelOperationReceipt{ + Schema: "axis.model-operation/v1", + ID: models.GenerateID("mo"), + Action: models.ModelOperationStart, + Status: models.ModelOperationCompleted, + Disposition: "already_serving", + Node: pick.Node, + Engine: pick.Runtime, + Model: pick.Model, + SnapshotSource: source, + StartedAt: now, + CompletedAt: now, + } + if snap != nil { + receipt.SnapshotAt = snap.Timestamp + if snap.Publication != nil { + receipt.PublicationID = snap.Publication.ID + } + } + return printOutput(cmd.OutOrStdout(), receipt, format) +} diff --git a/cmd/axis/model_daily_test.go b/cmd/axis/model_daily_test.go new file mode 100644 index 00000000..4d51f7b8 --- /dev/null +++ b/cmd/axis/model_daily_test.go @@ -0,0 +1,187 @@ +package main + +import ( + "bytes" + "context" + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/config" + "github.com/toasterbook88/axis/internal/models" +) + +func TestModelStartDailyPathPicksListeningOllama(t *testing.T) { + snap := &models.ClusterSnapshot{Nodes: []models.NodeFacts{ + { + Name: "zeta", + Status: models.StatusComplete, + Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral:latest"}}, + }, + { + Name: "alpha", + Status: models.StatusComplete, + Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral"}}, + }, + {Name: "partial", Status: models.StatusPartial, Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral"}}}, + }} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "alpha"}, {Name: "zeta"}}}) + var scripts []string + prevScript := runNodeScript + runNodeScript = func(_ context.Context, node models.NodeFacts, _ *config.NodeConfig, script string) (string, error) { + if node.Name != "alpha" { + t.Fatalf("ssh node=%s", node.Name) + } + scripts = append(scripts, script) + return `{"models":[{"name":"mistral","model":"mistral"}]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + prevRefresh := signalModelDaemonRefresh + signalModelDaemonRefresh = func(context.Context, string, string) error { return nil } + t.Cleanup(func() { signalModelDaemonRefresh = prevRefresh }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"mistral"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + got := buf.String() + if !strings.Contains(got, "picked alpha: ollama already listening\n") { + t.Fatalf("stdout=%q", got) + } + if !strings.Contains(got, "placed ollama model mistral on alpha operation ") { + t.Fatalf("stdout=%q", got) + } + if len(scripts) != 1 || !strings.Contains(scripts[0], `grep -q '"done_reason":"load"'`) { + t.Fatalf("scripts=%v", scripts) + } +} + +func TestModelStartDailyPathReportsResidentModelWithoutSSH(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + snap.Nodes[0].ResidentModels = []models.ResidentModel{{Name: "mistral", Runtime: "llama.cpp", Port: 8080}} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + var calls int + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return "", nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + runner := &fakeModelRunner{} + prevRunner := defaultModelRunner + defaultModelRunner = runner + t.Cleanup(func() { defaultModelRunner = prevRunner }) + + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"mistral", "--format", "text"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + if buf.String() != "picked storage: llama.cpp already serving mistral\n" { + t.Fatalf("stdout=%q", buf.String()) + } + if calls != 0 || len(runner.started) != 0 { + t.Fatalf("calls=%d started=%v", calls, runner.started) + } +} + +func TestModelStartDailyPathJSONReportsAlreadyServing(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + snap.Nodes[0].ResidentModels = []models.ResidentModel{{Name: "mistral", Runtime: "ollama"}} + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + cmd := modelStartCmd() + var buf bytes.Buffer + cmd.SetOut(&buf) + cmd.SetArgs([]string{"mistral", "--format", "json"}) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + got := buf.String() + if !strings.Contains(got, `"disposition": "already_serving"`) || !strings.Contains(got, `"node": "storage"`) || !strings.Contains(got, `"engine": "ollama"`) { + t.Fatalf("receipt=%s", got) + } +} + +func TestModelStartDailyPathRefusesAFlagBoard(t *testing.T) { + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"mistral", "--node", "storage", "--ollama-model", "mistral"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "picks the node") || !strings.Contains(err.Error(), "--node") { + t.Fatalf("err=%v", err) + } +} + +func TestModelStartDailyPathRefusesWhenNoRuntimeIsPresent(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + var calls int + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return "", nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"mistral"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "no complete node already has a runtime") { + t.Fatalf("err=%v", err) + } + if calls != 0 { + t.Fatalf("ssh calls=%d", calls) + } +} + +func TestOllamaStartRefusesBeforeSSHWhenTheServerIsNotListening(t *testing.T) { + snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + stubModelSnapshot(t, snap) + stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) + var calls int + prevScript := runNodeScript + runNodeScript = func(context.Context, models.NodeFacts, *config.NodeConfig, string) (string, error) { + calls++ + return `{"models":[{"name":"mistral"}]}`, nil + } + t.Cleanup(func() { runNodeScript = prevScript }) + cmd := modelStartCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) + err := cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "no listening ollama") { + t.Fatalf("err=%v", err) + } + if calls != 0 { + t.Fatalf("ssh calls=%d", calls) + } + + snap.Nodes[0].Status = models.StatusPartial + snap.Nodes[0].Ollama = &models.OllamaInfo{Listening: true} + cmd = modelStopCmd() + cmd.SetOut(&bytes.Buffer{}) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"--node", "storage", "--ollama-model", "mistral"}) + err = cmd.Execute() + if err == nil || !strings.Contains(err.Error(), "refusing ollama ssh") { + t.Fatalf("stop err=%v", err) + } + if calls != 0 { + t.Fatalf("stop ssh calls=%d", calls) + } +} diff --git a/cmd/axis/model_mlx_test.go b/cmd/axis/model_mlx_test.go index 9f11cc79..7f2aac66 100644 --- a/cmd/axis/model_mlx_test.go +++ b/cmd/axis/model_mlx_test.go @@ -20,8 +20,8 @@ import ( func TestMLXHelpSaysHTTPAPIIsNotForProduction(t *testing.T) { cmd := modelStartCmd() - if !strings.Contains(cmd.Short, "not for production") { - t.Fatalf("short=%q", cmd.Short) + if !strings.Contains(cmd.Long, "not for production") { + t.Fatalf("long=%q", cmd.Long) } flag := cmd.Flags().Lookup("mlx-model") if flag == nil || !strings.Contains(flag.Usage, "not for production") { diff --git a/cmd/axis/model_observation_test.go b/cmd/axis/model_observation_test.go index e850ec0e..db55dad9 100644 --- a/cmd/axis/model_observation_test.go +++ b/cmd/axis/model_observation_test.go @@ -195,6 +195,7 @@ func TestObservationPersistFailureDoesNotRollBackStart(t *testing.T) { func TestOllamaAndMLXStartsDoNotRecordLlamaObservation(t *testing.T) { snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Running: true, Listening: true} snap.Nodes[0].Resources.MemoryTopology = models.MemoryTopologyUnified snap.Nodes[0].Tools = append(snap.Nodes[0].Tools, models.ToolInfo{Name: "mlx_lm.server", Path: "/usr/local/bin/mlx_lm.server"}) diff --git a/cmd/axis/model_ollama_test.go b/cmd/axis/model_ollama_test.go index d6b490bb..063848be 100644 --- a/cmd/axis/model_ollama_test.go +++ b/cmd/axis/model_ollama_test.go @@ -24,6 +24,7 @@ func TestOllamaModelAndWeightsAreMutuallyExclusive(t *testing.T) { func TestOllamaStartPlacesOnExistingLoopbackServer(t *testing.T) { snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Running: true, Listening: true} stubModelSnapshot(t, snap) stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) @@ -101,6 +102,8 @@ func TestOllamaStartPlacesOnExistingLoopbackServer(t *testing.T) { func TestOllamaStopUnloadsWithoutKillingTheServer(t *testing.T) { snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Listening: true} stubModelSnapshot(t, snap) stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) runner := &fakeModelRunner{} @@ -163,6 +166,7 @@ func TestOllamaGenerationStopDoesNotReachProcessKill(t *testing.T) { snap.Nodes[0].ResidentModels[0].Runtime = "ollama" snap.Nodes[0].ResidentModels[0].Name = "mistral" snap.Nodes[0].ResidentModels[0].Port = 11434 + snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Listening: true} want := modelinventory.FromSnapshot(snap, "daemon-cache").Instances[0] if want.Engine != "ollama" { t.Fatalf("engine=%s", want.Engine) @@ -204,6 +208,8 @@ func TestOllamaGenerationStopDoesNotReachProcessKill(t *testing.T) { func TestOllamaStartTextNamesTheModel(t *testing.T) { snap := testSnap() + snap.Nodes[0].Status = models.StatusComplete + snap.Nodes[0].Ollama = &models.OllamaInfo{Installed: true, Listening: true} stubModelSnapshot(t, snap) stubModelConfig(t, &config.Config{Nodes: []config.NodeConfig{{Name: "storage"}}}) prevScript := runNodeScript diff --git a/cmd/axis/model_run_profile.go b/cmd/axis/model_run_profile.go index 55461ff3..e9ecc470 100644 --- a/cmd/axis/model_run_profile.go +++ b/cmd/axis/model_run_profile.go @@ -11,6 +11,24 @@ import ( "github.com/toasterbook88/axis/internal/models" ) +func requireDailyModelStart(cmd *cobra.Command) error { + var used []string + for _, name := range []string{ + "node", "weights", "port", "from-plan", "n-gpu-layers", "ctx-size", + "batch-size", "ubatch-size", "threads", "main-gpu", "ollama-model", + "ollama-keep-alive", "ollama-num-ctx", "mlx-model", "prefill-step-size", + "prompt-cache-bytes", "kv-bits", + } { + if cmd.Flags().Changed(name) { + used = append(used, "--"+name) + } + } + if len(used) == 0 { + return nil + } + return fmt.Errorf("axis model start picks the node; do not combine it with %s", strings.Join(used, ", ")) +} + func requireModelStartIdentity(cmd *cobra.Command) error { if cmd.Flags().Changed("from-plan") { return nil diff --git a/cmd/axis/task.go b/cmd/axis/task.go index 8be4e937..92d8112a 100644 --- a/cmd/axis/task.go +++ b/cmd/axis/task.go @@ -791,7 +791,7 @@ func buildContextBlock(snap *models.ClusterSnapshot, reqs models.TaskRequirement Be precise. Use real node names and tools above. Placement is advisory: execute via `+"`axis task run --script/--exec`"+` (guarded, reserved, confirmed) -or model lifecycle via `+"`axis model start --node --weights --port

`"+`.`, +or model lifecycle via `+"`axis model start `"+` (picks the node that already has the runtime).`, sourceOrLive(source), best.Name, ramSummary, pressure, gpuLine(gpuSummary), contextHint(reqs), toolsList(best), clusterSummaryLine(snap), task, extraLines) diff --git a/cmd/axis/task_context_enrichment_test.go b/cmd/axis/task_context_enrichment_test.go index cbb49316..4c469789 100644 --- a/cmd/axis/task_context_enrichment_test.go +++ b/cmd/axis/task_context_enrichment_test.go @@ -126,7 +126,7 @@ func TestBuildContextBlockNextActionNamesRealPaths(t *testing.T) { out := buildContextBlock(snap, models.TaskRequirements{}, "run task", "live", nil, nil) for _, want := range []string{ "axis task run --script/--exec", - "axis model start --node", + "axis model start ", "axis mcp serve", } { if !strings.Contains(out, want) { diff --git a/cmd/axis/testdata/task_context_turboquant.golden b/cmd/axis/testdata/task_context_turboquant.golden index 2fd61ac7..29f0615f 100644 --- a/cmd/axis/testdata/task_context_turboquant.golden +++ b/cmd/axis/testdata/task_context_turboquant.golden @@ -12,4 +12,4 @@ AXIS CLUSTER CONTEXT (paste as system prompt): Be precise. Use real node names and tools above. Placement is advisory: execute via `axis task run --script/--exec` (guarded, reserved, confirmed) -or model lifecycle via `axis model start --node --weights --port

`. \ No newline at end of file +or model lifecycle via `axis model start ` (picks the node that already has the runtime). \ No newline at end of file diff --git a/docs/current-state.md b/docs/current-state.md index c11d38ed..d959b7bd 100644 --- a/docs/current-state.md +++ b/docs/current-state.md @@ -108,7 +108,7 @@ Top-level commands currently registered in the binary: | `axis agent` | Agentic tool-calling assistant | Cluster tools + Layer-4 guarded `run_shell` / `run_on_node` / `axis_run_task`; injects nearest `AGENTS.md` into the system prompt when present; `--auto-approve` for safe commands; `--system` appends to system prompt | | `axis llm` | Removed | Prints `use: axis ai route` | | `axis model list\|inspect` | Inspect resident model instances | Daemon cache by default; `--live` explicitly performs a fresh cluster collection; text, JSON, and YAML output | -| `axis model start\|stop` | Manage llama-server, Ollama, and MLX | llama-server start requires a node, a port, and a weight path on an observed local volume; MLX start requires a node, a port, and a local model directory; Ollama load and unload require a node and a model name | +| `axis model start\|stop` | Manage llama-server, Ollama, and MLX | `axis model start ` picks one complete node that already has the runtime and prints that choice. Explicit llama-server start still requires a node, a port, and a weight path on an observed local volume. MLX start still requires a node, a port, and a local model directory. A pinned Ollama load or unload still requires a node and a model name | | `axis cluster` | Fleet snapshot | `status` (cache-first for 5 minutes; `--live` sweeps), `summary` | | `axis node` | This machine | `facts` (localhost). Root `axis facts` still works | diff --git a/internal/chat/system.go b/internal/chat/system.go index 491f281c..b417b156 100644 --- a/internal/chat/system.go +++ b/internal/chat/system.go @@ -59,7 +59,7 @@ func BuildSystemPrompt(cluster *ClusterSummaryForPrompt, extra string) string { b.WriteString("- You have first-class tools: axis_status, axis_facts, axis_place, axis_summary, axis_reservations. Prefer them over guessing — they read the live fact plane.\n") b.WriteString("- For placement questions, call axis_place with the task description; report the chosen node and reasoning verbatim. Placement is advisory.\n") b.WriteString("- Mutating actions (run_shell, run_on_node, axis_run_task, write_file, edit_file) go through safety checks and operator confirmation. Never assume approval.\n") - b.WriteString("- For models: axis model list shows what is resident where; axis model start/stop manage llama-server instances; axis model query prompts a resident model directly.\n") + b.WriteString("- For models: axis model list shows what is resident where; axis model start picks the node that already has the runtime; explicit start/stop flags pin llama-server, Ollama, or MLX; axis model query prompts a resident model directly.\n") b.WriteString("- CLI equivalents the user can run: axis facts, axis status, axis task place, axis task context, axis task run, axis doctor.\n") if cluster != nil { diff --git a/internal/modellife/ollama.go b/internal/modellife/ollama.go index 386fbd43..0d30765b 100644 --- a/internal/modellife/ollama.go +++ b/internal/modellife/ollama.go @@ -4,11 +4,15 @@ import ( "encoding/json" "fmt" "strings" + + "github.com/toasterbook88/axis/internal/models" ) // OllamaLoadScript preloads a model on the Ollama server that is already -// listening on 127.0.0.1:11434. It does not exec ollama serve. The last -// command's stdout is GET /api/ps, which is the load fact. +// listening on 127.0.0.1:11434. It does not exec ollama serve. The POST body +// omits prompt. Current Ollama schedules the runner and, when prompt is empty, +// returns done_reason load without calling Completion. The script accepts that +// body only, then prints GET /api/ps, which is the load fact. func OllamaLoadScript(model, keepAlive string, numCtx *int) (string, error) { model = strings.TrimSpace(model) if model == "" { @@ -36,7 +40,20 @@ func OllamaLoadScript(model, keepAlive string, numCtx *int) (string, error) { } probe := "curl -fsS --max-time 5 http://127.0.0.1:11434/api/ps" post := "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate -H 'Content-Type: application/json' -d " + shellSingleQuote(string(raw)) - return probe + " >/dev/null && " + post + " >/dev/null && " + probe, nil + loadOnly := `grep -q '"done_reason":"load"'` + return probe + " >/dev/null && " + post + " | " + loadOnly + " && " + probe, nil +} + +// OllamaServerReady reports whether facts show a complete node whose Ollama +// server is already listening. Placement and unload refuse before SSH otherwise. +func OllamaServerReady(node models.NodeFacts) error { + if node.Status != models.StatusComplete { + return fmt.Errorf("node %s status is %q; refusing ollama ssh", node.Name, node.Status) + } + if node.Ollama == nil || !node.Ollama.Listening { + return fmt.Errorf("node %s has no listening ollama server; refusing ollama ssh", node.Name) + } + return nil } // OllamaUnloadScript asks the existing server to drop a model, then prints diff --git a/internal/modellife/ollama_test.go b/internal/modellife/ollama_test.go index 5d7149c5..aa44d2a2 100644 --- a/internal/modellife/ollama_test.go +++ b/internal/modellife/ollama_test.go @@ -4,6 +4,8 @@ import ( "encoding/json" "strings" "testing" + + "github.com/toasterbook88/axis/internal/models" ) func TestOllamaLoadScriptIsLoopbackGenerateWithoutPrompt(t *testing.T) { @@ -17,6 +19,9 @@ func TestOllamaLoadScriptIsLoopbackGenerateWithoutPrompt(t *testing.T) { if !strings.Contains(script, "curl -fsS --max-time 30 -X POST http://127.0.0.1:11434/api/generate") { t.Fatalf("post missing:\n%s", script) } + if !strings.Contains(script, `grep -q '"done_reason":"load"'`) { + t.Fatalf("load must accept only the Ollama load response:\n%s", script) + } if strings.Contains(script, "/v1/models") || strings.Contains(script, "ollama serve") || strings.Contains(script, "OLLAMA_HOST") || strings.Contains(script, "OLLAMA_KEEP_ALIVE") { t.Fatalf("script widens or starts ollama:\n%s", script) } @@ -56,6 +61,20 @@ func TestOllamaLoadScriptAddsOnlyKeepAliveAndNumCtx(t *testing.T) { } } +func TestOllamaServerReadyRefusesIncompleteOrSilentNodes(t *testing.T) { + err := OllamaServerReady(models.NodeFacts{Name: "storage", Status: models.StatusPartial, Ollama: &models.OllamaInfo{Listening: true}}) + if err == nil || !strings.Contains(err.Error(), "refusing ollama ssh") { + t.Fatalf("partial err=%v", err) + } + err = OllamaServerReady(models.NodeFacts{Name: "storage", Status: models.StatusComplete}) + if err == nil || !strings.Contains(err.Error(), "no listening ollama") { + t.Fatalf("silent err=%v", err) + } + if err := OllamaServerReady(models.NodeFacts{Name: "storage", Status: models.StatusComplete, Ollama: &models.OllamaInfo{Listening: true}}); err != nil { + t.Fatal(err) + } +} + func TestOllamaUnloadScriptKeepsAliveZeroAndDoesNotKill(t *testing.T) { script, err := OllamaUnloadScript("mistral") if err != nil { diff --git a/internal/modellife/serve_pick.go b/internal/modellife/serve_pick.go new file mode 100644 index 00000000..043d813b --- /dev/null +++ b/internal/modellife/serve_pick.go @@ -0,0 +1,124 @@ +package modellife + +import ( + "fmt" + "sort" + "strings" + + "github.com/toasterbook88/axis/internal/models" +) + +// ServingPick is the one complete node whose facts already show a runtime for +// the named model. Already means the model is resident and must not be started +// again. Otherwise the runtime is an Ollama server that is already listening +// and lists the model. +type ServingPick struct { + Node string + Runtime string + Model string + Already bool + warmth float64 +} + +// Sentence is the operator line for a pick. +func (p ServingPick) Sentence() string { + if p.Already { + return fmt.Sprintf("picked %s: %s already serving %s", p.Node, p.Runtime, p.Model) + } + return fmt.Sprintf("picked %s: %s already listening", p.Node, p.Runtime) +} + +// PickServingNode chooses one complete node that already has a runtime for model. +// A resident model outranks an Ollama server that only lists it. Equal resident +// warmth then breaks by node name. A node that only has a llama-server binary, +// or an Ollama server that does not list the model, is not a candidate. +func PickServingNode(nodes []models.NodeFacts, model string) (ServingPick, error) { + model = strings.TrimSpace(model) + if model == "" { + return ServingPick{}, fmt.Errorf("model name is required") + } + var resident []ServingPick + var listening []ServingPick + for _, node := range nodes { + if node.Status != models.StatusComplete || strings.TrimSpace(node.Name) == "" { + continue + } + if pick, ok := residentServingPick(node, model); ok { + resident = append(resident, pick) + continue + } + if pick, ok := ollamaListeningPick(node, model); ok { + listening = append(listening, pick) + } + } + if len(resident) > 0 { + sort.SliceStable(resident, func(i, j int) bool { + if resident[i].warmth != resident[j].warmth { + return resident[i].warmth > resident[j].warmth + } + return resident[i].Node < resident[j].Node + }) + return resident[0], nil + } + if len(listening) > 0 { + sort.SliceStable(listening, func(i, j int) bool { + return listening[i].Node < listening[j].Node + }) + return listening[0], nil + } + return ServingPick{}, fmt.Errorf("no complete node already has a runtime for %q", model) +} + +func residentServingPick(node models.NodeFacts, model string) (ServingPick, bool) { + var found ServingPick + ok := false + for _, res := range node.ResidentModels { + if !knownServingRuntime(res.Runtime) || !servingNameMatches(res.Runtime, res.Name, model) { + continue + } + if ok && res.WarmthScore <= found.warmth { + continue + } + found = ServingPick{ + Node: node.Name, + Runtime: res.Runtime, + Model: model, + Already: true, + warmth: res.WarmthScore, + } + ok = true + } + return found, ok +} + +func ollamaListeningPick(node models.NodeFacts, model string) (ServingPick, bool) { + if node.Ollama == nil || !node.Ollama.Listening { + return ServingPick{}, false + } + for _, listed := range node.Ollama.Models { + if ollamaNameMatches(listed, model) { + return ServingPick{ + Node: node.Name, + Runtime: models.EngineOllama, + Model: model, + }, true + } + } + return ServingPick{}, false +} + +func knownServingRuntime(runtime string) bool { + switch runtime { + case models.EngineOllama, models.EngineLlamaCpp, models.EngineMLX, "apple-foundation-models": + return true + default: + return false + } +} + +func servingNameMatches(runtime, listed, requested string) bool { + if runtime == models.EngineOllama { + return ollamaNameMatches(listed, requested) + } + return strings.EqualFold(strings.TrimSpace(listed), strings.TrimSpace(requested)) +} diff --git a/internal/modellife/serve_pick_test.go b/internal/modellife/serve_pick_test.go new file mode 100644 index 00000000..7e4c6b9b --- /dev/null +++ b/internal/modellife/serve_pick_test.go @@ -0,0 +1,104 @@ +package modellife + +import ( + "strings" + "testing" + + "github.com/toasterbook88/axis/internal/models" +) + +func TestPickServingNodePrefersResidentOverListening(t *testing.T) { + nodes := []models.NodeFacts{ + { + Name: "library", + Status: models.StatusComplete, + Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral"}}, + }, + { + Name: "warm", + Status: models.StatusComplete, + ResidentModels: []models.ResidentModel{{ + Name: "mistral:latest", Runtime: "ollama", WarmthScore: 0.2, + }}, + }, + } + got, err := PickServingNode(nodes, "mistral") + if err != nil { + t.Fatal(err) + } + if got.Node != "warm" || !got.Already || got.Runtime != "ollama" { + t.Fatalf("pick=%#v", got) + } + if got.Sentence() != "picked warm: ollama already serving mistral" { + t.Fatalf("sentence=%q", got.Sentence()) + } +} + +func TestPickServingNodeBreaksResidentTiesByWarmthThenName(t *testing.T) { + nodes := []models.NodeFacts{ + {Name: "b", Status: models.StatusComplete, ResidentModels: []models.ResidentModel{{Name: "mistral", Runtime: "llama.cpp", WarmthScore: 0.9}}}, + {Name: "a", Status: models.StatusComplete, ResidentModels: []models.ResidentModel{{Name: "mistral", Runtime: "llama.cpp", WarmthScore: 0.1}}}, + {Name: "c", Status: models.StatusComplete, ResidentModels: []models.ResidentModel{{Name: "mistral", Runtime: "mlx", WarmthScore: 0.9}}}, + } + got, err := PickServingNode(nodes, "Mistral") + if err != nil { + t.Fatal(err) + } + if got.Node != "b" || got.Runtime != "llama.cpp" || !got.Already { + t.Fatalf("warmth tie pick=%#v", got) + } + nodes[0].ResidentModels[0].WarmthScore = 0.4 + got, err = PickServingNode(nodes, "mistral") + if err != nil { + t.Fatal(err) + } + if got.Node != "c" || got.Runtime != "mlx" { + t.Fatalf("equal warmth pick=%#v", got) + } +} + +func TestPickServingNodeUsesListeningOllamaLibrary(t *testing.T) { + nodes := []models.NodeFacts{ + {Name: "partial", Status: models.StatusPartial, Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral"}}}, + {Name: "zeta", Status: models.StatusComplete, Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral:latest"}}}, + {Name: "alpha", Status: models.StatusComplete, Ollama: &models.OllamaInfo{Listening: true, Models: []string{"other"}}}, + {Name: "beta", Status: models.StatusComplete, Tools: []models.ToolInfo{{Name: "llama-server", Path: "/usr/bin/llama-server"}}}, + {Name: "mid", Status: models.StatusComplete, Ollama: &models.OllamaInfo{Installed: true, Running: true, Models: []string{"mistral"}}}, + } + got, err := PickServingNode(nodes, "mistral") + if err != nil { + t.Fatal(err) + } + if got.Node != "zeta" || got.Already || got.Sentence() != "picked zeta: ollama already listening" { + t.Fatalf("pick=%#v sentence=%q", got, got.Sentence()) + } +} + +func TestPickServingNodeRefusesWhenNoRuntimeIsAlreadyThere(t *testing.T) { + _, err := PickServingNode(nil, "mistral") + if err == nil || !strings.Contains(err.Error(), `no complete node already has a runtime for "mistral"`) { + t.Fatalf("err=%v", err) + } + _, err = PickServingNode([]models.NodeFacts{{Name: "storage", Status: models.StatusComplete}}, " ") + if err == nil || !strings.Contains(err.Error(), "model name is required") { + t.Fatalf("blank err=%v", err) + } +} + +func TestPickServingNodeIgnoresUnknownResidentRuntime(t *testing.T) { + nodes := []models.NodeFacts{{ + Name: "storage", + Status: models.StatusComplete, + ResidentModels: []models.ResidentModel{{ + Name: "mistral", Runtime: "custom-engine", + }}, + Ollama: &models.OllamaInfo{Listening: true, Models: []string{"mistral"}}, + }} + got, err := PickServingNode(nodes, "mistral") + if err != nil { + t.Fatal(err) + } + if got.Already || got.Runtime != "ollama" { + t.Fatalf("pick=%#v", got) + } +}