Skip to content

Commit e79ce2a

Browse files
committed
feat(mcp): register scheduling tools
Assisted-by: Hephaestus:openai/gpt-5.5 Signed-off-by: Owen Adirah <owenadira@gmail.com>
1 parent 757922d commit e79ce2a

4 files changed

Lines changed: 161 additions & 0 deletions

File tree

pkg/mcp/localaitools/fakes_test.go

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,10 @@ type fakeClient struct {
4242
upgradeBackend func(string) (string, error)
4343
systemInfo func() (*SystemInfo, error)
4444
listNodes func() ([]Node, error)
45+
listScheduling func() ([]ModelSchedulingConfig, error)
46+
getScheduling func(string) (*ModelSchedulingConfig, error)
47+
setScheduling func(SetSchedulingRequest) (*ModelSchedulingConfig, error)
48+
deleteScheduling func(string) error
4549
setNodeVRAMBudget func(string, string) error
4650
vramEstimate func(VRAMEstimateRequest) (*vram.EstimateResult, error)
4751
toggleModelState func(string, modeladmin.Action) error
@@ -230,6 +234,38 @@ func (f *fakeClient) ListNodes(_ context.Context) ([]Node, error) {
230234
return nil, nil
231235
}
232236

237+
func (f *fakeClient) ListScheduling(_ context.Context) ([]ModelSchedulingConfig, error) {
238+
f.record("ListScheduling", nil)
239+
if f.listScheduling != nil {
240+
return f.listScheduling()
241+
}
242+
return []ModelSchedulingConfig{}, nil
243+
}
244+
245+
func (f *fakeClient) GetScheduling(_ context.Context, modelName string) (*ModelSchedulingConfig, error) {
246+
f.record("GetScheduling", modelName)
247+
if f.getScheduling != nil {
248+
return f.getScheduling(modelName)
249+
}
250+
return &ModelSchedulingConfig{ModelName: modelName}, nil
251+
}
252+
253+
func (f *fakeClient) SetScheduling(_ context.Context, req SetSchedulingRequest) (*ModelSchedulingConfig, error) {
254+
f.record("SetScheduling", req)
255+
if f.setScheduling != nil {
256+
return f.setScheduling(req)
257+
}
258+
return &ModelSchedulingConfig{ModelName: req.ModelName, MinReplicas: req.MinReplicas, MaxReplicas: req.MaxReplicas, SpreadAll: req.SpreadAll}, nil
259+
}
260+
261+
func (f *fakeClient) DeleteScheduling(_ context.Context, modelName string) error {
262+
f.record("DeleteScheduling", modelName)
263+
if f.deleteScheduling != nil {
264+
return f.deleteScheduling(modelName)
265+
}
266+
return nil
267+
}
268+
233269
func (f *fakeClient) SetNodeVRAMBudget(_ context.Context, nodeID, budget string) error {
234270
f.record("SetNodeVRAMBudget", []any{nodeID, budget})
235271
if f.setNodeVRAMBudget != nil {

pkg/mcp/localaitools/server.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@ func NewServer(client LocalAIClient, opts Options) *mcp.Server {
4747
registerBackendTools(srv, client, opts)
4848
registerConfigTools(srv, client, opts)
4949
registerSystemTools(srv, client, opts)
50+
registerSchedulingTools(srv, client, opts)
5051
registerStateTools(srv, client, opts)
5152
registerBrandingTools(srv, client, opts)
5253
registerVoiceProfileTools(srv, client, opts)

pkg/mcp/localaitools/server_test.go

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -92,6 +92,7 @@ var expectedFullCatalog = sortedStrings(
9292
ToolListInstalledModels,
9393
ToolListKnownBackends,
9494
ToolListNodes,
95+
ToolListScheduling,
9596
ToolListVoiceProfiles,
9697
ToolLoadModel,
9798
ToolReloadModels,
@@ -105,6 +106,9 @@ var expectedFullCatalog = sortedStrings(
105106
ToolCreateVoiceProfile,
106107
ToolDeleteVoiceProfile,
107108
ToolSetNodeVRAMBudget,
109+
ToolSetScheduling,
110+
ToolDeleteScheduling,
111+
ToolGetScheduling,
108112
)
109113

110114
// expectedReadOnlyCatalog is the tool set when DisableMutating=true. Sorted.
@@ -123,6 +127,8 @@ var expectedReadOnlyCatalog = sortedStrings(
123127
ToolListInstalledModels,
124128
ToolListKnownBackends,
125129
ToolListNodes,
130+
ToolListScheduling,
131+
ToolGetScheduling,
126132
ToolListVoiceProfiles,
127133
ToolSystemInfo,
128134
ToolVRAMEstimate,
@@ -165,6 +171,10 @@ var _ = Describe("Tool dispatch", func() {
165171
{ToolListKnownBackends, struct{}{}, "ListKnownBackends"},
166172
{ToolSystemInfo, struct{}{}, "SystemInfo"},
167173
{ToolListNodes, struct{}{}, "ListNodes"},
174+
{ToolListScheduling, struct{}{}, "ListScheduling"},
175+
{ToolGetScheduling, DeleteSchedulingRequest{ModelName: "qwen"}, "GetScheduling"},
176+
{ToolSetScheduling, SetSchedulingRequest{ModelName: "qwen", MinReplicas: 1, MaxReplicas: 2}, "SetScheduling"},
177+
{ToolDeleteScheduling, DeleteSchedulingRequest{ModelName: "qwen"}, "DeleteScheduling"},
168178
{ToolListVoiceProfiles, struct{}{}, "ListVoiceProfiles"},
169179
{ToolInstallModel, InstallModelRequest{ModelName: "test/foo"}, "InstallModel"},
170180
{ToolImportModelURI, ImportModelURIRequest{URI: "Qwen/Qwen3-4B-GGUF"}, "ImportModelURI"},
@@ -219,6 +229,40 @@ var _ = Describe("Tool error surfacing", func() {
219229
})
220230
})
221231

232+
var _ = Describe("Scheduling tool behavior", func() {
233+
It("round-trips list/get/set/delete over the MCP surface", func() {
234+
fc := &fakeClient{
235+
listScheduling: func() ([]ModelSchedulingConfig, error) {
236+
return []ModelSchedulingConfig{{ModelName: "qwen", MinReplicas: 1, MaxReplicas: 2}}, nil
237+
},
238+
getScheduling: func(modelName string) (*ModelSchedulingConfig, error) {
239+
return &ModelSchedulingConfig{ModelName: modelName, SpreadAll: true}, nil
240+
},
241+
setScheduling: func(req SetSchedulingRequest) (*ModelSchedulingConfig, error) {
242+
return &ModelSchedulingConfig{ModelName: req.ModelName, MinReplicas: req.MinReplicas, MaxReplicas: req.MaxReplicas}, nil
243+
},
244+
}
245+
ctx, sess, done := connectInMemory(fc, Options{})
246+
DeferCleanup(done)
247+
248+
listRes := callTool(ctx, sess, ToolListScheduling, struct{}{})
249+
Expect(listRes.IsError).To(BeFalse(), resultText(listRes))
250+
Expect(resultText(listRes)).To(ContainSubstring(`"model_name": "qwen"`))
251+
252+
getRes := callTool(ctx, sess, ToolGetScheduling, DeleteSchedulingRequest{ModelName: "qwen"})
253+
Expect(getRes.IsError).To(BeFalse(), resultText(getRes))
254+
Expect(resultText(getRes)).To(ContainSubstring(`"spread_all": true`))
255+
256+
setRes := callTool(ctx, sess, ToolSetScheduling, SetSchedulingRequest{ModelName: "qwen", MinReplicas: 1, MaxReplicas: 2})
257+
Expect(setRes.IsError).To(BeFalse(), resultText(setRes))
258+
Expect(resultText(setRes)).To(ContainSubstring(`"max_replicas": 2`))
259+
260+
deleteRes := callTool(ctx, sess, ToolDeleteScheduling, DeleteSchedulingRequest{ModelName: "qwen"})
261+
Expect(deleteRes.IsError).To(BeFalse(), resultText(deleteRes))
262+
Expect(resultText(deleteRes)).To(ContainSubstring(`"deleted": "qwen"`))
263+
})
264+
})
265+
222266
var _ = Describe("Argument validation", func() {
223267
type validationCase struct {
224268
desc string
@@ -235,6 +279,9 @@ var _ = Describe("Argument validation", func() {
235279
{"toggle_model_state rejects unknown action", ToolToggleModelState, map[string]any{"name": "foo", "action": "noop"}, "action must be one of"},
236280
{"edit_model_config rejects empty patch", ToolEditModelConfig, map[string]any{"name": "foo", "patch": map[string]any{}}, "patch is required"},
237281
{"create_voice_profile requires consent", ToolCreateVoiceProfile, CreateVoiceProfileRequest{Name: "Voice", Transcript: "words", AudioBase64: "UklGRg=="}, "consent_confirmed must be true"},
282+
{"set_scheduling requires model_name", ToolSetScheduling, SetSchedulingRequest{}, "model_name is required"},
283+
{"set_scheduling rejects invalid replica range", ToolSetScheduling, SetSchedulingRequest{ModelName: "qwen", MinReplicas: 3, MaxReplicas: 1}, "min_replicas must be <= max_replicas"},
284+
{"delete_scheduling requires model_name", ToolDeleteScheduling, DeleteSchedulingRequest{}, "model_name is required"},
238285
}
239286

240287
for _, c := range cases {
Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,77 @@
1+
package localaitools
2+
3+
import (
4+
"context"
5+
6+
"github.com/modelcontextprotocol/go-sdk/mcp"
7+
)
8+
9+
func registerSchedulingTools(s *mcp.Server, client LocalAIClient, opts Options) {
10+
mcp.AddTool(s, &mcp.Tool{
11+
Name: ToolListScheduling,
12+
Description: "List distributed per-model scheduling configs (only meaningful in distributed mode).",
13+
}, func(ctx context.Context, _ *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, any, error) {
14+
configs, err := client.ListScheduling(ctx)
15+
if err != nil {
16+
return errorResult(err), nil, nil
17+
}
18+
return jsonResult(configs), nil, nil
19+
})
20+
21+
mcp.AddTool(s, &mcp.Tool{
22+
Name: ToolGetScheduling,
23+
Description: "Get the distributed scheduling config for one model, or null when none is configured.",
24+
}, func(ctx context.Context, _ *mcp.CallToolRequest, args DeleteSchedulingRequest) (*mcp.CallToolResult, any, error) {
25+
if args.ModelName == "" {
26+
return errorResultf("model_name is required"), nil, nil
27+
}
28+
config, err := client.GetScheduling(ctx, args.ModelName)
29+
if err != nil {
30+
return errorResult(err), nil, nil
31+
}
32+
return jsonResult(config), nil, nil
33+
})
34+
35+
if opts.DisableMutating {
36+
return
37+
}
38+
39+
mcp.AddTool(s, &mcp.Tool{
40+
Name: ToolSetScheduling,
41+
Description: "Create or update a distributed per-model scheduling config. Requires user confirmation per safety rule 1.",
42+
}, func(ctx context.Context, _ *mcp.CallToolRequest, args SetSchedulingRequest) (*mcp.CallToolResult, any, error) {
43+
if args.ModelName == "" {
44+
return errorResultf("model_name is required"), nil, nil
45+
}
46+
if args.SpreadAll && (args.MinReplicas != 0 || args.MaxReplicas != 0) {
47+
return errorResultf("spread_all and min_replicas/max_replicas are mutually exclusive"), nil, nil
48+
}
49+
if args.MinReplicas < 0 {
50+
return errorResultf("min_replicas must be >= 0"), nil, nil
51+
}
52+
if args.MaxReplicas < 0 {
53+
return errorResultf("max_replicas must be >= 0"), nil, nil
54+
}
55+
if args.MaxReplicas > 0 && args.MinReplicas > args.MaxReplicas {
56+
return errorResultf("min_replicas must be <= max_replicas"), nil, nil
57+
}
58+
config, err := client.SetScheduling(ctx, args)
59+
if err != nil {
60+
return errorResult(err), nil, nil
61+
}
62+
return jsonResult(config), nil, nil
63+
})
64+
65+
mcp.AddTool(s, &mcp.Tool{
66+
Name: ToolDeleteScheduling,
67+
Description: "Delete a distributed per-model scheduling config. Requires user confirmation per safety rule 1.",
68+
}, func(ctx context.Context, _ *mcp.CallToolRequest, args DeleteSchedulingRequest) (*mcp.CallToolResult, any, error) {
69+
if args.ModelName == "" {
70+
return errorResultf("model_name is required"), nil, nil
71+
}
72+
if err := client.DeleteScheduling(ctx, args.ModelName); err != nil {
73+
return errorResult(err), nil, nil
74+
}
75+
return jsonResult(map[string]string{"deleted": args.ModelName}), nil, nil
76+
})
77+
}

0 commit comments

Comments
 (0)