From eaa857933305444e952d174edeb994f706dc89e8 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Mon, 24 Aug 2026 18:22:59 +0800 Subject: [PATCH 01/12] feat: expand provider and model management Signed-off-by: masteryyh --- .../pkg/agentloop/testhelper_test.go | 4 +- .../agenty-core/pkg/application/provider.go | 242 +++++++-- .../pkg/application/provider_test.go | 267 ++++++++-- .../pkg/application/testhelper_test.go | 28 +- .../pkg/domain/catalog/available_model.go | 22 + .../agenty-core/pkg/domain/catalog/model.go | 31 +- .../pkg/domain/catalog/model_test.go | 51 +- .../pkg/domain/catalog/provider.go | 32 +- .../pkg/domain/catalog/provider_test.go | 8 +- .../pkg/domain/shared/reasoning.go | 30 ++ .../pkg/domain/shared/reasoning_test.go | 34 +- .../pkg/infra/catalogdata/catalog.go | 114 ++++ .../pkg/infra/catalogdata/catalog_test.go | 57 ++ .../pkg/infra/catalogdata/providers.json | 279 ++++++++++ .../pkg/infra/initialize/initialize.go | 8 +- .../pkg/infra/initialize/initialize_test.go | 31 +- .../agenty-core/pkg/infra/llm/anthropic.go | 5 +- packages/agenty-core/pkg/infra/llm/convert.go | 28 +- .../agenty-core/pkg/infra/llm/convert_test.go | 146 +++--- packages/agenty-core/pkg/infra/llm/factory.go | 20 +- packages/agenty-core/pkg/infra/llm/google.go | 26 +- .../agenty-core/pkg/infra/llm/openai_chat.go | 5 +- .../pkg/infra/llm/openai_responses.go | 65 ++- .../pkg/infra/modelcatalog/lister.go | 487 ++++++++++++++++++ .../pkg/infra/modelcatalog/lister_test.go | 229 ++++++++ .../pkg/infra/rpc/adapter/provider.go | 23 +- .../agenty-core/pkg/infra/storage/catalog.go | 267 +++++++++- .../pkg/infra/storage/catalog_test.go | 170 +++++- .../test/e2e/agenty_client_test.go | 9 + .../agenty-core/test/e2e/contracts_test.go | 89 ++-- packages/agenty-core/test/e2e/journey_test.go | 35 +- .../test/e2e/provider_fixture_test.go | 4 +- 32 files changed, 2467 insertions(+), 379 deletions(-) create mode 100644 packages/agenty-core/pkg/domain/catalog/available_model.go create mode 100644 packages/agenty-core/pkg/infra/catalogdata/catalog.go create mode 100644 packages/agenty-core/pkg/infra/catalogdata/catalog_test.go create mode 100644 packages/agenty-core/pkg/infra/catalogdata/providers.json create mode 100644 packages/agenty-core/pkg/infra/modelcatalog/lister.go create mode 100644 packages/agenty-core/pkg/infra/modelcatalog/lister_test.go diff --git a/packages/agenty-core/pkg/agentloop/testhelper_test.go b/packages/agenty-core/pkg/agentloop/testhelper_test.go index b108755..5006f06 100644 --- a/packages/agenty-core/pkg/agentloop/testhelper_test.go +++ b/packages/agenty-core/pkg/agentloop/testhelper_test.go @@ -75,9 +75,7 @@ func cloneProvider(provider *catalog.Provider) *catalog.Provider { copy := *provider copy.Models = slices.Clone(provider.Models) for index := range copy.Models { - copy.Models[index].ReasoningEffortMapping = maps.Clone( - copy.Models[index].ReasoningEffortMapping, - ) + copy.Models[index].ReasoningEfforts = slices.Clone(copy.Models[index].ReasoningEfforts) } copy.Metadata = maps.Clone(provider.Metadata) return © diff --git a/packages/agenty-core/pkg/application/provider.go b/packages/agenty-core/pkg/application/provider.go index c17b6ae..ea00d49 100644 --- a/packages/agenty-core/pkg/application/provider.go +++ b/packages/agenty-core/pkg/application/provider.go @@ -3,10 +3,14 @@ package application import ( "context" "errors" + "log/slog" + "strings" + "sync" "time" "github.com/masteryyh/agenty-core/pkg/domain/catalog" "github.com/masteryyh/agenty-core/pkg/domain/shared" + "github.com/masteryyh/agenty-core/pkg/infra/modelcatalog" "github.com/masteryyh/agenty-core/pkg/infra/storage" ) @@ -19,6 +23,8 @@ type providerRepository interface { List(ctx context.Context) ([]*catalog.Provider, error) Save(ctx context.Context, provider *catalog.Provider) error Delete(ctx context.Context, code shared.Code) error + NeedsModelDiscovery(ctx context.Context, code shared.Code) (bool, error) + ReplaceModels(ctx context.Context, code shared.Code, models []catalog.Model, expiresAt time.Time) error } func NewProviderService(repo providerRepository) *ProviderService { @@ -26,11 +32,12 @@ func NewProviderService(repo providerRepository) *ProviderService { } type ProviderInput struct { - Name string `json:"name"` - Type catalog.APIType `json:"type"` - BaseURL string `json:"baseUrl,omitempty"` - APIKey string `json:"apiKey,omitempty"` - Metadata shared.Metadata `json:"metadata,omitempty"` + Name string `json:"name"` + Type catalog.APIType `json:"type"` + BaseURL string `json:"baseUrl,omitempty"` + APIKey string `json:"apiKey,omitempty"` + FreeFormTool bool `json:"freeFormTool,omitempty"` + Metadata shared.Metadata `json:"metadata,omitempty"` } func (s *ProviderService) Create(ctx context.Context, code string, in ProviderInput) (*catalog.Provider, error) { @@ -56,6 +63,7 @@ func (s *ProviderService) Create(ctx context.Context, code string, in ProviderIn p.BaseURL = in.BaseURL p.APIKey = in.APIKey + p.FreeFormTool = supportsFreeFormTool(in.Type, in.FreeFormTool) p.Metadata = in.Metadata if err := s.repo.Save(ctx, p); err != nil { @@ -80,7 +88,7 @@ func (s *ProviderService) Get(ctx context.Context, code string) (*catalog.Provid return p, nil } -func (s *ProviderService) List(ctx context.Context) ([]*catalog.Provider, error) { +func (s *ProviderService) List(ctx context.Context, targetCodes ...string) ([]*catalog.Provider, error) { providers, err := s.repo.List(ctx) if err != nil { return nil, Internal("failed to list providers: " + err.Error()) @@ -88,15 +96,132 @@ func (s *ProviderService) List(ctx context.Context) ([]*catalog.Provider, error) if providers == nil { providers = make([]*catalog.Provider, 0) } + targetCode := "" + if len(targetCodes) > 0 { + targetCode = strings.TrimSpace(targetCodes[0]) + } + + toDiscover := make([]shared.Code, 0) + for _, provider := range providers { + if provider == nil || (targetCode != "" && provider.Code.String() != targetCode) { + continue + } + if strings.TrimSpace(provider.APIKey) == "" { + continue + } + needsDiscovery, err := s.repo.NeedsModelDiscovery(ctx, provider.Code) + if err != nil { + return nil, Internal("failed to inspect model cache for provider " + provider.Code.String() + ": " + err.Error()) + } + if needsDiscovery { + toDiscover = append(toDiscover, provider.Code) + } + } + + var waitGroup sync.WaitGroup + for _, code := range toDiscover { + waitGroup.Add(1) + go func(providerCode shared.Code) { + defer waitGroup.Done() + if _, err := s.ListModels(ctx, providerCode.String()); err != nil { + slog.WarnContext(ctx, "failed to discover provider models", "providerCode", providerCode, "error", err) + } + }(code) + } + waitGroup.Wait() + if len(toDiscover) > 0 { + providers, err = s.repo.List(ctx) + if err != nil { + return nil, Internal("failed to list providers after model discovery: " + err.Error()) + } + } return providers, nil } +func (s *ProviderService) ListModels(ctx context.Context, code string) ([]catalog.AvailableModel, error) { + codeVal, err := shared.NewCode(code) + if err != nil { + return nil, Validation(err.Error()) + } + + provider, err := s.repo.Get(ctx, codeVal) + if err != nil { + if errors.Is(err, storage.ErrProviderNotFound) { + return nil, NotFound("provider " + code + " not found") + } + return nil, Internal("failed to get provider: " + err.Error()) + } + if provider == nil { + return nil, Internal("provider repository returned an empty provider") + } + if strings.TrimSpace(provider.APIKey) == "" { + return nil, Validation("provider " + code + " has no API key") + } + needsDiscovery, err := s.repo.NeedsModelDiscovery(ctx, codeVal) + if err != nil { + return nil, Internal("failed to inspect model cache for provider " + code + ": " + err.Error()) + } + if !needsDiscovery { + return availableModelsFromCatalog(provider.Models), nil + } + models, err := modelcatalog.List(ctx, *provider) + if err != nil { + return nil, Internal("failed to list models for provider " + code + ": " + err.Error()) + } + if models == nil { + models = make([]catalog.AvailableModel, 0) + } + if err := s.repo.ReplaceModels( + ctx, + codeVal, + catalogModelsFromAvailable(models), + time.Now().UTC().Add(catalog.ModelDiscoveryCacheTTL), + ); err != nil { + return nil, Internal("failed to cache models for provider " + code + ": " + err.Error()) + } + return models, nil +} + +func catalogModelsFromAvailable(models []catalog.AvailableModel) []catalog.Model { + now := time.Now().UTC() + result := make([]catalog.Model, 0, len(models)) + for _, available := range models { + result = append(result, catalog.Model{ + Code: available.Code, + Name: available.Name, + ContextWindow: available.ContextWindow, + MaxOutputTokens: available.MaxOutputTokens, + MultiModal: available.MultiModal, + ReasoningEfforts: available.ReasoningEfforts, + CreatedAt: now, + UpdatedAt: now, + }) + } + return result +} + +func availableModelsFromCatalog(models []catalog.Model) []catalog.AvailableModel { + result := make([]catalog.AvailableModel, 0, len(models)) + for _, model := range models { + result = append(result, catalog.AvailableModel{ + Code: model.Code, + Name: model.Name, + ContextWindow: model.ContextWindow, + MaxOutputTokens: model.MaxOutputTokens, + MultiModal: model.MultiModal, + ReasoningEfforts: model.ReasoningEfforts, + }) + } + return result +} + type ProviderUpdate struct { - Name *string `json:"name,omitempty"` - Type *catalog.APIType `json:"type,omitempty"` - BaseURL *string `json:"baseUrl,omitempty"` - APIKey *string `json:"apiKey,omitempty"` - Metadata *shared.Metadata `json:"metadata,omitempty"` + Name *string `json:"name,omitempty"` + Type *catalog.APIType `json:"type,omitempty"` + BaseURL *string `json:"baseUrl,omitempty"` + APIKey *string `json:"apiKey,omitempty"` + FreeFormTool *bool `json:"freeFormTool,omitempty"` + Metadata *shared.Metadata `json:"metadata,omitempty"` } func (s *ProviderService) Update(ctx context.Context, code string, upd ProviderUpdate) (*catalog.Provider, error) { @@ -112,6 +237,19 @@ func (s *ProviderService) Update(ctx context.Context, code string, upd ProviderU } return nil, Internal("failed to get provider: " + err.Error()) } + if p.Builtin { + if upd.Name != nil || upd.Type != nil || upd.BaseURL != nil || upd.FreeFormTool != nil || upd.Metadata != nil { + return nil, Validation("built-in provider metadata is read-only; only the API key can be changed") + } + if upd.APIKey == nil { + return p, nil + } + p.APIKey = *upd.APIKey + if err := s.repo.Save(ctx, p); err != nil { + return nil, Internal("failed to save provider API key: " + err.Error()) + } + return p, nil + } if upd.Name != nil { p.Name = *upd.Name @@ -128,6 +266,10 @@ func (s *ProviderService) Update(ctx context.Context, code string, upd ProviderU if upd.APIKey != nil { p.APIKey = *upd.APIKey } + if upd.FreeFormTool != nil { + p.FreeFormTool = *upd.FreeFormTool + } + p.FreeFormTool = supportsFreeFormTool(p.Type, p.FreeFormTool) if upd.Metadata != nil { p.Metadata = *upd.Metadata } @@ -139,12 +281,27 @@ func (s *ProviderService) Update(ctx context.Context, code string, upd ProviderU return p, nil } +func supportsFreeFormTool(apiType catalog.APIType, enabled bool) bool { + return apiType == catalog.APIOpenAI && enabled +} + func (s *ProviderService) Delete(ctx context.Context, code string) error { codeVal, err := shared.NewCode(code) if err != nil { return Validation(err.Error()) } + p, err := s.repo.Get(ctx, codeVal) + if err != nil { + if errors.Is(err, storage.ErrProviderNotFound) { + return NotFound("provider " + code + " not found") + } + return Internal("failed to get provider: " + err.Error()) + } + if p.Builtin { + return Validation("built-in provider is read-only; only the API key can be changed") + } + if err := s.repo.Delete(ctx, codeVal); err != nil { if errors.Is(err, storage.ErrProviderNotFound) { return NotFound("provider " + code + " not found") @@ -155,15 +312,13 @@ func (s *ProviderService) Delete(ctx context.Context, code string) error { } type ModelInput struct { - Name string `json:"name"` - ContextWindow int `json:"contextWindow,omitempty"` - // MaxOutputTokens is retained for wire compatibility with older clients. - // The core ignores it and persists the single global output limit. - MaxOutputTokens int64 `json:"maxOutputTokens"` - MultiModal bool `json:"multiModal,omitempty"` - Light bool `json:"light,omitempty"` - ReasoningEffortMapping map[string]shared.ReasoningEffort `json:"reasoningEffortMapping,omitempty"` - IsDefault bool `json:"isDefault,omitempty"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow,omitempty"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal,omitempty"` + Light bool `json:"light,omitempty"` + Reasoning *bool `json:"reasoning,omitempty"` + IsDefault bool `json:"isDefault,omitempty"` } func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode string, in ModelInput) (*catalog.Provider, error) { @@ -176,15 +331,6 @@ func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode if err != nil { return nil, Validation(err.Error()) } - for nativeEffort, agentyEffort := range in.ReasoningEffortMapping { - if nativeEffort == "" { - return nil, Validation("reasoning effort mapping contains an empty native effort") - } - if !agentyEffort.Valid() { - return nil, Validation("invalid agenty reasoning effort: " + string(agentyEffort)) - } - } - p, err := s.repo.Get(ctx, ps) if err != nil { if errors.Is(err, storage.ErrProviderNotFound) { @@ -192,19 +338,34 @@ func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode } return nil, Internal("failed to get provider: " + err.Error()) } + if p.Builtin { + return nil, Validation("built-in provider models are read-only") + } now := time.Now().UTC() + maxOutputTokens := in.MaxOutputTokens + if maxOutputTokens <= 0 { + maxOutputTokens = catalog.DefaultMaxOutputTokens + } + reasoning := true + if in.Reasoning != nil { + reasoning = *in.Reasoning + } + reasoningEfforts := make([]shared.ReasoningEffort, 0) + if reasoning { + reasoningEfforts = shared.StandardReasoningEfforts() + } p.AddModel(catalog.Model{ - Code: ms, - Name: in.Name, - ContextWindow: in.ContextWindow, - MaxOutputTokens: catalog.DefaultMaxOutputTokens, - MultiModal: in.MultiModal, - Light: in.Light, - ReasoningEffortMapping: in.ReasoningEffortMapping, - IsDefault: in.IsDefault, - CreatedAt: now, - UpdatedAt: now, + Code: ms, + Name: in.Name, + ContextWindow: in.ContextWindow, + MaxOutputTokens: maxOutputTokens, + MultiModal: in.MultiModal, + Light: in.Light, + ReasoningEfforts: reasoningEfforts, + IsDefault: in.IsDefault, + CreatedAt: now, + UpdatedAt: now, }) p.UpdatedAt = now @@ -232,6 +393,9 @@ func (s *ProviderService) RemoveModel(ctx context.Context, providerCode, modelCo } return nil, Internal("failed to get provider: " + err.Error()) } + if p.Builtin { + return nil, Validation("built-in provider models are read-only") + } if _, err := p.Model(ms); err != nil { if errors.Is(err, catalog.ErrModelNotFound) { diff --git a/packages/agenty-core/pkg/application/provider_test.go b/packages/agenty-core/pkg/application/provider_test.go index 812bb85..e4dc006 100644 --- a/packages/agenty-core/pkg/application/provider_test.go +++ b/packages/agenty-core/pkg/application/provider_test.go @@ -2,6 +2,10 @@ package application_test import ( "context" + "net/http" + "net/http/httptest" + "slices" + "sync/atomic" "testing" "github.com/masteryyh/agenty-core/pkg/application" @@ -14,10 +18,11 @@ func TestProviderCreateAndGet(t *testing.T) { ctx := context.Background() p, err := providerSvc.Create(ctx, "anthropic", application.ProviderInput{ - Name: "Anthropic", - Type: catalog.APIAnthropic, - BaseURL: "https://api.anthropic.com", - APIKey: "sk-ant-test", + Name: "Anthropic", + Type: catalog.APIAnthropic, + BaseURL: "https://api.anthropic.com", + APIKey: "sk-ant-test", + FreeFormTool: true, }) if err != nil { t.Fatalf("Create: %v", err) @@ -36,6 +41,35 @@ func TestProviderCreateAndGet(t *testing.T) { if got.Type != catalog.APIAnthropic { t.Errorf("type = %s", got.Type) } + if got.FreeFormTool { + t.Error("freeFormTool = true for Anthropic provider, want ignored") + } +} + +func TestProviderCreateIgnoresFreeFormToolForNonOpenAI(t *testing.T) { + for _, test := range []struct { + code string + name string + apiType catalog.APIType + }{ + {code: "anthropic", name: "Anthropic", apiType: catalog.APIAnthropic}, + {code: "google", name: "Google", apiType: catalog.APIGemini}, + } { + t.Run(test.name, func(t *testing.T) { + _, providerSvc, _ := newServices(t) + provider, err := providerSvc.Create(t.Context(), test.code, application.ProviderInput{ + Name: test.name, + Type: test.apiType, + FreeFormTool: true, + }) + if err != nil { + t.Fatalf("Create: %v", err) + } + if provider.FreeFormTool { + t.Fatalf("freeFormTool = true for %s provider, want ignored", test.apiType) + } + }) + } } func TestProviderCreateInvalidType(t *testing.T) { @@ -78,11 +112,131 @@ func TestProviderList(t *testing.T) { } } +func TestProviderListModels(t *testing.T) { + requestCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestCount++ + if r.URL.Path != "/models" { + t.Errorf("path = %q, want /models", r.URL.Path) + } + if got := r.Header.Get("Authorization"); got != "Bearer test-key" { + t.Errorf("authorization = %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-test"}]}`)) + })) + defer server.Close() + + repo := newProviderRepositoryFake() + provider, err := catalog.NewProvider("openai", "OpenAI", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + provider.BaseURL = server.URL + provider.APIKey = "test-key" + repo.providers[provider.Code] = provider + service := application.NewProviderService(repo) + + models, err := service.ListModels(t.Context(), "openai") + if err != nil { + t.Fatalf("ListModels: %v", err) + } + if len(models) != 1 || models[0].Code != "gpt-test" { + t.Fatalf("models = %#v", models) + } + if models[0].Name != "gpt-test" || models[0].ContextWindow != catalog.DefaultAvailableModelContextWindow || models[0].MaxOutputTokens != catalog.DefaultAvailableModelMaxOutputTokens { + t.Fatalf("model defaults = %#v", models[0]) + } + if models[0].ReasoningEfforts == nil || len(models[0].ReasoningEfforts) != 0 { + t.Fatalf("reasoning efforts = %#v, want empty", models[0].ReasoningEfforts) + } + + cached, err := service.ListModels(t.Context(), "openai") + if err != nil { + t.Fatalf("ListModels cached: %v", err) + } + if len(cached) != 1 || cached[0].Code != models[0].Code { + t.Fatalf("cached models = %#v", cached) + } + if requestCount != 1 { + t.Fatalf("model list requests = %d, want 1", requestCount) + } +} + +func TestProviderListModelsRequiresAPIKey(t *testing.T) { + repo := newProviderRepositoryFake() + provider, err := catalog.NewProvider("openai", "OpenAI", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + repo.providers[provider.Code] = provider + service := application.NewProviderService(repo) + + _, listErr := service.ListModels(t.Context(), "openai") + if code := appErrorCode(listErr); code != application.CodeValidation { + t.Fatalf("error code = %v, want validation", code) + } +} + +func TestProviderListAutomaticallyDiscoversOnlyRequestedProvider(t *testing.T) { + var firstRequests atomic.Int32 + firstServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + firstRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-first"}]}`)) + })) + defer firstServer.Close() + + var secondRequests atomic.Int32 + secondServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + secondRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-second"}]}`)) + })) + defer secondServer.Close() + + repo := newProviderRepositoryFake() + for _, item := range []struct { + code string + baseURL string + }{ + {code: "first", baseURL: firstServer.URL}, + {code: "second", baseURL: secondServer.URL}, + } { + provider, err := catalog.NewProvider(item.code, item.code, catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + provider.BaseURL = item.baseURL + provider.APIKey = "test-key" + repo.providers[provider.Code] = provider + } + service := application.NewProviderService(repo) + + providers, err := service.List(t.Context(), "first") + if err != nil { + t.Fatalf("List targeted: %v", err) + } + if len(providers) != 2 || len(providers[0].Models)+len(providers[1].Models) != 1 { + t.Fatalf("targeted providers = %#v", providers) + } + if firstRequests.Load() != 1 || secondRequests.Load() != 0 { + t.Fatalf("targeted requests = (%d, %d), want (1, 0)", firstRequests.Load(), secondRequests.Load()) + } + + if _, err := service.List(t.Context()); err != nil { + t.Fatalf("List all: %v", err) + } + if firstRequests.Load() != 1 || secondRequests.Load() != 1 { + t.Fatalf("all requests = (%d, %d), want (1, 1)", firstRequests.Load(), secondRequests.Load()) + } +} + func TestProviderUpdate(t *testing.T) { _, providerSvc, _ := newServices(t) ctx := t.Context() if _, err := providerSvc.Create(ctx, "openai", application.ProviderInput{ - Name: "OpenAI", Type: catalog.APIOpenAI, BaseURL: "https://old.example", APIKey: "old-key", + Name: "OpenAI", Type: catalog.APIOpenAI, BaseURL: "https://old.example", APIKey: "old-key", FreeFormTool: true, Metadata: shared.Metadata{"region": "us"}, }); err != nil { t.Fatal(err) @@ -101,12 +255,61 @@ func TestProviderUpdate(t *testing.T) { if updated.APIKey != "old-key" || updated.Type != catalog.APIOpenAI || updated.Metadata["region"] != "us" { t.Errorf("unset fields changed: %+v", updated) } + if !updated.FreeFormTool { + t.Error("freeFormTool = false, want true") + } invalid := catalog.APIType("invalid") _, err = providerSvc.Update(ctx, "openai", application.ProviderUpdate{Type: &invalid}) if code := appErrorCode(err); code != application.CodeValidation { t.Errorf("invalid type code = %v, want validation", code) } + anthropic := catalog.APIAnthropic + updated, err = providerSvc.Update(ctx, "openai", application.ProviderUpdate{Type: &anthropic}) + if err != nil { + t.Fatalf("change provider type: %v", err) + } + if updated.FreeFormTool { + t.Error("freeFormTool = true after changing to Anthropic, want ignored") + } +} + +func TestBuiltinProviderAllowsOnlyAPIKeyUpdate(t *testing.T) { + repo := newProviderRepositoryFake() + provider, err := catalog.NewProvider("openai", "OpenAI", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + provider.Builtin = true + provider.Official = true + provider.BaseURL = "https://api.openai.com/v1" + provider.Models = []catalog.Model{{Code: "gpt-5-mini", Name: "GPT-5 mini", ContextWindow: 400_000, MaxOutputTokens: 128_000}} + repo.providers[provider.Code] = provider + providerSvc := application.NewProviderService(repo) + ctx := t.Context() + if _, err := providerSvc.Update(ctx, "openai", application.ProviderUpdate{APIKey: ptr("secret")}); err != nil { + t.Fatalf("API key update: %v", err) + } + updated, err := providerSvc.Get(ctx, "openai") + if err != nil { + t.Fatal(err) + } + if updated.APIKey != "secret" { + t.Fatalf("API key = %q", updated.APIKey) + } + + if _, err := providerSvc.Update(ctx, "openai", application.ProviderUpdate{Name: ptr("Changed")}); appErrorCode(err) != application.CodeValidation { + t.Fatalf("metadata update error = %v, want validation", err) + } + if _, err := providerSvc.Update(ctx, "openai", application.ProviderUpdate{FreeFormTool: ptr(true)}); appErrorCode(err) != application.CodeValidation { + t.Fatalf("freeFormTool update error = %v, want validation", err) + } + if _, err := providerSvc.AddModel(ctx, "openai", "other", application.ModelInput{Name: "Other"}); appErrorCode(err) != application.CodeValidation { + t.Fatalf("builtin AddModel error = %v, want validation", err) + } + if err := providerSvc.Delete(ctx, "openai"); appErrorCode(err) != application.CodeValidation { + t.Fatalf("builtin Delete error = %v, want validation", err) + } } func TestProviderAddModelAndRemoveModel(t *testing.T) { @@ -121,10 +324,6 @@ func TestProviderAddModelAndRemoveModel(t *testing.T) { Name: "Claude Opus 4.8", ContextWindow: 200_000, MaxOutputTokens: 32_000, - ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "low": shared.ReasoningLow, - "high": shared.ReasoningHigh, - }, }) if err != nil { t.Fatalf("AddModel: %v", err) @@ -132,11 +331,11 @@ func TestProviderAddModelAndRemoveModel(t *testing.T) { if len(p.Models) != 1 { t.Fatalf("provider has %d models, want 1", len(p.Models)) } - if effort, ok := p.Models[0].MapReasoningEffort("high"); !ok || effort != shared.ReasoningHigh { - t.Errorf("mapped high effort = %q, %v; want high, true", effort, ok) + if !p.Models[0].SupportsReasoningEffort(shared.ReasoningMax) { + t.Error("default reasoning levels do not include max") } - if p.Models[0].MaxOutputTokens != catalog.DefaultMaxOutputTokens { - t.Errorf("max output tokens = %d, want %d", p.Models[0].MaxOutputTokens, catalog.DefaultMaxOutputTokens) + if p.Models[0].MaxOutputTokens != 32_000 { + t.Errorf("max output tokens = %d, want %d", p.Models[0].MaxOutputTokens, 32_000) } // AddModel is upsert: re-adding the same code replaces. @@ -182,6 +381,7 @@ func TestProviderAddModelAndRemoveModel(t *testing.T) { Name: "Haiku", MaxOutputTokens: 8_000, Light: true, + Reasoning: ptr(false), }); err != nil { t.Fatal(err) } @@ -195,6 +395,9 @@ func TestProviderAddModelAndRemoveModel(t *testing.T) { if p.Models[0].Code.String() != "claude-haiku-4-5" { t.Errorf("remaining model = %s, want claude-haiku-4-5", p.Models[0].Code) } + if p.Models[0].SupportsReasoning() { + t.Error("non-reasoning model supports reasoning") + } // Removing again surfaces not-found. _, err = providerSvc.RemoveModel(ctx, "anthropic", "claude-opus-4-8") @@ -211,31 +414,33 @@ func mustModelCodeForTest(value string) shared.ModelCode { return modelCode } -func TestProviderAddModelRejectsInvalidReasoningEffortMapping(t *testing.T) { +func TestProviderAddModelDefaultsReasoningAndAllowsExplicitDisable(t *testing.T) { _, providerSvc, _ := newServices(t) ctx := t.Context() if _, err := providerSvc.Create(ctx, "openai", application.ProviderInput{Name: "OpenAI", Type: catalog.APIOpenAI}); err != nil { t.Fatal(err) } - tests := []struct { - name string - mapping map[string]shared.ReasoningEffort - }{ - {name: "empty native effort", mapping: map[string]shared.ReasoningEffort{"": shared.ReasoningLow}}, - {name: "invalid agenty effort", mapping: map[string]shared.ReasoningEffort{"minimal": "minimal"}}, + reasoning, err := providerSvc.AddModel(ctx, "openai", "reasoning", application.ModelInput{Name: "Reasoning"}) + if err != nil { + t.Fatal(err) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - _, err := providerSvc.AddModel(ctx, "openai", "gpt-5", application.ModelInput{ - Name: "GPT-5", - MaxOutputTokens: 64_000, - ReasoningEffortMapping: tt.mapping, - }) - if code := appErrorCode(err); code != application.CodeValidation { - t.Errorf("code = %v, want validation", code) - } - }) + if got := reasoning.Models[0].ReasoningEfforts; !slices.Equal(got, shared.StandardReasoningEfforts()) { + t.Fatalf("default reasoning efforts = %v", got) + } + + disabled, err := providerSvc.AddModel(ctx, "openai", "non-reasoning", application.ModelInput{ + Name: "Non-reasoning", Reasoning: ptr(false), + }) + if err != nil { + t.Fatal(err) + } + model, err := disabled.Model("non-reasoning") + if err != nil { + t.Fatal(err) + } + if model.ReasoningEfforts == nil || len(model.ReasoningEfforts) != 0 { + t.Fatalf("disabled reasoning efforts = %v, want empty", model.ReasoningEfforts) } } diff --git a/packages/agenty-core/pkg/application/testhelper_test.go b/packages/agenty-core/pkg/application/testhelper_test.go index 362b4fb..8778949 100644 --- a/packages/agenty-core/pkg/application/testhelper_test.go +++ b/packages/agenty-core/pkg/application/testhelper_test.go @@ -3,11 +3,11 @@ package application_test import ( "context" "errors" - "maps" "slices" "sort" "sync" "testing" + "time" "github.com/google/uuid" @@ -135,11 +135,35 @@ func (r *providerRepositoryFake) Delete(_ context.Context, code shared.Code) err return nil } +func (r *providerRepositoryFake) NeedsModelDiscovery(_ context.Context, code shared.Code) (bool, error) { + provider, ok := r.providers[code] + if !ok { + return false, storage.ErrProviderNotFound + } + return len(provider.Models) == 0, nil +} + +func (r *providerRepositoryFake) ReplaceModels( + _ context.Context, + code shared.Code, + models []catalog.Model, + _ time.Time, +) error { + provider, ok := r.providers[code] + if !ok { + return storage.ErrProviderNotFound + } + updated := cloneProvider(provider) + updated.Models = slices.Clone(models) + r.providers[code] = updated + return nil +} + func cloneProvider(p *catalog.Provider) *catalog.Provider { copy := *p copy.Models = slices.Clone(p.Models) for i := range copy.Models { - copy.Models[i].ReasoningEffortMapping = maps.Clone(copy.Models[i].ReasoningEffortMapping) + copy.Models[i].ReasoningEfforts = slices.Clone(copy.Models[i].ReasoningEfforts) } copy.Metadata = cloneMetadata(p.Metadata) return © diff --git a/packages/agenty-core/pkg/domain/catalog/available_model.go b/packages/agenty-core/pkg/domain/catalog/available_model.go new file mode 100644 index 0000000..dc0c2d2 --- /dev/null +++ b/packages/agenty-core/pkg/domain/catalog/available_model.go @@ -0,0 +1,22 @@ +package catalog + +import ( + "time" + + "github.com/masteryyh/agenty-core/pkg/domain/shared" +) + +const ( + DefaultAvailableModelContextWindow = 256_000 + DefaultAvailableModelMaxOutputTokens int64 = 65_536 + ModelDiscoveryCacheTTL = 8 * time.Hour +) + +type AvailableModel struct { + Code shared.ModelCode `json:"code"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal"` + ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts"` +} diff --git a/packages/agenty-core/pkg/domain/catalog/model.go b/packages/agenty-core/pkg/domain/catalog/model.go index 7c01d31..c1b4c4f 100644 --- a/packages/agenty-core/pkg/domain/catalog/model.go +++ b/packages/agenty-core/pkg/domain/catalog/model.go @@ -9,20 +9,20 @@ import ( const DefaultMaxOutputTokens int64 = 8_192 type Model struct { - Code shared.ModelCode `json:"code"` - Name string `json:"name"` - ContextWindow int `json:"contextWindow"` - MaxOutputTokens int64 `json:"maxOutputTokens"` - MultiModal bool `json:"multiModal"` - Light bool `json:"light"` - ReasoningEffortMapping map[string]shared.ReasoningEffort `json:"reasoningEffortMapping,omitempty"` - IsDefault bool `json:"isDefault"` - CreatedAt time.Time `json:"createdAt"` - UpdatedAt time.Time `json:"updatedAt"` + Code shared.ModelCode `json:"code"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal"` + Light bool `json:"light"` + ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts"` + IsDefault bool `json:"isDefault"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` } func (m *Model) SupportsReasoning() bool { - for _, effort := range m.ReasoningEffortMapping { + for _, effort := range m.ReasoningEfforts { if effort.Enabled() { return true } @@ -31,15 +31,10 @@ func (m *Model) SupportsReasoning() bool { } func (m *Model) SupportsReasoningEffort(effort shared.ReasoningEffort) bool { - for _, mappedEffort := range m.ReasoningEffortMapping { - if mappedEffort == effort { + for _, supported := range m.ReasoningEfforts { + if supported == effort { return true } } return false } - -func (m *Model) MapReasoningEffort(nativeEffort string) (shared.ReasoningEffort, bool) { - effort, ok := m.ReasoningEffortMapping[nativeEffort] - return effort, ok -} diff --git a/packages/agenty-core/pkg/domain/catalog/model_test.go b/packages/agenty-core/pkg/domain/catalog/model_test.go index 29e6d08..c535b79 100644 --- a/packages/agenty-core/pkg/domain/catalog/model_test.go +++ b/packages/agenty-core/pkg/domain/catalog/model_test.go @@ -6,55 +6,30 @@ import ( "github.com/masteryyh/agenty-core/pkg/domain/shared" ) -func TestModel_ReasoningEffortMapping(t *testing.T) { +func TestModelReasoningCapabilities(t *testing.T) { t.Parallel() - model := Model{ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "none": shared.ReasoningOff, - "minimal": shared.ReasoningLow, - "low": shared.ReasoningLow, - "high": shared.ReasoningHigh, + model := Model{ReasoningEfforts: []shared.ReasoningEffort{ + shared.ReasoningLow, + shared.ReasoningMedium, + shared.ReasoningHigh, }} - - tests := []struct { - name string - nativeEffort string - want shared.ReasoningEffort - wantOK bool - }{ - {name: "off", nativeEffort: "none", want: shared.ReasoningOff, wantOK: true}, - {name: "first low alias", nativeEffort: "minimal", want: shared.ReasoningLow, wantOK: true}, - {name: "second low alias", nativeEffort: "low", want: shared.ReasoningLow, wantOK: true}, - {name: "missing", nativeEffort: "xhigh", wantOK: false}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - got, ok := model.MapReasoningEffort(tt.nativeEffort) - if got != tt.want || ok != tt.wantOK { - t.Errorf("MapReasoningEffort(%q) = %q, %v; want %q, %v", tt.nativeEffort, got, ok, tt.want, tt.wantOK) - } - }) - } - if !model.SupportsReasoning() { - t.Error("SupportsReasoning() = false, want true") + t.Fatal("SupportsReasoning() = false, want true") } - if !model.SupportsReasoningEffort(shared.ReasoningLow) { - t.Error("SupportsReasoningEffort(low) = false, want true") + if !model.SupportsReasoningEffort(shared.ReasoningHigh) { + t.Fatal("SupportsReasoningEffort(high) = false, want true") } - if model.SupportsReasoningEffort(shared.ReasoningMax) { - t.Error("SupportsReasoningEffort(max) = true, want false") + if model.SupportsReasoningEffort(shared.ReasoningXHigh) { + t.Fatal("SupportsReasoningEffort(xhigh) = true, want false") } } -func TestModel_SupportsReasoningRequiresEnabledEffort(t *testing.T) { +func TestModelWithoutReasoningEffortsDoesNotSupportReasoning(t *testing.T) { t.Parallel() - model := Model{ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "none": shared.ReasoningOff, - }} + model := Model{ReasoningEfforts: []shared.ReasoningEffort{}} if model.SupportsReasoning() { - t.Error("SupportsReasoning() = true for an off-only mapping") + t.Fatal("SupportsReasoning() = true, want false") } } diff --git a/packages/agenty-core/pkg/domain/catalog/provider.go b/packages/agenty-core/pkg/domain/catalog/provider.go index fd631ee..f96e084 100644 --- a/packages/agenty-core/pkg/domain/catalog/provider.go +++ b/packages/agenty-core/pkg/domain/catalog/provider.go @@ -8,19 +8,26 @@ import ( ) var ( - ErrModelNotFound = errors.New("catalog: model not found") + ErrModelNotFound = errors.New("catalog: model not found") + ErrBuiltinProviderReadOnly = errors.New("catalog: built-in provider is read-only") ) type Provider struct { - Code shared.Code `json:"code"` - Name string `json:"name"` - Type APIType `json:"type"` - BaseURL string `json:"baseUrl"` - APIKey string `json:"apiKey"` - Models []Model `json:"models"` - Metadata shared.Metadata `json:"metadata,omitempty"` - CreatedAt time.Time `json:"createdAt"` - UpdatedAt time.Time `json:"updatedAt"` + Code shared.Code `json:"code"` + Name string `json:"name"` + Type APIType `json:"type"` + BaseURL string `json:"baseUrl"` + APIKey string `json:"apiKey"` + Builtin bool `json:"builtin"` + Official bool `json:"official"` + FreeFormTool bool `json:"freeFormTool"` + ModelsURL string `json:"modelsUrl,omitempty"` + TokenCountURL string `json:"tokenCountUrl,omitempty"` + Models []Model `json:"models"` + ModelsCached bool `json:"modelsCached,omitempty"` + Metadata shared.Metadata `json:"metadata,omitempty"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` } func NewProvider(code, name string, apiType APIType) (*Provider, error) { @@ -53,7 +60,10 @@ func (p *Provider) Model(code shared.ModelCode) (*Model, error) { } func (p *Provider) AddModel(m Model) { - m.MaxOutputTokens = DefaultMaxOutputTokens + m.ReasoningEfforts = shared.NormalizeReasoningEfforts(m.ReasoningEfforts) + if m.MaxOutputTokens <= 0 { + m.MaxOutputTokens = DefaultMaxOutputTokens + } for i := range p.Models { if p.Models[i].Code == m.Code { p.Models[i] = m diff --git a/packages/agenty-core/pkg/domain/catalog/provider_test.go b/packages/agenty-core/pkg/domain/catalog/provider_test.go index b6ad4af..b96e7d4 100644 --- a/packages/agenty-core/pkg/domain/catalog/provider_test.go +++ b/packages/agenty-core/pkg/domain/catalog/provider_test.go @@ -24,11 +24,9 @@ func TestProvider_ModelLifecycle(t *testing.T) { } p.AddModel(Model{ - Code: "model-b", - Name: "B2", - ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "high": shared.ReasoningHigh, - }, + Code: "model-b", + Name: "B2", + ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningHigh}, }) if len(p.Models) != 2 { t.Fatalf("models = %d, want 2 after upsert", len(p.Models)) diff --git a/packages/agenty-core/pkg/domain/shared/reasoning.go b/packages/agenty-core/pkg/domain/shared/reasoning.go index f015184..cb94ca3 100644 --- a/packages/agenty-core/pkg/domain/shared/reasoning.go +++ b/packages/agenty-core/pkg/domain/shared/reasoning.go @@ -23,3 +23,33 @@ func (r ReasoningEffort) Valid() bool { func (r ReasoningEffort) Enabled() bool { return r != "" && r != ReasoningOff } + +func StandardReasoningEfforts() []ReasoningEffort { + return []ReasoningEffort{ + ReasoningLow, + ReasoningMedium, + ReasoningHigh, + ReasoningXHigh, + ReasoningMax, + } +} + +// NormalizeReasoningEfforts keeps supported Agenty levels in their canonical order. +// A nil input means the upstream did not report capabilities and receives the defaults; +// an explicit empty slice identifies a non-reasoning model. +func NormalizeReasoningEfforts(efforts []ReasoningEffort) []ReasoningEffort { + if efforts == nil { + return StandardReasoningEfforts() + } + + normalized := make([]ReasoningEffort, 0, len(efforts)) + for _, supported := range StandardReasoningEfforts() { + for _, effort := range efforts { + if effort == supported { + normalized = append(normalized, supported) + break + } + } + } + return normalized +} diff --git a/packages/agenty-core/pkg/domain/shared/reasoning_test.go b/packages/agenty-core/pkg/domain/shared/reasoning_test.go index f4c5a5c..6b9c48a 100644 --- a/packages/agenty-core/pkg/domain/shared/reasoning_test.go +++ b/packages/agenty-core/pkg/domain/shared/reasoning_test.go @@ -1,6 +1,9 @@ package shared -import "testing" +import ( + "slices" + "testing" +) func TestReasoningEffort_ValidAndEnabled(t *testing.T) { t.Parallel() @@ -37,3 +40,32 @@ func TestReasoningEffort_ValidAndEnabled(t *testing.T) { }) } } + +func TestNormalizeReasoningEfforts(t *testing.T) { + t.Parallel() + + wantDefault := []ReasoningEffort{ + ReasoningLow, + ReasoningMedium, + ReasoningHigh, + ReasoningXHigh, + ReasoningMax, + } + if got := NormalizeReasoningEfforts(nil); !slices.Equal(got, wantDefault) { + t.Fatalf("NormalizeReasoningEfforts(nil) = %v, want %v", got, wantDefault) + } + if got := NormalizeReasoningEfforts([]ReasoningEffort{}); got == nil || len(got) != 0 { + t.Fatalf("NormalizeReasoningEfforts(empty) = %v, want empty non-nil slice", got) + } + + got := NormalizeReasoningEfforts([]ReasoningEffort{ + ReasoningMax, + "minimal", + ReasoningLow, + ReasoningLow, + }) + want := []ReasoningEffort{ReasoningLow, ReasoningMax} + if !slices.Equal(got, want) { + t.Fatalf("NormalizeReasoningEfforts(filtered) = %v, want %v", got, want) + } +} diff --git a/packages/agenty-core/pkg/infra/catalogdata/catalog.go b/packages/agenty-core/pkg/infra/catalogdata/catalog.go new file mode 100644 index 0000000..6b5b56b --- /dev/null +++ b/packages/agenty-core/pkg/infra/catalogdata/catalog.go @@ -0,0 +1,114 @@ +package catalogdata + +import ( + _ "embed" + "errors" + "fmt" + "net/url" + "strings" + + json "github.com/bytedance/sonic" + + "github.com/masteryyh/agenty-core/pkg/domain/catalog" + "github.com/masteryyh/agenty-core/pkg/domain/shared" +) + +//go:embed providers.json +var providersJSON []byte + +func LoadProviders() ([]*catalog.Provider, error) { + var providers []*catalog.Provider + if err := json.Unmarshal(providersJSON, &providers); err != nil { + return nil, fmt.Errorf("catalogdata: decode embedded providers: %w", err) + } + if len(providers) == 0 { + return nil, errors.New("catalogdata: embedded providers are empty") + } + + seen := make(map[shared.Code]struct{}, len(providers)) + for index, provider := range providers { + if provider == nil { + return nil, fmt.Errorf("catalogdata: provider %d is null", index) + } + if err := validateProvider(provider); err != nil { + return nil, fmt.Errorf("catalogdata: provider %d: %w", index, err) + } + if _, ok := seen[provider.Code]; ok { + return nil, fmt.Errorf("catalogdata: duplicate provider code %q", provider.Code) + } + + seen[provider.Code] = struct{}{} + provider.Builtin = true + if provider.Models == nil { + provider.Models = make([]catalog.Model, 0) + } + for modelIndex := range provider.Models { + provider.Models[modelIndex].ReasoningEfforts = shared.NormalizeReasoningEfforts( + provider.Models[modelIndex].ReasoningEfforts, + ) + } + } + return providers, nil +} + +func validateProvider(provider *catalog.Provider) error { + if !provider.Code.Valid() { + return errors.New("invalid provider code") + } + if strings.TrimSpace(provider.Name) == "" { + return errors.New("provider name is empty") + } + if !provider.Type.Valid() { + return fmt.Errorf("invalid API type %q", provider.Type) + } + if strings.TrimSpace(provider.BaseURL) == "" { + return errors.New("base URL is empty") + } + if err := validateRelativeURL(provider.ModelsURL); err != nil { + return fmt.Errorf("models URL: %w", err) + } + if err := validateRelativeURL(provider.TokenCountURL); err != nil { + return fmt.Errorf("token count URL: %w", err) + } + for index := range provider.Models { + model := &provider.Models[index] + if model.Code.IsZero() || strings.TrimSpace(model.Name) == "" { + return fmt.Errorf("model %d has an empty code or name", index) + } + if model.ContextWindow <= 0 || model.MaxOutputTokens <= 0 { + return fmt.Errorf("model %q has invalid token limits", model.Code) + } + if err := validateReasoningEfforts(model); err != nil { + return err + } + } + return nil +} + +func validateReasoningEfforts(model *catalog.Model) error { + seen := make(map[shared.ReasoningEffort]struct{}, len(model.ReasoningEfforts)) + for _, effort := range model.ReasoningEfforts { + if !effort.Valid() || !effort.Enabled() { + return fmt.Errorf("model %q has invalid reasoning effort %q", model.Code, effort) + } + if _, ok := seen[effort]; ok { + return fmt.Errorf("model %q repeats reasoning effort %q", model.Code, effort) + } + seen[effort] = struct{}{} + } + return nil +} + +func validateRelativeURL(raw string) error { + if strings.TrimSpace(raw) == "" { + return nil + } + parsed, err := url.Parse(raw) + if err != nil { + return err + } + if parsed.IsAbs() || parsed.Host != "" { + return errors.New("must be relative to the provider base URL") + } + return nil +} diff --git a/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go b/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go new file mode 100644 index 0000000..2f9fd95 --- /dev/null +++ b/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go @@ -0,0 +1,57 @@ +package catalogdata + +import "testing" + +func TestLoadProviders(t *testing.T) { + providers, err := LoadProviders() + if err != nil { + t.Fatalf("LoadProviders: %v", err) + } + if len(providers) != 5 { + t.Fatalf("providers = %d, want 5", len(providers)) + } + + byCode := make(map[string]struct { + name string + typeName string + }) + for _, provider := range providers { + if !provider.Builtin { + t.Errorf("provider %s is not built in", provider.Code) + } + if provider.Code != "openrouter" && !provider.Official { + t.Errorf("provider %s is not official", provider.Code) + } + byCode[provider.Code.String()] = struct { + name string + typeName string + }{provider.Name, string(provider.Type)} + } + if got := byCode["openai_legacy"]; got.name != "OpenAI (Legacy API)" || got.typeName != "openai_completions" { + t.Errorf("legacy OpenAI provider = %+v", got) + } + openRouter := providers[2] + if openRouter.Code != "openrouter" || openRouter.Name != "OpenRouter" || openRouter.Type != "openai" { + t.Errorf("OpenRouter provider = %+v", openRouter) + } + if openRouter.Official || openRouter.BaseURL != "https://openrouter.ai/api/v1" || openRouter.ModelsURL != "models" { + t.Errorf("OpenRouter discovery config = %+v", openRouter) + } + if openRouter.Models == nil || len(openRouter.Models) != 0 { + t.Errorf("OpenRouter embedded models = %#v, want empty", openRouter.Models) + } + openAI := providers[0] + if openAI.Code != "openai" || !openAI.FreeFormTool { + t.Errorf("OpenAI freeFormTool = %v, want true", openAI.FreeFormTool) + } + for _, provider := range providers { + for _, model := range provider.Models { + if model.ContextWindow <= 0 || model.MaxOutputTokens <= 0 { + t.Errorf("model %s/%s has invalid token limits", provider.Code, model.Code) + } + if model.ReasoningEfforts == nil { + t.Errorf("model %s/%s has nil reasoning efforts", provider.Code, model.Code) + } + } + } +} diff --git a/packages/agenty-core/pkg/infra/catalogdata/providers.json b/packages/agenty-core/pkg/infra/catalogdata/providers.json new file mode 100644 index 0000000..299b6a9 --- /dev/null +++ b/packages/agenty-core/pkg/infra/catalogdata/providers.json @@ -0,0 +1,279 @@ +[ + { + "name": "OpenAI", + "code": "openai", + "type": "openai", + "baseUrl": "https://api.openai.com/v1", + "official": true, + "freeFormTool": true, + "modelsUrl": "models", + "tokenCountUrl": "responses/input_tokens", + "models": [ + { + "name": "GPT-5.6 Sol", + "code": "gpt-5.6-sol", + "contextWindow": 1050000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": true + }, + { + "name": "GPT-5.6 Terra", + "code": "gpt-5.6-terra", + "contextWindow": 1050000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": false + }, + { + "name": "GPT-5.6 Luna", + "code": "gpt-5.6-luna", + "contextWindow": 1050000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": true, + "isDefault": false + }, + { + "name": "GPT-5.5", + "code": "gpt-5.5", + "contextWindow": 1050000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh" + ], + "multiModal": true, + "light": false, + "isDefault": false + } + ] + }, + { + "name": "OpenAI (Legacy API)", + "code": "openai_legacy", + "type": "openai_completions", + "baseUrl": "https://api.openai.com/v1", + "official": true, + "modelsUrl": "models", + "tokenCountUrl": "", + "models": [ + { + "name": "GPT-4o", + "code": "gpt-4o", + "contextWindow": 128000, + "maxOutputTokens": 16384, + "reasoningEfforts": [], + "multiModal": true, + "light": false, + "isDefault": true + }, + { + "name": "GPT-4.1", + "code": "gpt-4.1", + "contextWindow": 1047576, + "maxOutputTokens": 32768, + "reasoningEfforts": [], + "multiModal": true, + "light": false + }, + { + "name": "GPT-4o mini", + "code": "gpt-4o-mini", + "contextWindow": 128000, + "maxOutputTokens": 16384, + "reasoningEfforts": [], + "multiModal": true, + "light": true + } + ] + }, + { + "name": "OpenRouter", + "code": "openrouter", + "type": "openai", + "baseUrl": "https://openrouter.ai/api/v1", + "modelsUrl": "models", + "tokenCountUrl": "", + "models": [] + }, + { + "name": "Anthropic", + "code": "anthropic", + "type": "anthropic", + "baseUrl": "https://api.anthropic.com", + "official": true, + "modelsUrl": "v1/models", + "tokenCountUrl": "v1/messages/count_tokens", + "models": [ + { + "name": "Claude Fable 5", + "code": "claude-fable-5", + "contextWindow": 1000000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": false + }, + { + "name": "Claude Opus 5", + "code": "claude-opus-5", + "contextWindow": 1000000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": true + }, + { + "name": "Claude Opus 4.6", + "code": "claude-opus-4-6", + "contextWindow": 1000000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": false + }, + { + "name": "Claude Sonnet 5", + "code": "claude-sonnet-5", + "contextWindow": 1000000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": false + }, + { + "name": "Claude Sonnet 4.6", + "code": "claude-sonnet-4-6", + "contextWindow": 1000000, + "maxOutputTokens": 128000, + "reasoningEfforts": [ + "low", + "medium", + "high", + "xhigh", + "max" + ], + "multiModal": true, + "light": false, + "isDefault": false + }, + { + "name": "Claude Haiku 4.5", + "code": "claude-haiku-4-5", + "contextWindow": 200000, + "maxOutputTokens": 64000, + "reasoningEfforts": [], + "multiModal": true, + "light": true, + "isDefault": false + } + ] + }, + { + "name": "Google", + "code": "google", + "type": "gemini", + "baseUrl": "https://generativelanguage.googleapis.com/v1beta", + "official": true, + "modelsUrl": "models", + "tokenCountUrl": "models/{model}:countTokens", + "models": [ + { + "name": "Gemini 3.7 Flash", + "code": "gemini-3.7-flash", + "contextWindow": 1048576, + "maxOutputTokens": 65536, + "reasoningEfforts": [ + "low", + "medium", + "high" + ], + "multiModal": true, + "light": true, + "isDefault": true + }, + { + "name": "Gemini 3.5 Flash Lite", + "code": "gemini-3.5-flash-lite", + "contextWindow": 1048576, + "maxOutputTokens": 65536, + "reasoningEfforts": [ + "low", + "medium", + "high" + ], + "multiModal": true, + "light": true, + "isDefault": false + }, + { + "name": "Gemini 3.1 Pro (Preview)", + "code": "gemini-3.1-pro-preview", + "contextWindow": 1048576, + "maxOutputTokens": 65536, + "reasoningEfforts": [ + "low", + "medium", + "high" + ], + "multiModal": true, + "light": false, + "isDefault": false + } + ] + } +] diff --git a/packages/agenty-core/pkg/infra/initialize/initialize.go b/packages/agenty-core/pkg/infra/initialize/initialize.go index 24a5e28..a700081 100644 --- a/packages/agenty-core/pkg/infra/initialize/initialize.go +++ b/packages/agenty-core/pkg/infra/initialize/initialize.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" + "github.com/masteryyh/agenty-core/pkg/infra/catalogdata" "github.com/masteryyh/agenty-core/pkg/infra/config" "github.com/masteryyh/agenty-core/pkg/infra/storage" ) @@ -35,16 +36,19 @@ func OpenRepositories(ctx context.Context) (*Repositories, error) { if err != nil { return nil, err } + builtinProviders, err := catalogdata.LoadProviders() + if err != nil { + return nil, err + } db, err := storage.OpenDB(paths.DatabaseFile) if err != nil { return nil, err } - return &Repositories{ Conversation: storage.NewConversationRepository(db, paths.SessionsDir), Agent: storage.NewAgentRepository(paths.AgentsDir), - Catalog: storage.NewCatalogRepository(paths.ProvidersDir), + Catalog: storage.NewCatalogRepository(paths.ProvidersDir, builtinProviders...), db: db, }, nil } diff --git a/packages/agenty-core/pkg/infra/initialize/initialize_test.go b/packages/agenty-core/pkg/infra/initialize/initialize_test.go index 44fac59..5958ed0 100644 --- a/packages/agenty-core/pkg/infra/initialize/initialize_test.go +++ b/packages/agenty-core/pkg/infra/initialize/initialize_test.go @@ -27,6 +27,35 @@ func TestOpenRepositoriesEndToEnd(t *testing.T) { defer repos.Close() ctx := context.Background() + builtinProviders, err := repos.Catalog.List(ctx) + if err != nil { + t.Fatalf("List built-in providers: %v", err) + } + if len(builtinProviders) != 5 { + t.Fatalf("built-in providers = %d, want 5", len(builtinProviders)) + } + openRouter, err := repos.Catalog.Get(ctx, mustCode("openrouter")) + if err != nil { + t.Fatalf("Get OpenRouter provider: %v", err) + } + if openRouter.Models == nil || len(openRouter.Models) != 0 { + t.Fatalf("OpenRouter embedded models = %#v, want empty", openRouter.Models) + } + builtin, err := repos.Catalog.Get(ctx, mustCode("openai_legacy")) + if err != nil { + t.Fatalf("Get built-in provider: %v", err) + } + builtin.APIKey = "secret" + if err := repos.Catalog.Save(ctx, builtin); err != nil { + t.Fatalf("Save built-in API key: %v", err) + } + credentialData, err := os.ReadFile(filepath.Join(tmpDir, "providers", "openai_legacy.json")) + if err != nil { + t.Fatalf("read built-in credentials: %v", err) + } + if string(credentialData) != "{\n \"apiKey\": \"secret\"\n}" { + t.Fatalf("built-in credentials = %s", credentialData) + } // Directory structure, config and SQLite database were created. for _, dir := range []string{"sessions", "agents", "providers"} { @@ -45,7 +74,7 @@ func TestOpenRepositoriesEndToEnd(t *testing.T) { } // Create and persist the catalog and an agent with its default session model. - provider, err := catalog.NewProvider("anthropic", "Anthropic", catalog.APIAnthropic) + provider, err := catalog.NewProvider("integration-anthropic", "Anthropic", catalog.APIAnthropic) if err != nil { t.Fatal(err) } diff --git a/packages/agenty-core/pkg/infra/llm/anthropic.go b/packages/agenty-core/pkg/infra/llm/anthropic.go index aedc37b..61bd8be 100644 --- a/packages/agenty-core/pkg/infra/llm/anthropic.go +++ b/packages/agenty-core/pkg/infra/llm/anthropic.go @@ -103,10 +103,7 @@ func (caller *anthropicCaller) params(request modelRequest) (anthropic.MessageNe if err != nil { return anthropic.MessageNewParams{}, err } - effort, err := nativeReasoningEffort(caller.model, request.ReasoningEffort) - if err != nil { - return anthropic.MessageNewParams{}, err - } + effort := modelReasoningEffort(caller.model, request.ReasoningEffort) messages := make([]anthropic.MessageParam, 0, len(request.Messages)) for index, message := range request.Messages { diff --git a/packages/agenty-core/pkg/infra/llm/convert.go b/packages/agenty-core/pkg/infra/llm/convert.go index 4048674..d612f07 100644 --- a/packages/agenty-core/pkg/infra/llm/convert.go +++ b/packages/agenty-core/pkg/infra/llm/convert.go @@ -3,7 +3,6 @@ package llm import ( "encoding/base64" "fmt" - "sort" "strings" json "github.com/bytedance/sonic" @@ -32,30 +31,11 @@ func validateRequest(request modelRequest) error { return nil } -func nativeReasoningEffort(model catalog.Model, effort shared.ReasoningEffort) (string, error) { - if effort == "" || effort == shared.ReasoningOff { - return "", nil - } - if native, ok := model.ReasoningEffortMapping[string(effort)]; ok && native == effort { - return string(effort), nil - } - - matches := make([]string, 0, 1) - for native, mapped := range model.ReasoningEffortMapping { - if mapped == effort { - matches = append(matches, native) - } - } - sort.Strings(matches) - - switch len(matches) { - case 0: - return "", invalidRequest("model %q does not support reasoning effort %q", model.Code, effort) - case 1: - return matches[0], nil - default: - return "", invalidRequest("model %q maps reasoning effort %q ambiguously to %s", model.Code, effort, strings.Join(matches, ", ")) +func modelReasoningEffort(model catalog.Model, effort shared.ReasoningEffort) string { + if effort == "" || effort == shared.ReasoningOff || !model.SupportsReasoning() { + return "" } + return string(effort) } func systemPrompt(request modelRequest) (string, error) { diff --git a/packages/agenty-core/pkg/infra/llm/convert_test.go b/packages/agenty-core/pkg/infra/llm/convert_test.go index cec22b1..0c7415c 100644 --- a/packages/agenty-core/pkg/infra/llm/convert_test.go +++ b/packages/agenty-core/pkg/infra/llm/convert_test.go @@ -19,43 +19,58 @@ import ( "github.com/masteryyh/agenty-core/pkg/domain/shared" ) -func TestNativeReasoningEffort(t *testing.T) { +func TestModelReasoningEffort(t *testing.T) { t.Parallel() model := testModel() tests := []struct { - name string - effort shared.ReasoningEffort - want string - wantErr bool + name string + effort shared.ReasoningEffort + want string }{ {name: "empty", effort: "", want: ""}, {name: "off", effort: shared.ReasoningOff, want: ""}, {name: "exact", effort: shared.ReasoningLow, want: "low"}, - {name: "mapped", effort: shared.ReasoningHigh, want: "HIGH"}, - {name: "unsupported", effort: shared.ReasoningMax, wantErr: true}, + {name: "high", effort: shared.ReasoningHigh, want: "high"}, + {name: "max", effort: shared.ReasoningMax, want: "max"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - got, err := nativeReasoningEffort(model, tt.effort) - if tt.wantErr { - if !errors.Is(err, ErrInvalidRequest) { - t.Fatalf("nativeReasoningEffort() error = %v, want ErrInvalidRequest", err) - } - return - } - if err != nil { - t.Fatalf("nativeReasoningEffort() error = %v", err) - } + got := modelReasoningEffort(model, tt.effort) if got != tt.want { - t.Fatalf("nativeReasoningEffort() = %q, want %q", got, tt.want) + t.Fatalf("modelReasoningEffort() = %q, want %q", got, tt.want) } }) } } +func TestModelReasoningEffortSendsUnsupportedLevelToUpstream(t *testing.T) { + model := catalog.Model{ + Code: "gpt-5-mini", + ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}, + } + got := modelReasoningEffort(model, shared.ReasoningMax) + if got != "max" { + t.Fatalf("modelReasoningEffort() = %q, want max", got) + } +} + +func TestModelReasoningEffortIgnoresNonReasoningModel(t *testing.T) { + model := catalog.Model{Code: "gpt-4o", ReasoningEfforts: []shared.ReasoningEffort{}} + got := modelReasoningEffort(model, shared.ReasoningHigh) + if got != "" { + t.Fatalf("modelReasoningEffort() = %q, want empty", got) + } +} + +func TestGoogleThinkingLevelPassesStandardEffortThrough(t *testing.T) { + if got := googleThinkingLevel("max"); got != genai.ThinkingLevel("MAX") { + t.Fatalf("googleThinkingLevel(max) = %q, want MAX", got) + } +} + func TestProviderRequestConversions(t *testing.T) { t.Parallel() @@ -70,7 +85,7 @@ func TestProviderRequestConversions(t *testing.T) { t.Run("OpenAI Responses", func(t *testing.T) { t.Parallel() - params, err := (&openAIResponsesCaller{model: modelWithReasoningNative("high")}).params(request) + params, err := (&openAIResponsesCaller{model: reasoningModel()}).params(request) if err != nil { t.Fatalf("convert request: %v", err) } @@ -99,7 +114,7 @@ func TestProviderRequestConversions(t *testing.T) { t.Run("OpenAI Chat Completions", func(t *testing.T) { t.Parallel() - params, err := (&openAIChatCaller{model: modelWithReasoningNative("high")}).params(request) + params, err := (&openAIChatCaller{model: reasoningModel()}).params(request) if err != nil { t.Fatalf("convert request: %v", err) } @@ -121,7 +136,7 @@ func TestProviderRequestConversions(t *testing.T) { t.Run("Anthropic Messages", func(t *testing.T) { t.Parallel() - params, err := (&anthropicCaller{model: modelWithReasoningNative("high")}).params(request) + params, err := (&anthropicCaller{model: reasoningModel()}).params(request) if err != nil { t.Fatalf("convert request: %v", err) } @@ -147,7 +162,7 @@ func TestProviderRequestConversions(t *testing.T) { t.Run("Google GenAI", func(t *testing.T) { t.Parallel() - contents, config, err := (&googleCaller{model: modelWithReasoningNative("HIGH")}).params(request) + contents, config, err := (&googleCaller{model: reasoningModel()}).params(request) if err != nil { t.Fatalf("convert request: %v", err) } @@ -166,12 +181,10 @@ func TestProviderRequestConversions(t *testing.T) { }) } -func modelWithReasoningNative(native string) catalog.Model { +func reasoningModel() catalog.Model { return catalog.Model{ - Code: "test-model", - ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - native: shared.ReasoningHigh, - }, + Code: "test-model", + ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningHigh}, } } @@ -305,37 +318,31 @@ func TestApplyPatchToolRegistrations(t *testing.T) { definitions := []modelToolDefinition{ testNamedTool("read_file"), - testNamedTool("write_file"), - testNamedTool("patch_file"), - testNamedTool("delete_file"), testApplyPatchTool(), } - native, err := openAIResponsesTools(definitions, true) + filesystem, err := openAIResponsesTools(definitions, true, false) if err != nil { t.Fatal(err) } - if len(native) != 2 || native[0].OfFunction == nil || native[0].OfFunction.Name != "read_file" || - native[1].OfApplyPatch == nil { - t.Fatalf("native Responses tools = %#v, want read_file and native apply_patch", native) + if len(filesystem) != 1 || filesystem[0].OfFunction == nil || filesystem[0].OfFunction.Name != "read_file" { + t.Fatalf("non-free-form Responses tools = %#v, want read_file only", filesystem) } - compatible, err := openAIResponsesTools(definitions, false) + freeForm, err := openAIResponsesTools(definitions, true, true) if err != nil { t.Fatal(err) } - if len(compatible) != 2 || compatible[0].OfFunction == nil || compatible[0].OfFunction.Name != "read_file" || - compatible[1].OfCustom == nil || compatible[1].OfCustom.Name != "apply_patch" { - t.Fatalf("compatible Responses tools = %#v, want read_file and custom apply_patch", compatible) + if len(freeForm) != 2 || freeForm[0].OfFunction == nil || freeForm[0].OfFunction.Name != "read_file" || + freeForm[1].OfCustom == nil || freeForm[1].OfCustom.Name != "apply_patch" { + t.Fatalf("free-form Responses tools = %#v, want read_file and custom apply_patch", freeForm) } chat, err := openAIChatTools(definitions) if err != nil { t.Fatal(err) } - if names := openAIChatToolNames(chat); !slices.Equal(names, []string{ - "read_file", "write_file", "patch_file", "delete_file", - }) { + if names := openAIChatToolNames(chat); !slices.Equal(names, []string{"read_file"}) { t.Errorf("OpenAI Chat tools = %q", names) } @@ -343,16 +350,16 @@ func TestApplyPatchToolRegistrations(t *testing.T) { if err != nil { t.Fatal(err) } - if len(anthropicDefinitions) != 4 { - t.Errorf("Anthropic tools = %d, want 4 original filesystem tools", len(anthropicDefinitions)) + if len(anthropicDefinitions) != 1 { + t.Errorf("Anthropic tools = %d, want read_file only", len(anthropicDefinitions)) } googleDefinitions, err := googleTools(definitions) if err != nil { t.Fatal(err) } - if len(googleDefinitions) != 4 { - t.Errorf("Google tools = %#v, want 4 original filesystem tools", googleDefinitions) + if len(googleDefinitions) != 1 { + t.Errorf("Google tools = %#v, want read_file only", googleDefinitions) } } @@ -840,6 +847,15 @@ func TestApplyPatchMessageConversionsAcrossProviders(t *testing.T) { t.Fatalf("native Responses history = %#v", nativeItems) } + freeFormItems, err := openAIResponsesMessages([]conversation.Message{assistant, result}, true, true) + if err != nil { + t.Fatal(err) + } + if len(freeFormItems) != 2 || freeFormItems[0].OfCustomToolCall == nil || + freeFormItems[1].OfCustomToolCallOutput == nil { + t.Fatalf("free-form Responses history = %#v", freeFormItems) + } + compatibleItems, err := openAIResponsesMessages([]conversation.Message{assistant, result}, false) if err != nil { t.Fatal(err) @@ -1057,24 +1073,25 @@ func TestNewCallerConfiguresNativeOpenAIResponsesToolsByProviderIdentity(t *test t.Parallel() tests := []struct { - name string - provider catalog.Provider - want bool + name string + provider catalog.Provider + wantNative bool + wantFreeForm bool }{ { name: "built-in OpenAI with SDK default URL", provider: catalog.Provider{ - Code: "openai", Type: catalog.APIOpenAI, APIKey: "test-key", + Code: "openai", Type: catalog.APIOpenAI, APIKey: "test-key", Official: true, FreeFormTool: true, }, - want: true, + wantNative: true, wantFreeForm: true, }, { name: "built-in OpenAI with official URL", provider: catalog.Provider{ - Code: "openai", Type: catalog.APIOpenAI, APIKey: "test-key", + Code: "openai", Type: catalog.APIOpenAI, APIKey: "test-key", Official: true, BaseURL: "https://api.openai.com/v1/", }, - want: true, + wantNative: true, }, { name: "OpenRouter Responses compatibility", @@ -1110,19 +1127,32 @@ func TestNewCallerConfiguresNativeOpenAIResponsesToolsByProviderIdentity(t *test if !ok { t.Fatalf("NewCaller() = %T, want *openAIResponsesCaller", caller) } - if responsesCaller.nativeOpenAI != tt.want { - t.Errorf("nativeOpenAI = %v, want %v", responsesCaller.nativeOpenAI, tt.want) + if responsesCaller.nativeOpenAI != tt.wantNative { + t.Errorf("nativeOpenAI = %v, want %v", responsesCaller.nativeOpenAI, tt.wantNative) + } + if responsesCaller.freeFormTool != tt.wantFreeForm { + t.Errorf("freeFormTool = %v, want %v", responsesCaller.freeFormTool, tt.wantFreeForm) } }) } } +func TestNativeOpenAIResponsesProviderRequiresOfficialResponsesAPI(t *testing.T) { + if !nativeOpenAIResponsesProvider(catalog.Provider{Type: catalog.APIOpenAI, Official: true}) { + t.Fatal("official OpenAI Responses provider was not recognized") + } + if nativeOpenAIResponsesProvider(catalog.Provider{Type: catalog.APIOpenAICompletions, Official: true}) { + t.Fatal("official OpenAI Chat Completions provider was recognized as Responses") + } + if nativeOpenAIResponsesProvider(catalog.Provider{Type: catalog.APIOpenAI}) { + t.Fatal("compatible Responses provider was recognized as official OpenAI") + } +} + func testModel() catalog.Model { return catalog.Model{ - Code: "test-model", - ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "low": shared.ReasoningLow, "HIGH": shared.ReasoningHigh, - }, + Code: "test-model", + ReasoningEfforts: shared.StandardReasoningEfforts(), } } diff --git a/packages/agenty-core/pkg/infra/llm/factory.go b/packages/agenty-core/pkg/infra/llm/factory.go index 0c0c286..d8ad5cb 100644 --- a/packages/agenty-core/pkg/infra/llm/factory.go +++ b/packages/agenty-core/pkg/infra/llm/factory.go @@ -4,7 +4,6 @@ import ( "context" "fmt" "net/http" - "net/url" "strings" "github.com/anthropics/anthropic-sdk-go" @@ -61,6 +60,7 @@ func NewCaller( client: &client, model: model, nativeOpenAI: nativeOpenAIResponsesProvider(provider), + freeFormTool: provider.FreeFormTool, }, nil case catalog.APIOpenAICompletions: client := newOpenAIClient(provider, config) @@ -80,23 +80,7 @@ func NewCaller( } func nativeOpenAIResponsesProvider(provider catalog.Provider) bool { - if provider.Code.String() != "openai" { - return false - } - baseURL := strings.TrimSpace(provider.BaseURL) - if baseURL == "" { - return true - } - - parsed, err := url.Parse(baseURL) - if err != nil { - return false - } - - return strings.EqualFold(parsed.Scheme, "https") && - strings.EqualFold(parsed.Hostname(), "api.openai.com") && - strings.TrimRight(parsed.EscapedPath(), "/") == "/v1" && - parsed.RawQuery == "" && parsed.Fragment == "" + return provider.Official && provider.Type == catalog.APIOpenAI } func newOpenAIClient(provider catalog.Provider, config factoryConfig) openai.Client { diff --git a/packages/agenty-core/pkg/infra/llm/google.go b/packages/agenty-core/pkg/infra/llm/google.go index b5d2684..27ccdd4 100644 --- a/packages/agenty-core/pkg/infra/llm/google.go +++ b/packages/agenty-core/pkg/infra/llm/google.go @@ -88,10 +88,7 @@ func (caller *googleCaller) params(request modelRequest) ([]*genai.Content, *gen if err != nil { return nil, nil, err } - effort, err := nativeReasoningEffort(caller.model, request.ReasoningEffort) - if err != nil { - return nil, nil, err - } + effort := modelReasoningEffort(caller.model, request.ReasoningEffort) toolNames := googleToolNames(request.Messages) contents := make([]*genai.Content, 0, len(request.Messages)) @@ -123,11 +120,7 @@ func (caller *googleCaller) params(request modelRequest) ([]*genai.Content, *gen budget := int32(request.ReasoningBudgetTokens) thinking.ThinkingBudget = &budget } else { - level, err := googleThinkingLevel(effort) - if err != nil { - return nil, nil, err - } - thinking.ThinkingLevel = level + thinking.ThinkingLevel = googleThinkingLevel(effort) } config.ThinkingConfig = thinking } @@ -304,19 +297,8 @@ func googleMessage(message conversation.Message, toolNames map[string]string) (* return genai.NewContentFromParts(parts, role), nil } -func googleThinkingLevel(effort string) (genai.ThinkingLevel, error) { - switch strings.ToUpper(effort) { - case "MINIMAL": - return genai.ThinkingLevelMinimal, nil - case "LOW": - return genai.ThinkingLevelLow, nil - case "MEDIUM": - return genai.ThinkingLevelMedium, nil - case "HIGH": - return genai.ThinkingLevelHigh, nil - default: - return "", invalidRequest("Google does not support native thinking level %q", effort) - } +func googleThinkingLevel(effort string) genai.ThinkingLevel { + return genai.ThinkingLevel(strings.ToUpper(effort)) } func googleResponse(result *genai.GenerateContentResponse) (*modelResponse, error) { diff --git a/packages/agenty-core/pkg/infra/llm/openai_chat.go b/packages/agenty-core/pkg/infra/llm/openai_chat.go index 60b9196..f532ff2 100644 --- a/packages/agenty-core/pkg/infra/llm/openai_chat.go +++ b/packages/agenty-core/pkg/infra/llm/openai_chat.go @@ -116,10 +116,7 @@ func (caller *openAIChatCaller) params(request modelRequest) (openai.ChatComplet if err != nil { return openai.ChatCompletionNewParams{}, err } - effort, err := nativeReasoningEffort(caller.model, request.ReasoningEffort) - if err != nil { - return openai.ChatCompletionNewParams{}, err - } + effort := modelReasoningEffort(caller.model, request.ReasoningEffort) messages := make([]openai.ChatCompletionMessageParamUnion, 0, len(request.Messages)+1) if prompt != "" { diff --git a/packages/agenty-core/pkg/infra/llm/openai_responses.go b/packages/agenty-core/pkg/infra/llm/openai_responses.go index f581c3d..f1b9ac4 100644 --- a/packages/agenty-core/pkg/infra/llm/openai_responses.go +++ b/packages/agenty-core/pkg/infra/llm/openai_responses.go @@ -20,6 +20,7 @@ type openAIResponsesCaller struct { client *openai.Client model catalog.Model nativeOpenAI bool + freeFormTool bool } func (caller *openAIResponsesCaller) Invoke(ctx context.Context, request modelRequest) (*modelResponse, error) { @@ -165,17 +166,14 @@ func (caller *openAIResponsesCaller) params(request modelRequest) (responses.Res if err != nil { return responses.ResponseNewParams{}, err } - effort, err := nativeReasoningEffort(caller.model, request.ReasoningEffort) - if err != nil { - return responses.ResponseNewParams{}, err - } + effort := modelReasoningEffort(caller.model, request.ReasoningEffort) - input, err := openAIResponsesMessages(request.Messages, caller.nativeOpenAI) + input, err := openAIResponsesMessages(request.Messages, caller.nativeOpenAI, caller.freeFormTool) if err != nil { return responses.ResponseNewParams{}, err } - tools, err := openAIResponsesTools(request.Tools, caller.nativeOpenAI) + tools, err := openAIResponsesTools(request.Tools, caller.nativeOpenAI, caller.freeFormTool) if err != nil { return responses.ResponseNewParams{}, err } @@ -200,13 +198,22 @@ func (caller *openAIResponsesCaller) params(request modelRequest) (responses.Res return params, nil } -func openAIResponsesTools(definitions []modelToolDefinition, nativeOpenAI bool) ([]responses.ToolUnionParam, error) { +func openAIResponsesTools( + definitions []modelToolDefinition, + nativeOpenAI bool, + freeFormTool ...bool, +) ([]responses.ToolUnionParam, error) { + useFreeFormTool := !nativeOpenAI + if len(freeFormTool) > 0 { + useFreeFormTool = freeFormTool[0] + } + tools := make([]responses.ToolUnionParam, 0, len(definitions)) for _, definition := range definitions { - if isReplacedFileTool(definition.Name) { + if definition.Type == agentloop.ToolTypeApplyPatch && !useFreeFormTool { continue } - tool, err := openAIResponsesToolDefinition(definition, nativeOpenAI) + tool, err := openAIResponsesToolDefinitionWithFreeForm(definition, nativeOpenAI, useFreeFormTool) if err != nil { return nil, err } @@ -218,7 +225,12 @@ func openAIResponsesTools(definitions []modelToolDefinition, nativeOpenAI bool) func openAIResponsesMessages( messages []conversation.Message, nativeOpenAI bool, + freeFormTool ...bool, ) (responses.ResponseInputParam, error) { + useFreeFormTool := !nativeOpenAI + if len(freeFormTool) > 0 { + useFreeFormTool = freeFormTool[0] + } callSources := openAIResponsesCallSources(messages) input := make(responses.ResponseInputParam, 0, len(messages)) for index, message := range messages { @@ -228,6 +240,7 @@ func openAIResponsesMessages( items, err := openAIResponsesMessageWithNativeCallIDs( message, nativeOpenAI, + useFreeFormTool, callSources, ) if err != nil { @@ -264,16 +277,15 @@ func openAIResponsesCallSources(messages []conversation.Message) map[string]open return sources } -func isReplacedFileTool(name string) bool { - switch name { - case "write_file", "patch_file", "delete_file": - return true - default: - return false - } +func openAIResponsesToolDefinition(tool modelToolDefinition, nativeOpenAI bool) (responses.ToolUnionParam, error) { + return openAIResponsesToolDefinitionWithFreeForm(tool, nativeOpenAI, !nativeOpenAI) } -func openAIResponsesToolDefinition(tool modelToolDefinition, nativeOpenAI bool) (responses.ToolUnionParam, error) { +func openAIResponsesToolDefinitionWithFreeForm( + tool modelToolDefinition, + nativeOpenAI bool, + freeFormTool bool, +) (responses.ToolUnionParam, error) { toolType, err := providerToolType(tool) if err != nil { return responses.ToolUnionParam{}, err @@ -286,7 +298,7 @@ func openAIResponsesToolDefinition(tool modelToolDefinition, nativeOpenAI bool) }}, nil } if toolType == agentloop.ToolTypeApplyPatch { - if nativeOpenAI { + if !freeFormTool { return responses.ToolUnionParam{OfApplyPatch: &responses.ApplyPatchToolParam{}}, nil } return responses.ToolUnionParam{OfCustom: &responses.CustomToolParam{ @@ -310,13 +322,22 @@ func openAIResponsesToolDefinition(tool modelToolDefinition, nativeOpenAI bool) return responses.ToolUnionParam{OfFunction: &converted}, nil } -func openAIResponsesMessage(message conversation.Message, nativeOpenAI bool) (responses.ResponseInputParam, error) { - return openAIResponsesMessageWithNativeCallIDs(message, nativeOpenAI, nil) +func openAIResponsesMessage( + message conversation.Message, + nativeOpenAI bool, + freeFormTool ...bool, +) (responses.ResponseInputParam, error) { + useFreeFormTool := !nativeOpenAI + if len(freeFormTool) > 0 { + useFreeFormTool = freeFormTool[0] + } + return openAIResponsesMessageWithNativeCallIDs(message, nativeOpenAI, useFreeFormTool, nil) } func openAIResponsesMessageWithNativeCallIDs( message conversation.Message, nativeOpenAI bool, + freeFormTool bool, callSources map[string]openAIResponsesCallSource, ) (responses.ResponseInputParam, error) { role := responses.EasyInputMessageRole(message.Role) @@ -400,7 +421,7 @@ func openAIResponsesMessageWithNativeCallIDs( return nil, unsupportedContent("OpenAI Responses apply patch call requires assistant role") } flush() - item, err := openAIResponsesApplyPatchCall(value, nativeOpenAI) + item, err := openAIResponsesApplyPatchCall(value, nativeOpenAI && !freeFormTool) if err != nil { return nil, err } @@ -412,7 +433,7 @@ func openAIResponsesMessageWithNativeCallIDs( if err != nil { return nil, err } - if source == conversation.ApplyPatchSourceNative && nativeOpenAI { + if source == conversation.ApplyPatchSourceNative && nativeOpenAI && !freeFormTool { status := "completed" if value.IsError { status = "failed" diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister.go b/packages/agenty-core/pkg/infra/modelcatalog/lister.go new file mode 100644 index 0000000..09c26d3 --- /dev/null +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister.go @@ -0,0 +1,487 @@ +// Package modelcatalog retrieves and normalizes provider model listings. +package modelcatalog + +import ( + "context" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + json "github.com/bytedance/sonic" + + "github.com/masteryyh/agenty-core/pkg/domain/catalog" + "github.com/masteryyh/agenty-core/pkg/domain/shared" +) + +const ( + defaultAnthropicVersion = "2023-06-01" + defaultModelsURL = "models" + defaultAnthropicModels = "v1/models" + maxResponseBytes = 16 << 20 + maxPaginationPages = 100 +) + +type httpDoer interface { + Do(*http.Request) (*http.Response, error) +} + +// Lister calls a provider's configured model endpoint and adapts its response +// to the provider-neutral AvailableModel shape. +type Lister struct { + client httpDoer +} + +// NewLister creates a model lister. A bounded timeout is used when no client is +// supplied so an unavailable provider cannot block the core indefinitely. +func NewLister(client *http.Client) *Lister { + if client == nil { + client = &http.Client{Timeout: 30 * time.Second} + } + return &Lister{client: client} +} + +var defaultLister = NewLister(nil) + +// List retrieves and normalizes the model catalog using the built-in HTTP +// client. ProviderService uses this entry point so callers do not need to +// construct or inject an infrastructure dependency. +func List(ctx context.Context, provider catalog.Provider) ([]catalog.AvailableModel, error) { + return defaultLister.List(ctx, provider) +} + +// List retrieves all available models for provider. Provider-specific pages are +// followed transparently until the upstream indicates that the list is done. +func (l *Lister) List(ctx context.Context, provider catalog.Provider) ([]catalog.AvailableModel, error) { + if l == nil || l.client == nil { + return nil, fmt.Errorf("model catalog HTTP client is not configured") + } + + switch provider.Type { + case catalog.APIOpenAI, catalog.APIOpenAICompletions: + return l.listOpenAICompatible(ctx, provider) + case catalog.APIAnthropic: + return l.listAnthropic(ctx, provider) + case catalog.APIGemini: + return l.listGemini(ctx, provider) + default: + return nil, fmt.Errorf("unsupported provider API type %q", provider.Type) + } +} + +type anthropicListResponse struct { + Data []anthropicModel `json:"data"` + HasMore bool `json:"has_more"` + LastID string `json:"last_id"` +} + +type anthropicModel struct { + ID string `json:"id"` + DisplayName string `json:"display_name"` + MaxInputTokens int `json:"max_input_tokens"` + MaxTokens int64 `json:"max_tokens"` + Capabilities *anthropicCapabilities `json:"capabilities"` +} + +type capabilitySupport struct { + Supported bool `json:"supported"` +} + +type anthropicCapabilities struct { + ImageInput *capabilitySupport `json:"image_input"` + PDFInput *capabilitySupport `json:"pdf_input"` + Thinking *capabilitySupport `json:"thinking"` + Effort *anthropicEffort `json:"effort"` +} + +type anthropicEffort struct { + Supported bool `json:"supported"` + Low capabilitySupport `json:"low"` + Medium capabilitySupport `json:"medium"` + High capabilitySupport `json:"high"` + XHigh capabilitySupport `json:"xhigh"` + Max capabilitySupport `json:"max"` +} + +func (l *Lister) listAnthropic(ctx context.Context, provider catalog.Provider) ([]catalog.AvailableModel, error) { + models := make([]catalog.AvailableModel, 0) + seenCursors := make(map[string]struct{}) + var afterID string + + for range maxPaginationPages { + query := url.Values{} + if afterID != "" { + query.Set("after_id", afterID) + } + + var response anthropicListResponse + if err := l.getJSON(ctx, provider, query, &response); err != nil { + return nil, err + } + + for index, item := range response.Data { + efforts := anthropicReasoningEfforts(item.Capabilities) + multiModal := item.Capabilities != nil && + ((item.Capabilities.ImageInput != nil && item.Capabilities.ImageInput.Supported) || + (item.Capabilities.PDFInput != nil && item.Capabilities.PDFInput.Supported)) + model, err := normalizeModel( + provider, + len(models)+index, + item.ID, + item.DisplayName, + item.MaxInputTokens, + item.MaxTokens, + multiModal, + efforts, + ) + if err != nil { + return nil, err + } + models = append(models, model) + } + + if !response.HasMore { + return models, nil + } + if response.LastID == "" { + return nil, fmt.Errorf("provider %q returned has_more without last_id", provider.Code) + } + if _, ok := seenCursors[response.LastID]; ok { + return nil, fmt.Errorf("provider %q returned a repeated pagination cursor", provider.Code) + } + seenCursors[response.LastID] = struct{}{} + afterID = response.LastID + } + + return nil, fmt.Errorf("provider %q exceeded model pagination limit", provider.Code) +} + +func anthropicReasoningEfforts(capabilities *anthropicCapabilities) []shared.ReasoningEffort { + if capabilities == nil { + return []shared.ReasoningEffort{} + } + if capabilities.Effort != nil { + effort := capabilities.Effort + supported := make([]shared.ReasoningEffort, 0, len(shared.StandardReasoningEfforts())) + if effort.Low.Supported { + supported = append(supported, shared.ReasoningLow) + } + if effort.Medium.Supported { + supported = append(supported, shared.ReasoningMedium) + } + if effort.High.Supported { + supported = append(supported, shared.ReasoningHigh) + } + if effort.XHigh.Supported { + supported = append(supported, shared.ReasoningXHigh) + } + if effort.Max.Supported { + supported = append(supported, shared.ReasoningMax) + } + if len(supported) > 0 { + return supported + } + if effort.Supported { + return shared.StandardReasoningEfforts() + } + return []shared.ReasoningEffort{} + } + if capabilities.Thinking != nil && capabilities.Thinking.Supported { + return shared.StandardReasoningEfforts() + } + return []shared.ReasoningEffort{} +} + +type openRouterResponse struct { + Data []openRouterModel `json:"data"` +} + +type openRouterModel struct { + ID string `json:"id"` + Name string `json:"name"` + ContextLength int `json:"context_length"` + MaxCompletionTokens int64 `json:"max_completion_tokens"` + Architecture openRouterArchitecture `json:"architecture"` + TopProvider openRouterTopProvider `json:"top_provider"` + Reasoning *openRouterReasoning `json:"reasoning"` +} + +type openRouterArchitecture struct { + InputModalities []string `json:"input_modalities"` +} + +type openRouterTopProvider struct { + ContextLength int `json:"context_length"` + MaxCompletionTokens int64 `json:"max_completion_tokens"` +} + +type openRouterReasoning struct { + SupportedEfforts []string `json:"supported_efforts"` +} + +func (l *Lister) listOpenAICompatible(ctx context.Context, provider catalog.Provider) ([]catalog.AvailableModel, error) { + var response openRouterResponse + if err := l.getJSON(ctx, provider, nil, &response); err != nil { + return nil, err + } + + models := make([]catalog.AvailableModel, 0, len(response.Data)) + for index, item := range response.Data { + contextWindow := item.ContextLength + if contextWindow <= 0 { + contextWindow = item.TopProvider.ContextLength + } + maxOutputTokens := item.MaxCompletionTokens + if maxOutputTokens <= 0 { + maxOutputTokens = item.TopProvider.MaxCompletionTokens + } + var efforts []shared.ReasoningEffort + if item.Reasoning != nil { + if item.Reasoning.SupportedEfforts == nil { + efforts = shared.StandardReasoningEfforts() + } else { + efforts = parseReasoningEfforts(item.Reasoning.SupportedEfforts) + } + } else { + efforts = []shared.ReasoningEffort{} + } + model, err := normalizeModel( + provider, + index, + item.ID, + item.Name, + contextWindow, + maxOutputTokens, + containsNonTextInput(item.Architecture.InputModalities), + efforts, + ) + if err != nil { + return nil, err + } + models = append(models, model) + } + return models, nil +} + +type geminiListResponse struct { + Models []geminiModel `json:"models"` + NextPageToken string `json:"nextPageToken"` +} + +type geminiModel struct { + Name string `json:"name"` + BaseModelID string `json:"baseModelId"` + DisplayName string `json:"displayName"` + InputTokenLimit int `json:"inputTokenLimit"` + OutputTokenLimit int64 `json:"outputTokenLimit"` + SupportedGenerationMethods []string `json:"supportedGenerationMethods"` + Thinking *bool `json:"thinking"` +} + +func (l *Lister) listGemini(ctx context.Context, provider catalog.Provider) ([]catalog.AvailableModel, error) { + models := make([]catalog.AvailableModel, 0) + seenTokens := make(map[string]struct{}) + var pageToken string + + for range maxPaginationPages { + query := url.Values{} + if pageToken != "" { + query.Set("pageToken", pageToken) + } + + var response geminiListResponse + if err := l.getJSON(ctx, provider, query, &response); err != nil { + return nil, err + } + + for index, item := range response.Models { + code := strings.TrimPrefix(item.BaseModelID, "models/") + if code == "" { + code = strings.TrimPrefix(item.Name, "models/") + } + efforts := []shared.ReasoningEffort{} + if item.Thinking != nil && *item.Thinking { + efforts = shared.StandardReasoningEfforts() + } + model, err := normalizeModel( + provider, + len(models)+index, + code, + item.DisplayName, + item.InputTokenLimit, + item.OutputTokenLimit, + false, + efforts, + ) + if err != nil { + return nil, err + } + models = append(models, model) + } + + if response.NextPageToken == "" { + return models, nil + } + if _, ok := seenTokens[response.NextPageToken]; ok { + return nil, fmt.Errorf("provider %q returned a repeated pagination token", provider.Code) + } + seenTokens[response.NextPageToken] = struct{}{} + pageToken = response.NextPageToken + } + + return nil, fmt.Errorf("provider %q exceeded model pagination limit", provider.Code) +} + +func (l *Lister) getJSON(ctx context.Context, provider catalog.Provider, query url.Values, target any) error { + endpoint, err := endpointURL(provider) + if err != nil { + return err + } + if len(query) > 0 { + values := endpoint.Query() + for key, items := range query { + values.Del(key) + for _, item := range items { + values.Add(key, item) + } + } + endpoint.RawQuery = values.Encode() + } + + request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil) + if err != nil { + return fmt.Errorf("build model list request: %w", err) + } + request.Header.Set("Accept", "application/json") + switch provider.Type { + case catalog.APIAnthropic: + request.Header.Set("x-api-key", provider.APIKey) + request.Header.Set("anthropic-version", defaultAnthropicVersion) + case catalog.APIGemini: + values := request.URL.Query() + values.Set("key", provider.APIKey) + request.URL.RawQuery = values.Encode() + default: + request.Header.Set("Authorization", "Bearer "+provider.APIKey) + } + + response, err := l.client.Do(request) + if err != nil { + return fmt.Errorf("request provider model list: %w", err) + } + defer response.Body.Close() + + body, err := io.ReadAll(io.LimitReader(response.Body, maxResponseBytes+1)) + if err != nil { + return fmt.Errorf("read provider model list response: %w", err) + } + if int64(len(body)) > maxResponseBytes { + return fmt.Errorf("provider model list response exceeds %d bytes", maxResponseBytes) + } + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { + return fmt.Errorf("provider model list returned HTTP %d", response.StatusCode) + } + if err := json.Unmarshal(body, target); err != nil { + return fmt.Errorf("decode provider model list response: %w", err) + } + return nil +} + +func endpointURL(provider catalog.Provider) (*url.URL, error) { + baseURL, err := url.Parse(strings.TrimSpace(provider.BaseURL)) + if err != nil { + return nil, fmt.Errorf("invalid provider base URL: %w", err) + } + if baseURL.Scheme != "http" && baseURL.Scheme != "https" || baseURL.Host == "" { + return nil, fmt.Errorf("provider base URL must be an absolute HTTP(S) URL") + } + + modelsURL := strings.TrimSpace(provider.ModelsURL) + if modelsURL == "" { + modelsURL = defaultModelsURL + if provider.Type == catalog.APIAnthropic { + modelsURL = defaultAnthropicModels + } + } + relative, err := url.Parse(modelsURL) + if err != nil { + return nil, fmt.Errorf("invalid provider models URL: %w", err) + } + if relative.IsAbs() || relative.Host != "" || strings.HasPrefix(relative.Path, "//") { + return nil, fmt.Errorf("provider models URL must be relative to the base URL") + } + + baseURL.Path = strings.TrimRight(baseURL.Path, "/") + "/" + strings.TrimLeft(relative.Path, "/") + baseURL.RawPath = "" + if relative.RawQuery != "" { + baseURL.RawQuery = relative.RawQuery + } + baseURL.Fragment = "" + return baseURL, nil +} + +func normalizeModel( + provider catalog.Provider, + index int, + code string, + name string, + contextWindow int, + maxOutputTokens int64, + multiModal bool, + reasoningEfforts []shared.ReasoningEffort, +) (catalog.AvailableModel, error) { + code = strings.TrimSpace(code) + modelCode, err := shared.NewModelCode(code) + if err != nil { + return catalog.AvailableModel{}, fmt.Errorf("provider %q model %d has invalid id: %w", provider.Code, index, err) + } + name = strings.TrimSpace(name) + if name == "" { + name = code + } + if contextWindow <= 0 { + contextWindow = catalog.DefaultAvailableModelContextWindow + } + if maxOutputTokens <= 0 { + maxOutputTokens = catalog.DefaultAvailableModelMaxOutputTokens + } + if reasoningEfforts == nil { + reasoningEfforts = []shared.ReasoningEffort{} + } + return catalog.AvailableModel{ + Code: modelCode, + Name: name, + ContextWindow: contextWindow, + MaxOutputTokens: maxOutputTokens, + MultiModal: multiModal, + ReasoningEfforts: reasoningEfforts, + }, nil +} + +func parseReasoningEfforts(values []string) []shared.ReasoningEffort { + available := make(map[shared.ReasoningEffort]struct{}, len(values)) + for _, value := range values { + effort := shared.ReasoningEffort(strings.ToLower(strings.TrimSpace(value))) + if effort.Valid() && effort.Enabled() { + available[effort] = struct{}{} + } + } + result := make([]shared.ReasoningEffort, 0, len(available)) + for _, effort := range shared.StandardReasoningEfforts() { + if _, ok := available[effort]; ok { + result = append(result, effort) + } + } + return result +} + +func containsNonTextInput(modalities []string) bool { + for _, modality := range modalities { + if strings.TrimSpace(strings.ToLower(modality)) != "" && strings.ToLower(modality) != "text" { + return true + } + } + return false +} diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go new file mode 100644 index 0000000..27b940a --- /dev/null +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go @@ -0,0 +1,229 @@ +package modelcatalog + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "reflect" + "testing" + + "github.com/masteryyh/agenty-core/pkg/domain/catalog" + "github.com/masteryyh/agenty-core/pkg/domain/shared" +) + +func TestListerOpenAICompatibleDefaultsAndMissingReasoning(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/models" { + t.Errorf("path = %q, want /v1/models", r.URL.Path) + } + if got := r.Header.Get("Authorization"); got != "Bearer openai-secret" { + t.Errorf("authorization = %q", got) + } + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"gpt-test"}]}`)) + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "openai", + Type: catalog.APIOpenAI, + BaseURL: server.URL + "/v1", + APIKey: "openai-secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + want := []catalog.AvailableModel{{ + Code: "gpt-test", + Name: "gpt-test", + ContextWindow: catalog.DefaultAvailableModelContextWindow, + MaxOutputTokens: catalog.DefaultAvailableModelMaxOutputTokens, + ReasoningEfforts: []shared.ReasoningEffort{}, + }} + if !reflect.DeepEqual(models, want) { + t.Fatalf("models = %#v, want %#v", models, want) + } +} + +func TestListerDeepSeekUsesOpenAICompatibleShape(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/models" { + t.Errorf("path = %q, want /models", r.URL.Path) + } + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"deepseek-chat","owned_by":"deepseek"}]}`)) + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "deepseek", + Type: catalog.APIOpenAICompletions, + BaseURL: server.URL, + APIKey: "deepseek-secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(models) != 1 || models[0].Code != "deepseek-chat" || models[0].ReasoningEfforts == nil || len(models[0].ReasoningEfforts) != 0 { + t.Fatalf("models = %#v", models) + } +} + +func TestListerAnthropicPaginationAndCapabilities(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got := r.Header.Get("x-api-key"); got != "anthropic-secret" { + t.Errorf("x-api-key = %q", got) + } + if got := r.Header.Get("anthropic-version"); got != defaultAnthropicVersion { + t.Errorf("anthropic-version = %q", got) + } + if r.URL.Query().Get("after_id") == "" { + _, _ = w.Write([]byte(`{ + "data":[{ + "id":"claude-opus", + "display_name":"", + "max_input_tokens":200000, + "max_tokens":64000, + "capabilities":{"image_input":{"supported":true},"effort":{"low":{"supported":true},"max":{"supported":true}}} + }], + "has_more":true, + "last_id":"claude-opus" + }`)) + return + } + if got := r.URL.Query().Get("after_id"); got != "claude-opus" { + t.Errorf("after_id = %q", got) + } + _, _ = w.Write([]byte(`{ + "data":[{"id":"claude-haiku","max_input_tokens":0,"max_tokens":0}], + "has_more":false, + "last_id":"claude-haiku" + }`)) + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "anthropic", + Type: catalog.APIAnthropic, + BaseURL: server.URL, + APIKey: "anthropic-secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(models) != 2 { + t.Fatalf("models = %#v", models) + } + if models[0].Name != "claude-opus" || models[0].ContextWindow != 200000 || models[0].MaxOutputTokens != 64000 || !models[0].MultiModal { + t.Errorf("first model = %#v", models[0]) + } + if !reflect.DeepEqual(models[0].ReasoningEfforts, []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningMax}) { + t.Errorf("first reasoning efforts = %#v", models[0].ReasoningEfforts) + } + if models[1].Name != "claude-haiku" || models[1].ContextWindow != catalog.DefaultAvailableModelContextWindow || models[1].MaxOutputTokens != catalog.DefaultAvailableModelMaxOutputTokens || len(models[1].ReasoningEfforts) != 0 { + t.Errorf("second model = %#v", models[1]) + } +} + +func TestListerOpenRouterFieldsAndReasoning(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + payload := map[string]any{ + "data": []map[string]any{ + { + "id": "openai/gpt-test", + "name": "", + "context_length": 128000, + "architecture": map[string]any{"input_modalities": []string{"text", "image"}}, + "top_provider": map[string]any{"max_completion_tokens": 32768}, + "reasoning": map[string]any{"supported_efforts": []string{"high", "minimal", "low"}}, + }, + {"id": "plain-model"}, + }, + } + if err := json.NewEncoder(w).Encode(payload); err != nil { + t.Errorf("encode response: %v", err) + } + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "openrouter", + Type: catalog.APIOpenAICompletions, + BaseURL: server.URL + "/api/v1", + APIKey: "openrouter-secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(models) != 2 { + t.Fatalf("models = %#v", models) + } + if models[0].Name != "openai/gpt-test" || models[0].ContextWindow != 128000 || models[0].MaxOutputTokens != 32768 || !models[0].MultiModal { + t.Errorf("first model = %#v", models[0]) + } + if !reflect.DeepEqual(models[0].ReasoningEfforts, []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}) { + t.Errorf("first reasoning efforts = %#v", models[0].ReasoningEfforts) + } + if models[1].ContextWindow != catalog.DefaultAvailableModelContextWindow || models[1].MaxOutputTokens != catalog.DefaultAvailableModelMaxOutputTokens || len(models[1].ReasoningEfforts) != 0 { + t.Errorf("second model = %#v", models[1]) + } +} + +func TestListerOpenRouterNullReasoningMeansAllStandardEfforts(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"data":[{"id":"reasoning-model","reasoning":{"supported_efforts":null}}]}`)) + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "openrouter", + Type: catalog.APIOpenAI, + BaseURL: server.URL, + APIKey: "secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + if !reflect.DeepEqual(models[0].ReasoningEfforts, shared.StandardReasoningEfforts()) { + t.Fatalf("reasoning efforts = %#v", models[0].ReasoningEfforts) + } +} + +func TestListerGeminiPaginationAndThinking(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got := r.URL.Query().Get("key"); got != "gemini-secret" { + t.Errorf("key = %q", got) + } + if r.URL.Query().Get("pageToken") == "" { + _, _ = w.Write([]byte(`{"models":[{"name":"models/gemini-3-flash","displayName":"Gemini 3 Flash","inputTokenLimit":128000,"outputTokenLimit":8192,"thinking":true}],"nextPageToken":"next"}`)) + return + } + _, _ = w.Write([]byte(`{"models":[{"name":"models/gemini-3-pro"}],"nextPageToken":""}`)) + })) + defer server.Close() + + models, err := NewLister(server.Client()).List(t.Context(), catalog.Provider{ + Code: "google", + Type: catalog.APIGemini, + BaseURL: server.URL, + APIKey: "gemini-secret", + }) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(models) != 2 || models[0].Code != "gemini-3-flash" || models[0].Name != "Gemini 3 Flash" || len(models[0].ReasoningEfforts) != len(shared.StandardReasoningEfforts()) { + t.Fatalf("models = %#v", models) + } + if models[1].Name != "gemini-3-pro" || len(models[1].ReasoningEfforts) != 0 { + t.Fatalf("second model = %#v", models[1]) + } +} + +func TestEndpointURLRejectsAbsoluteModelsURL(t *testing.T) { + _, err := endpointURL(catalog.Provider{ + Type: catalog.APIOpenAI, + BaseURL: "https://example.com/v1", + ModelsURL: "https://other.example/models", + }) + if err == nil { + t.Fatal("endpointURL error = nil, want error") + } +} diff --git a/packages/agenty-core/pkg/infra/rpc/adapter/provider.go b/packages/agenty-core/pkg/infra/rpc/adapter/provider.go index 70cde6c..eb7e882 100644 --- a/packages/agenty-core/pkg/infra/rpc/adapter/provider.go +++ b/packages/agenty-core/pkg/infra/rpc/adapter/provider.go @@ -13,6 +13,7 @@ func RegisterProviderHandlers(d *rpc.Dispatcher, svc *application.ProviderServic d.Register("provider.create", providerCreate(svc)) d.Register("provider.get", providerGet(svc)) d.Register("provider.list", providerList(svc)) + d.Register("provider.listModels", providerListModels(svc)) d.Register("provider.update", providerUpdate(svc)) d.Register("provider.delete", providerDelete(svc)) d.Register("provider.addModel", providerAddModel(svc)) @@ -24,6 +25,10 @@ type providerCreateParams struct { application.ProviderInput } +type providerListParams struct { + ProviderCode string `json:"providerCode,omitempty"` +} + func providerCreate(svc *application.ProviderService) rpc.Handler { return func(ctx context.Context, params json.RawMessage) (any, error) { var p providerCreateParams @@ -46,11 +51,25 @@ func providerGet(svc *application.ProviderService) rpc.Handler { func providerList(svc *application.ProviderService) rpc.Handler { return func(ctx context.Context, params json.RawMessage) (any, error) { - var p struct{} + var p providerListParams + if err := decodeParams(params, &p); err != nil { + return nil, rpc.InvalidParams("invalid params: " + err.Error()) + } + return wrap(svc.List(ctx, p.ProviderCode)) + } +} + +type providerListModelsParams struct { + ProviderCode string `json:"providerCode"` +} + +func providerListModels(svc *application.ProviderService) rpc.Handler { + return func(ctx context.Context, params json.RawMessage) (any, error) { + var p providerListModelsParams if err := decodeParams(params, &p); err != nil { return nil, rpc.InvalidParams("invalid params: " + err.Error()) } - return wrap(svc.List(ctx)) + return wrap(svc.ListModels(ctx, p.ProviderCode)) } } diff --git a/packages/agenty-core/pkg/infra/storage/catalog.go b/packages/agenty-core/pkg/infra/storage/catalog.go index 6d3843a..c9f82d3 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog.go +++ b/packages/agenty-core/pkg/infra/storage/catalog.go @@ -3,9 +3,13 @@ package storage import ( "context" "fmt" + "maps" "os" "path/filepath" + "slices" "strings" + "sync" + "time" json "github.com/bytedance/sonic" @@ -16,14 +20,53 @@ import ( var ErrProviderNotFound = catalog.ErrProviderNotFound type CatalogRepository struct { - providersDir string + providersDir string + builtinProviders map[shared.Code]*catalog.Provider + builtinOrder []shared.Code + cacheMu sync.RWMutex } -func NewCatalogRepository(providersDir string) *CatalogRepository { - return &CatalogRepository{providersDir: providersDir} +type modelDiscoveryCache struct { + ExpiresAt time.Time `json:"expiresAt"` + Models []catalog.Model `json:"models"` +} + +func NewCatalogRepository(providersDir string, builtinProviders ...*catalog.Provider) *CatalogRepository { + builtins := make(map[shared.Code]*catalog.Provider, len(builtinProviders)) + order := make([]shared.Code, 0, len(builtinProviders)) + for _, provider := range builtinProviders { + if provider == nil || !provider.Code.Valid() { + continue + } + copy := cloneProvider(provider) + copy.Builtin = true + builtins[copy.Code] = copy + order = append(order, copy.Code) + } + return &CatalogRepository{providersDir: providersDir, builtinProviders: builtins, builtinOrder: order} } func (r *CatalogRepository) Get(_ context.Context, code shared.Code) (*catalog.Provider, error) { + r.cacheMu.RLock() + defer r.cacheMu.RUnlock() + + return r.getLocked(code, true) +} + +func (r *CatalogRepository) getLocked(code shared.Code, includeDiscoveryCache bool) (*catalog.Provider, error) { + if builtin, ok := r.builtinProviders[code]; ok { + provider := cloneProvider(builtin) + apiKey, err := r.readAPIKey(code) + if err != nil { + return nil, err + } + provider.APIKey = apiKey + if includeDiscoveryCache { + r.applyModelDiscoveryCache(provider) + } + return provider, nil + } + providerPath := filepath.Join(r.providersDir, code.String()+".json") data, err := os.ReadFile(providerPath) if err != nil { @@ -37,13 +80,47 @@ func (r *CatalogRepository) Get(_ context.Context, code shared.Code) (*catalog.P if err := json.Unmarshal(data, &provider); err != nil { return nil, err } + provider.ModelsCached = false normalizeModels(&provider) + if includeDiscoveryCache { + r.applyModelDiscoveryCache(&provider) + } return &provider, nil } +// NeedsModelDiscovery reports whether a provider with no embedded/persisted +// models has no fresh discovery cache. Expired cache entries remain available +// through Get so callers can continue using the last known model metadata while +// a subsequent list operation refreshes it. +func (r *CatalogRepository) NeedsModelDiscovery(_ context.Context, code shared.Code) (bool, error) { + r.cacheMu.RLock() + defer r.cacheMu.RUnlock() + + provider, err := r.getLocked(code, false) + if err != nil { + return false, err + } + if len(provider.Models) > 0 { + return false, nil + } + + cache, err := r.readModelDiscoveryCache(code) + if err != nil { + return true, nil + } + return cache == nil || cache.ExpiresAt.IsZero() || !time.Now().UTC().Before(cache.ExpiresAt), nil +} + func (r *CatalogRepository) List(ctx context.Context) ([]*catalog.Provider, error) { - providers := make([]*catalog.Provider, 0) + providers := make([]*catalog.Provider, 0, len(r.builtinProviders)) + for _, code := range r.builtinOrder { + provider, err := r.Get(ctx, code) + if err != nil { + return nil, err + } + providers = append(providers, provider) + } entries, err := os.ReadDir(r.providersDir) if err != nil { if os.IsNotExist(err) { @@ -61,6 +138,9 @@ func (r *CatalogRepository) List(ctx context.Context) ([]*catalog.Provider, erro if err != nil { continue } + if _, builtin := r.builtinProviders[code]; builtin { + continue + } provider, err := r.Get(ctx, code) if err != nil { @@ -79,10 +159,20 @@ func (r *CatalogRepository) Save(_ context.Context, provider *catalog.Provider) if provider == nil || !provider.Code.Valid() { return fmt.Errorf("storage: invalid provider code") } + r.cacheMu.Lock() + defer r.cacheMu.Unlock() + if err := r.invalidateModelDiscoveryCacheLocked(provider.Code); err != nil { + return err + } + if _, builtin := r.builtinProviders[provider.Code]; builtin { + return r.saveAPIKey(provider.Code, provider.APIKey) + } if err := os.MkdirAll(r.providersDir, 0700); err != nil { return err } + // ModelsCached is a transient RPC hint, not part of provider config. + provider.ModelsCached = false normalizeModels(provider) providerData, err := json.MarshalIndent(provider, "", " ") if err != nil { @@ -94,6 +184,11 @@ func (r *CatalogRepository) Save(_ context.Context, provider *catalog.Provider) } func (r *CatalogRepository) Delete(_ context.Context, code shared.Code) error { + r.cacheMu.Lock() + defer r.cacheMu.Unlock() + if _, builtin := r.builtinProviders[code]; builtin { + return catalog.ErrBuiltinProviderReadOnly + } providerPath := filepath.Join(r.providersDir, code.String()+".json") if err := os.Remove(providerPath); err != nil { if os.IsNotExist(err) { @@ -101,6 +196,117 @@ func (r *CatalogRepository) Delete(_ context.Context, code shared.Code) error { } return err } + if err := r.invalidateModelDiscoveryCacheLocked(code); err != nil { + return err + } + return nil +} + +// ReplaceModels stores the latest discovered model list without changing the +// provider's hand-authored configuration. The list becomes visible through +// Get/List until its expiration time. +func (r *CatalogRepository) ReplaceModels( + _ context.Context, + code shared.Code, + models []catalog.Model, + expiresAt time.Time, +) error { + if !code.Valid() { + return fmt.Errorf("storage: invalid provider code") + } + if expiresAt.IsZero() { + return fmt.Errorf("storage: model discovery cache expiration is required") + } + + r.cacheMu.Lock() + defer r.cacheMu.Unlock() + if _, err := r.getLocked(code, false); err != nil { + return err + } + + normalized := slices.Clone(models) + for index := range normalized { + normalized[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( + normalized[index].ReasoningEfforts, + ) + if normalized[index].MaxOutputTokens <= 0 { + normalized[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens + } + } + if normalized == nil { + normalized = make([]catalog.Model, 0) + } + + data, err := json.MarshalIndent(modelDiscoveryCache{ + ExpiresAt: expiresAt.UTC(), + Models: normalized, + }, "", " ") + if err != nil { + return fmt.Errorf("storage: encode model discovery cache: %w", err) + } + cacheDir := filepath.Join(r.providersDir, ".models") + if err := os.MkdirAll(cacheDir, 0700); err != nil { + return fmt.Errorf("storage: create model discovery cache directory: %w", err) + } + cachePath := filepath.Join(cacheDir, code.String()+".json") + if err := os.WriteFile(cachePath, data, 0600); err != nil { + return fmt.Errorf("storage: write model discovery cache: %w", err) + } + return nil +} + +func (r *CatalogRepository) applyModelDiscoveryCache(provider *catalog.Provider) { + if len(provider.Models) > 0 { + return + } + cache, err := r.readModelDiscoveryCache(provider.Code) + if err != nil { + return + } + if cache == nil { + return + } + provider.Models = cache.Models + provider.ModelsCached = true +} + +func (r *CatalogRepository) readModelDiscoveryCache(code shared.Code) (*modelDiscoveryCache, error) { + data, err := os.ReadFile(r.modelDiscoveryCachePath(code)) + if err != nil { + return nil, err + } + + var cache modelDiscoveryCache + if err := json.Unmarshal(data, &cache); err != nil { + return nil, err + } + normalizeModelsForCache(&cache) + return &cache, nil +} + +func normalizeModelsForCache(cache *modelDiscoveryCache) { + if cache.Models == nil { + cache.Models = make([]catalog.Model, 0) + } + for index := range cache.Models { + cache.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( + cache.Models[index].ReasoningEfforts, + ) + if cache.Models[index].MaxOutputTokens <= 0 { + cache.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens + } + } +} + +func (r *CatalogRepository) modelDiscoveryCachePath(code shared.Code) string { + return filepath.Join(r.providersDir, ".models", code.String()+".json") +} + +func (r *CatalogRepository) invalidateModelDiscoveryCacheLocked(code shared.Code) error { + err := os.Remove(r.modelDiscoveryCachePath(code)) + if err != nil && !os.IsNotExist(err) { + return fmt.Errorf("storage: remove model discovery cache: %w", err) + } return nil } @@ -109,6 +315,57 @@ func normalizeModels(provider *catalog.Provider) { provider.Models = make([]catalog.Model, 0) } for index := range provider.Models { - provider.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens + provider.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( + provider.Models[index].ReasoningEfforts, + ) + if provider.Models[index].MaxOutputTokens <= 0 { + provider.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens + } + } +} + +type providerCredentials struct { + APIKey string `json:"apiKey"` +} + +func (r *CatalogRepository) readAPIKey(code shared.Code) (string, error) { + providerPath := filepath.Join(r.providersDir, code.String()+".json") + data, err := os.ReadFile(providerPath) + if err != nil { + if os.IsNotExist(err) { + return "", nil + } + return "", err + } + + var credentials providerCredentials + if err := json.Unmarshal(data, &credentials); err != nil { + return "", err + } + return credentials.APIKey, nil +} + +func (r *CatalogRepository) saveAPIKey(code shared.Code, apiKey string) error { + if err := os.MkdirAll(r.providersDir, 0700); err != nil { + return err + } + data, err := json.MarshalIndent(providerCredentials{APIKey: apiKey}, "", " ") + if err != nil { + return err + } + providerPath := filepath.Join(r.providersDir, code.String()+".json") + return os.WriteFile(providerPath, data, 0600) +} + +func cloneProvider(provider *catalog.Provider) *catalog.Provider { + copy := *provider + copy.Models = make([]catalog.Model, len(provider.Models)) + copy.Models = append(copy.Models[:0], provider.Models...) + for index := range copy.Models { + copy.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( + provider.Models[index].ReasoningEfforts, + ) } + copy.Metadata = maps.Clone(provider.Metadata) + return © } diff --git a/packages/agenty-core/pkg/infra/storage/catalog_test.go b/packages/agenty-core/pkg/infra/storage/catalog_test.go index 48964b2..3ea6688 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog_test.go +++ b/packages/agenty-core/pkg/infra/storage/catalog_test.go @@ -11,6 +11,7 @@ import ( "github.com/masteryyh/agenty-core/pkg/domain/catalog" "github.com/masteryyh/agenty-core/pkg/domain/shared" + "github.com/masteryyh/agenty-core/pkg/infra/catalogdata" ) func newCatalogRepo(t *testing.T) *CatalogRepository { @@ -38,28 +39,26 @@ func TestCatalogSaveAndGet(t *testing.T) { provider.APIKey = "sk-ant-test" model1 := catalog.Model{ - Code: mustCatalogModelCode(`org/claude\\claude-opus[fast]`), - Name: "Claude Opus 4.8", - ContextWindow: 200000, - MaxOutputTokens: 32000, - ReasoningEffortMapping: map[string]shared.ReasoningEffort{ - "low": shared.ReasoningLow, - "medium": shared.ReasoningMedium, - "high": shared.ReasoningHigh, - }, - CreatedAt: time.Now().UTC(), - UpdatedAt: time.Now().UTC(), + Code: mustCatalogModelCode(`org/claude\\claude-opus[fast]`), + Name: "Claude Opus 4.8", + ContextWindow: 200000, + MaxOutputTokens: 32000, + ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningMedium, shared.ReasoningHigh}, + CreatedAt: time.Now().UTC(), + UpdatedAt: time.Now().UTC(), } model2 := catalog.Model{ - Code: mustCatalogModelCode("claude-haiku-4-5"), - Name: "Claude Haiku 4.5", - ContextWindow: 200000, - MaxOutputTokens: 8000, - Light: true, - CreatedAt: time.Now().UTC(), - UpdatedAt: time.Now().UTC(), + Code: mustCatalogModelCode("claude-haiku-4-5"), + Name: "Claude Haiku 4.5", + ContextWindow: 200000, + MaxOutputTokens: 8000, + Light: true, + ReasoningEfforts: []shared.ReasoningEffort{}, + CreatedAt: time.Now().UTC(), + UpdatedAt: time.Now().UTC(), } provider.Models = []catalog.Model{model1, model2} + provider.ModelsCached = true if err := repo.Save(ctx, provider); err != nil { t.Fatalf("Save: %v", err) @@ -72,6 +71,9 @@ func TestCatalogSaveAndGet(t *testing.T) { if err := json.Unmarshal(providerData, &persistedProvider); err != nil { t.Fatalf("decode provider file: %v", err) } + if persistedProvider.ModelsCached { + t.Fatal("transient cache marker was persisted in provider config") + } if len(persistedProvider.Models) != 2 { t.Fatalf("persisted %d models, want 2", len(persistedProvider.Models)) } @@ -113,22 +115,142 @@ func TestCatalogSaveAndGet(t *testing.T) { if gotOpus != nil && !gotOpus.SupportsReasoning() { t.Errorf("opus SupportsReasoning = %v, want true", gotOpus.SupportsReasoning()) } - if gotOpus != nil && gotOpus.MaxOutputTokens != catalog.DefaultMaxOutputTokens { - t.Errorf("opus max output tokens = %d, want %d", gotOpus.MaxOutputTokens, catalog.DefaultMaxOutputTokens) + if gotOpus != nil && gotOpus.MaxOutputTokens != model1.MaxOutputTokens { + t.Errorf("opus max output tokens = %d, want %d", gotOpus.MaxOutputTokens, model1.MaxOutputTokens) } if gotHaiku != nil && gotHaiku.SupportsReasoning() { t.Errorf("haiku SupportsReasoning = %v, want false", gotHaiku.SupportsReasoning()) } - if gotOpus != nil { - if effort, ok := gotOpus.MapReasoningEffort("medium"); !ok || effort != shared.ReasoningMedium { - t.Errorf("mapped medium effort = %q, %v; want medium, true", effort, ok) - } + if gotOpus != nil && !gotOpus.SupportsReasoningEffort(shared.ReasoningMedium) { + t.Error("opus does not support medium reasoning effort") } if gotHaiku != nil && !gotHaiku.Light { t.Errorf("haiku Light = %v, want true", gotHaiku.Light) } } +func TestCatalogBuiltinProviderPersistsOnlyAPIKey(t *testing.T) { + repo := newCatalogRepo(t) + builtins, err := catalogdata.LoadProviders() + if err != nil { + t.Fatal(err) + } + repo = NewCatalogRepository(repo.providersDir, builtins...) + ctx := context.Background() + + provider, err := repo.Get(ctx, mustCode("openai_legacy")) + if err != nil { + t.Fatal(err) + } + if !provider.Builtin || provider.Name != "OpenAI (Legacy API)" { + t.Fatalf("builtin provider = %+v", provider) + } + provider.APIKey = "secret" + if err := repo.Save(ctx, provider); err != nil { + t.Fatalf("Save builtin: %v", err) + } + + data, err := os.ReadFile(filepath.Join(repo.providersDir, "openai_legacy.json")) + if err != nil { + t.Fatal(err) + } + if string(data) != "{\n \"apiKey\": \"secret\"\n}" { + t.Fatalf("builtin credentials = %s", data) + } + + loaded, err := repo.Get(ctx, provider.Code) + if err != nil { + t.Fatal(err) + } + if loaded.APIKey != "secret" || loaded.Name != provider.Name || len(loaded.Models) != len(provider.Models) { + t.Fatalf("loaded builtin = %+v", loaded) + } + for _, model := range loaded.Models { + if model.ReasoningEfforts == nil { + t.Errorf("model %s reasoning efforts is nil", model.Code) + } + } + if err := repo.Delete(ctx, provider.Code); err != catalog.ErrBuiltinProviderReadOnly { + t.Fatalf("Delete builtin = %v, want read-only", err) + } +} + +func TestCatalogModelDiscoveryCacheUsesExpirationAndSurvivesRestart(t *testing.T) { + dir := t.TempDir() + builtins, err := catalogdata.LoadProviders() + if err != nil { + t.Fatal(err) + } + repo := NewCatalogRepository(filepath.Join(dir, "providers"), builtins...) + ctx := context.Background() + code := mustCode("openrouter") + models := []catalog.Model{{ + Code: mustCatalogModelCode("openai/gpt-test"), + Name: "GPT Test", + ContextWindow: 128_000, + MaxOutputTokens: 16_384, + ReasoningEfforts: []shared.ReasoningEffort{}, + }} + expiresAt := time.Now().UTC().Add(catalog.ModelDiscoveryCacheTTL) + if err := repo.ReplaceModels(ctx, code, models, expiresAt); err != nil { + t.Fatalf("ReplaceModels: %v", err) + } + + loaded, err := repo.Get(ctx, code) + if err != nil { + t.Fatalf("Get cached provider: %v", err) + } + if len(loaded.Models) != 1 || loaded.Models[0].Code != models[0].Code { + t.Fatalf("cached models = %#v", loaded.Models) + } + if !loaded.ModelsCached { + t.Fatal("cached provider did not expose the transient cache marker") + } + + cachePath := filepath.Join(repo.providersDir, ".models", "openrouter.json") + cacheData, err := os.ReadFile(cachePath) + if err != nil { + t.Fatalf("read cache: %v", err) + } + var cache modelDiscoveryCache + if err := json.Unmarshal(cacheData, &cache); err != nil { + t.Fatalf("decode cache: %v", err) + } + if cache.ExpiresAt.Before(time.Now().UTC()) || len(cache.Models) != 1 { + t.Fatalf("cache = %+v", cache) + } + + restarted := NewCatalogRepository(repo.providersDir, builtins...) + restartedProvider, err := restarted.Get(ctx, code) + if err != nil { + t.Fatalf("Get after restart: %v", err) + } + if len(restartedProvider.Models) != 1 { + t.Fatalf("restarted models = %#v", restartedProvider.Models) + } + if !restartedProvider.ModelsCached { + t.Fatal("restarted provider did not expose the transient cache marker") + } + + if err := repo.ReplaceModels(ctx, code, models, time.Now().UTC().Add(-time.Minute)); err != nil { + t.Fatalf("ReplaceModels expired: %v", err) + } + expired, err := repo.Get(ctx, code) + if err != nil { + t.Fatalf("Get expired provider: %v", err) + } + if len(expired.Models) != 1 { + t.Fatalf("expired models = %#v, want stale cache", expired.Models) + } + needsDiscovery, err := repo.NeedsModelDiscovery(ctx, code) + if err != nil { + t.Fatalf("NeedsModelDiscovery: %v", err) + } + if !needsDiscovery { + t.Fatal("expired cache did not request discovery") + } +} + func TestCatalogList(t *testing.T) { repo := newCatalogRepo(t) ctx := context.Background() diff --git a/packages/agenty-core/test/e2e/agenty_client_test.go b/packages/agenty-core/test/e2e/agenty_client_test.go index 82db530..80c720a 100644 --- a/packages/agenty-core/test/e2e/agenty_client_test.go +++ b/packages/agenty-core/test/e2e/agenty_client_test.go @@ -103,6 +103,15 @@ func (c *agentyClient) ListProviders(ctx context.Context) ([]Provider, error) { ) } +func (c *agentyClient) ListProviderModels(ctx context.Context, providerCode string) ([]AvailableModel, error) { + return callResult[[]AvailableModel]( + ctx, + c.rpc, + "provider.listModels", + map[string]any{"providerCode": providerCode}, + ) +} + func (c *agentyClient) UpdateProvider(ctx context.Context, input ProviderUpdateInput) (Provider, error) { return callResult[Provider]( ctx, diff --git a/packages/agenty-core/test/e2e/contracts_test.go b/packages/agenty-core/test/e2e/contracts_test.go index b51e1f4..a0cd24a 100644 --- a/packages/agenty-core/test/e2e/contracts_test.go +++ b/packages/agenty-core/test/e2e/contracts_test.go @@ -29,6 +29,7 @@ var publicRPCMethods = []string{ "provider.create", "provider.get", "provider.list", + "provider.listModels", "provider.update", "provider.delete", "provider.addModel", @@ -126,24 +127,38 @@ type Agent struct { } type Model struct { - Code string `json:"code"` - Name string `json:"name"` - ContextWindow int `json:"contextWindow"` - MaxOutputTokens int64 `json:"maxOutputTokens"` - MultiModal bool `json:"multiModal"` - Light bool `json:"light"` - ReasoningEffortMapping map[string]string `json:"reasoningEffortMapping"` - IsDefault bool `json:"isDefault"` + Code string `json:"code"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal"` + Light bool `json:"light"` + ReasoningEfforts []string `json:"reasoningEfforts"` + IsDefault bool `json:"isDefault"` +} + +type AvailableModel struct { + Code string `json:"code"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal"` + ReasoningEfforts []string `json:"reasoningEfforts"` } type Provider struct { - Code string `json:"code"` - Name string `json:"name"` - Type string `json:"type"` - BaseURL string `json:"baseUrl"` - APIKey string `json:"apiKey"` - Models []Model `json:"models"` - Metadata map[string]any `json:"metadata"` + Code string `json:"code"` + Name string `json:"name"` + Type string `json:"type"` + BaseURL string `json:"baseUrl"` + APIKey string `json:"apiKey"` + Builtin bool `json:"builtin"` + Official bool `json:"official"` + FreeFormTool bool `json:"freeFormTool"` + ModelsURL string `json:"modelsUrl"` + TokenCountURL string `json:"tokenCountUrl"` + Models []Model `json:"models"` + Metadata map[string]any `json:"metadata"` } type Session struct { @@ -248,33 +263,35 @@ type AgentUpdateInput struct { } type ProviderCreateInput struct { - Code string `json:"code"` - Name string `json:"name"` - Type string `json:"type"` - BaseURL string `json:"baseUrl,omitempty"` - APIKey string `json:"apiKey,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` + Code string `json:"code"` + Name string `json:"name"` + Type string `json:"type"` + BaseURL string `json:"baseUrl,omitempty"` + APIKey string `json:"apiKey,omitempty"` + FreeFormTool bool `json:"freeFormTool,omitempty"` + Metadata map[string]any `json:"metadata,omitempty"` } type ProviderUpdateInput struct { - Code string `json:"code"` - Name *string `json:"name,omitempty"` - Type *string `json:"type,omitempty"` - BaseURL *string `json:"baseUrl,omitempty"` - APIKey *string `json:"apiKey,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` + Code string `json:"code"` + Name *string `json:"name,omitempty"` + Type *string `json:"type,omitempty"` + BaseURL *string `json:"baseUrl,omitempty"` + APIKey *string `json:"apiKey,omitempty"` + FreeFormTool *bool `json:"freeFormTool,omitempty"` + Metadata map[string]any `json:"metadata,omitempty"` } type ModelInput struct { - ProviderCode string `json:"providerCode"` - ModelCode string `json:"modelCode"` - Name string `json:"name"` - ContextWindow int `json:"contextWindow,omitempty"` - MaxOutputTokens int64 `json:"maxOutputTokens"` - MultiModal bool `json:"multiModal,omitempty"` - Light bool `json:"light,omitempty"` - ReasoningEffortMapping map[string]string `json:"reasoningEffortMapping,omitempty"` - IsDefault bool `json:"isDefault,omitempty"` + ProviderCode string `json:"providerCode"` + ModelCode string `json:"modelCode"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow,omitempty"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal,omitempty"` + Light bool `json:"light,omitempty"` + Reasoning *bool `json:"reasoning,omitempty"` + IsDefault bool `json:"isDefault,omitempty"` } type SessionCreateInput struct { diff --git a/packages/agenty-core/test/e2e/journey_test.go b/packages/agenty-core/test/e2e/journey_test.go index 2a54eb0..95fd645 100644 --- a/packages/agenty-core/test/e2e/journey_test.go +++ b/packages/agenty-core/test/e2e/journey_test.go @@ -4,6 +4,7 @@ package e2e_test import ( "context" + "net/http" "testing" ) @@ -11,6 +12,9 @@ func TestClientJourneyCoversPublicRPCSurfaceAcrossRestart(t *testing.T) { t.Parallel() fixture := newProviderFixture(t, func(request providerRequest) providerReply { + if request.Method == http.MethodGet { + return providerReply{Body: `{"object":"list","data":[{"id":"fixture-model"}]}`} + } if request.Call == 3 { return providerReply{WaitForCancel: true} } @@ -122,7 +126,20 @@ func TestClientJourneyCoversPublicRPCSurfaceAcrossRestart(t *testing.T) { } providers, err := first.ListProviders(ctx) requireNoError(t, err) - if len(providers) != 1 || providers[0].Code != "local-openai" { + var legacyProvider *Provider + var localProvider *Provider + for index := range providers { + if providers[index].Code == "openai_legacy" { + legacyProvider = &providers[index] + } + if providers[index].Code == "local-openai" { + localProvider = &providers[index] + } + } + if legacyProvider == nil || !legacyProvider.Builtin || !legacyProvider.Official || legacyProvider.Name != "OpenAI (Legacy API)" || legacyProvider.Type != "openai_completions" { + t.Fatalf("legacy provider = %+v", legacyProvider) + } + if localProvider == nil { t.Fatalf("providers = %+v", providers) } _, err = first.GetProvider(ctx, "local-openai") @@ -135,13 +152,10 @@ func TestClientJourneyCoversPublicRPCSurfaceAcrossRestart(t *testing.T) { ContextWindow: 128_000, MaxOutputTokens: 100_000, MultiModal: true, - ReasoningEffortMapping: map[string]string{ - "high": "high", - }, - IsDefault: true, + IsDefault: true, }) requireNoError(t, err) - if len(provider.Models) != 1 || provider.Models[0].MaxOutputTokens != 8_192 { + if len(provider.Models) != 1 || provider.Models[0].MaxOutputTokens != 100_000 { t.Fatalf("provider models = %+v", provider.Models) } _, err = first.AddModel(ctx, ModelInput{ @@ -257,8 +271,8 @@ func TestClientJourneyCoversPublicRPCSurfaceAcrossRestart(t *testing.T) { fixture.requests, 2, ) - if firstProviderRequest.Body["max_completion_tokens"] != float64(8_192) { - t.Fatalf("max completion tokens = %v, want 8192", firstProviderRequest.Body["max_completion_tokens"]) + if firstProviderRequest.Body["max_completion_tokens"] != float64(100_000) { + t.Fatalf("max completion tokens = %v, want 100000", firstProviderRequest.Body["max_completion_tokens"]) } if providerMessageCount(secondProviderRequest) <= providerMessageCount(firstProviderRequest) { t.Fatalf( @@ -287,6 +301,11 @@ func TestClientJourneyCoversPublicRPCSurfaceAcrossRestart(t *testing.T) { fixture.requests, 3, ) + models, err := second.ListProviderModels(ctx, "local-openai") + requireNoError(t, err) + if len(models) != 1 || models[0].Code != "fixture-model" || models[0].ContextWindow != 256_000 || models[0].MaxOutputTokens != 65_536 || len(models[0].ReasoningEfforts) != 0 { + t.Fatalf("discovered models = %+v", models) + } _, err = second.StartSession( ctx, cancelSession.ID, diff --git a/packages/agenty-core/test/e2e/provider_fixture_test.go b/packages/agenty-core/test/e2e/provider_fixture_test.go index fa90750..1dd68f8 100644 --- a/packages/agenty-core/test/e2e/provider_fixture_test.go +++ b/packages/agenty-core/test/e2e/provider_fixture_test.go @@ -41,7 +41,9 @@ func newProviderFixture(t *testing.T, responder providerResponder) *providerFixt fixture := &providerFixture{requests: make(chan providerRequest, 64)} fixture.server = httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { var body map[string]any - if err := json.NewDecoder(request.Body).Decode(&body); err != nil { + if request.Method == http.MethodGet { + body = map[string]any{} + } else if err := json.NewDecoder(request.Body).Decode(&body); err != nil { t.Errorf("decode provider request: %v", err) http.Error(writer, "invalid request", http.StatusBadRequest) return From a37ddb80d700dfb99416bc99a7f4700ec1674b73 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Mon, 24 Aug 2026 18:28:41 +0800 Subject: [PATCH 02/12] feat: improve CLI provider and model interfaces Signed-off-by: masteryyh --- packages/agenty-cli/src/App.tsx | 4 +- packages/agenty-cli/src/api/client.test.ts | 121 ++- packages/agenty-cli/src/api/client.ts | 70 +- .../agenty-cli/src/api/modelReference.test.ts | 65 ++ packages/agenty-cli/src/api/modelReference.ts | 72 ++ packages/agenty-cli/src/api/types.ts | 32 +- packages/agenty-cli/src/cli/agent.ts | 11 +- packages/agenty-cli/src/cli/init.ts | 41 +- packages/agenty-cli/src/cli/model.ts | 35 +- packages/agenty-cli/src/cli/provider.ts | 11 +- packages/agenty-cli/src/cli/utils.ts | 23 +- packages/agenty-cli/src/commands/registry.ts | 7 +- .../src/components/AgentOverlay.test.ts | 16 - .../src/components/AgentOverlay.tsx | 293 +++--- .../src/components/CommandPalette.tsx | 98 +- .../src/components/CommonComponents.test.tsx | 339 ++++++ .../src/components/ConfirmDialog.tsx | 98 ++ .../agenty-cli/src/components/FormPanel.tsx | 992 +++++++----------- packages/agenty-cli/src/components/List.tsx | 185 ++++ .../src/components/ModelOverlay.test.ts | 119 ++- .../src/components/ModelOverlay.tsx | 774 +++++--------- packages/agenty-cli/src/components/Panel.tsx | 44 + .../src/components/ProviderOverlay.test.ts | 55 + .../src/components/ProviderOverlay.tsx | 811 +++++++++----- .../src/components/ResponsiveLayout.test.tsx | 222 ++++ .../src/components/SelectOverlay.tsx | 71 +- .../src/components/StatusOverlay.tsx | 61 +- .../agenty-cli/src/components/Table.test.ts | 56 + packages/agenty-cli/src/components/Table.tsx | 238 +++++ .../agenty-cli/src/components/TreeList.tsx | 54 + .../src/components/WizardOverlay.tsx | 716 ++++++++----- .../src/components/providerRows.test.ts | 128 +++ .../agenty-cli/src/components/providerRows.ts | 152 +++ .../src/components/ui/Interactive.tsx | 105 ++ .../agenty-cli/src/components/ui/Select.tsx | 37 - .../agenty-cli/src/components/ui/Text.tsx | 7 +- .../src/components/ui/TextInput.tsx | 11 +- .../agenty-cli/src/components/ui/index.ts | 9 +- .../src/components/wizardRows.test.ts | 119 +++ .../agenty-cli/src/components/wizardRows.ts | 127 +++ .../src/components/wizardSetup.test.ts | 158 ++- .../agenty-cli/src/components/wizardSetup.ts | 121 ++- packages/agenty-cli/src/config.ts | 4 +- .../agenty-cli/src/consts/providerPresets.ts | 189 ++-- packages/agenty-cli/src/state/store.test.ts | 4 +- packages/agenty-cli/src/state/store.ts | 6 +- 46 files changed, 4558 insertions(+), 2353 deletions(-) create mode 100644 packages/agenty-cli/src/api/modelReference.test.ts create mode 100644 packages/agenty-cli/src/api/modelReference.ts delete mode 100644 packages/agenty-cli/src/components/AgentOverlay.test.ts create mode 100644 packages/agenty-cli/src/components/CommonComponents.test.tsx create mode 100644 packages/agenty-cli/src/components/ConfirmDialog.tsx create mode 100644 packages/agenty-cli/src/components/List.tsx create mode 100644 packages/agenty-cli/src/components/Panel.tsx create mode 100644 packages/agenty-cli/src/components/ProviderOverlay.test.ts create mode 100644 packages/agenty-cli/src/components/ResponsiveLayout.test.tsx create mode 100644 packages/agenty-cli/src/components/Table.test.ts create mode 100644 packages/agenty-cli/src/components/Table.tsx create mode 100644 packages/agenty-cli/src/components/TreeList.tsx create mode 100644 packages/agenty-cli/src/components/providerRows.test.ts create mode 100644 packages/agenty-cli/src/components/providerRows.ts create mode 100644 packages/agenty-cli/src/components/ui/Interactive.tsx delete mode 100644 packages/agenty-cli/src/components/ui/Select.tsx create mode 100644 packages/agenty-cli/src/components/wizardRows.test.ts create mode 100644 packages/agenty-cli/src/components/wizardRows.ts diff --git a/packages/agenty-cli/src/App.tsx b/packages/agenty-cli/src/App.tsx index eaf062d..ccae71f 100644 --- a/packages/agenty-cli/src/App.tsx +++ b/packages/agenty-cli/src/App.tsx @@ -117,10 +117,10 @@ function ChatView() { return; } try { - const m = await client.resolveModel(ref); + const m = await client.resolveModelInput(ref); await app.switchModel(m); } catch (e) { - app.notify(`model not found: ${ref} (${(e as Error).message})`, true); + app.notify(`model input failed: ${ref} (${(e as Error).message})`, true); } }; diff --git a/packages/agenty-cli/src/api/client.test.ts b/packages/agenty-cli/src/api/client.test.ts index 17007f5..989f80e 100644 --- a/packages/agenty-cli/src/api/client.test.ts +++ b/packages/agenty-cli/src/api/client.test.ts @@ -124,7 +124,7 @@ describe("AgentyClient session list", () => { let updatedWith: ModelDto | undefined; const client = new AgentyClient({} as StdioRPCClient); client.resolveAgent = async () => agent; - client.resolveModel = async () => requestedModel; + client.resolveModelInput = async () => requestedModel; client.getLastSessionByAgent = async () => existing; client.setSessionModel = async (_id, model) => { updatedWith = model; @@ -133,7 +133,7 @@ describe("AgentyClient session list", () => { const prepared = await client.prepareSession({ agentRef: "default", - modelRef: "openai/gpt-new", + modelInput: "openai/gpt-new", newSession: false, }); @@ -141,4 +141,121 @@ describe("AgentyClient session list", () => { expect(prepared.model).toBe(requestedModel); expect(prepared.session.currentModel).toEqual({ providerCode: "openai", modelCode: "gpt-new" }); }); + + test("resolves a persisted current model through its structured reference", async () => { + const agent = { code: "default", name: "Default" } as AgentDto; + const currentModel = { providerCode: "deepseek", modelCode: "deepseek-v4-pro" }; + const resolvedModel = { + code: "deepseek-v4-pro", + providerCode: "deepseek", + providerName: "DeepSeek", + name: "DeepSeek V4 Pro", + } as ModelDto; + const session = { + id: "session", + agentCode: "default", + currentModel, + rounds: [], + } as unknown as ChatSessionDto; + let requestedRef: unknown; + const client = new AgentyClient({} as StdioRPCClient); + client.resolveAgent = async () => agent; + client.getLastSessionByAgent = async () => session; + client.getModel = async (ref) => { + requestedRef = ref; + return resolvedModel; + }; + + const prepared = await client.prepareSession({ agentRef: "default", newSession: false }); + + expect(requestedRef).toEqual(currentModel); + expect(prepared.model).toBe(resolvedModel); + }); +}); + +describe("AgentyClient provider model discovery", () => { + test("resolves structured references within the requested provider", async () => { + let method = ""; + let params: unknown; + const rpc = { + call: async (name: string, input?: unknown) => { + method = name; + params = input; + return [{ + code: "openrouter", + name: "OpenRouter", + type: "openai", + baseUrl: "https://example.invalid", + apiKey: "configured", + models: [{ + code: "deepseek/deepseek-v4-pro", + name: "DeepSeek: DeepSeek V4 Pro 0423", + contextWindow: 128000, + maxOutputTokens: 8192, + multiModal: false, + light: false, + isDefault: false, + }], + createdAt: "2026-01-01T00:00:00Z", + updatedAt: "2026-01-01T00:00:00Z", + }]; + }, + } as unknown as StdioRPCClient; + const client = new AgentyClient(rpc); + + await expect(client.getModel({ + providerCode: "openrouter", + modelCode: "deepseek/deepseek-v4-pro", + })).resolves.toMatchObject({ + providerCode: "openrouter", + code: "deepseek/deepseek-v4-pro", + }); + expect(method).toBe("provider.list"); + expect(params).toEqual({ providerCode: "openrouter" }); + }); + + test("passes an optional target provider to core provider.list", async () => { + let method = ""; + let params: unknown; + const rpc = { + call: async (name: string, input?: unknown) => { + method = name; + params = input; + return []; + }, + } as unknown as StdioRPCClient; + const client = new AgentyClient(rpc); + + await expect(client.listProviders("openrouter")).resolves.toEqual([]); + expect(method).toBe("provider.list"); + expect(params).toEqual({ providerCode: "openrouter" }); + }); + + test("normalizes a null model list and reasoning capabilities", async () => { + let params: unknown; + const rpc = { + call: async (_method: string, input: unknown) => { + params = input; + return [{ + code: "gpt-test", + name: "GPT Test", + contextWindow: 256000, + maxOutputTokens: 65536, + multiModal: false, + reasoningEfforts: undefined, + }, null]; + }, + } as unknown as StdioRPCClient; + const client = new AgentyClient(rpc); + + await expect(client.listProviderModels("openai")).resolves.toEqual([{ + code: "gpt-test", + name: "GPT Test", + contextWindow: 256000, + maxOutputTokens: 65536, + multiModal: false, + reasoningEfforts: [], + }]); + expect(params).toEqual({ providerCode: "openai" }); + }); }); diff --git a/packages/agenty-cli/src/api/client.ts b/packages/agenty-cli/src/api/client.ts index 481e502..5bcb86d 100644 --- a/packages/agenty-cli/src/api/client.ts +++ b/packages/agenty-cli/src/api/client.ts @@ -1,6 +1,8 @@ import type { StdioRPCClient } from "../core/rpc"; +import { formatModelRef, resolveModelInput as resolveModelInputFromList } from "./modelReference"; import type { AgentDto, + AvailableModelDto, ChatMessageDto, ChatSessionDto, CompactionEvent, @@ -13,6 +15,7 @@ import type { InitializeCompleteInput, ModelDto, ModelProviderDto, + ModelRef, PagedResponse, ReasoningEffort, RoundDto, @@ -22,6 +25,7 @@ import type { UpdateModelDto, UpdateModelProviderDto, } from "./types"; +import { STANDARD_REASONING_EFFORTS } from "./types"; export interface PreparedSession { agent: AgentDto; @@ -108,8 +112,10 @@ export class AgentyClient { await this.rpc.call("agent.delete", { code }); } - async listProviders(): Promise { - const providers = await this.rpc.call | null>("provider.list"); + async listProviders(providerCode?: string): Promise { + const providers = providerCode + ? await this.rpc.call | null>("provider.list", { providerCode }) + : await this.rpc.call | null>("provider.list"); return (providers ?? []) .filter((provider): provider is ModelProviderDto => provider !== null) .map(normalizeProvider); @@ -119,6 +125,18 @@ export class AgentyClient { return paginate(await this.listProviders(), page, pageSize); } + async listProviderModels(providerCode: string): Promise { + const models = await this.rpc.call | null>("provider.listModels", { + providerCode, + }); + return (models ?? []) + .filter((model): model is AvailableModelDto => model !== null) + .map((model) => ({ + ...model, + reasoningEfforts: Array.isArray(model.reasoningEfforts) ? model.reasoningEfforts : [], + })); + } + async createProvider(input: CreateModelProviderDto): Promise { const provider = await this.rpc.call("provider.create", input); if (!provider) { @@ -157,22 +175,21 @@ export class AgentyClient { return model; } - async resolveModel(reference?: string): Promise { + async getModel(ref: ModelRef): Promise { + const providers = await this.listProviders(ref.providerCode); + const provider = providers.find((candidate) => candidate.code === ref.providerCode); + const model = provider?.models.find((candidate) => candidate.code === ref.modelCode); + if (!provider || !model) { + throw new Error(`model not found: ${formatModelRef(ref)}`); + } + return projectModel(provider, model); + } + + async resolveModelInput(reference?: string): Promise { if (!reference) { return this.getDefaultModel(); } - const models = await this.listModels(); - const lower = reference.toLowerCase(); - const matches = models.filter((model) => - model.code === reference || - model.name.toLowerCase() === lower || - `${model.providerCode}/${model.code}`.toLowerCase() === lower || - `${model.providerName}/${model.name}`.toLowerCase() === lower, - ); - if (matches.length !== 1) { - throw new Error(matches.length === 0 ? `model not found: ${reference}` : `model reference is ambiguous: ${reference}`); - } - return matches[0]; + return resolveModelInputFromList(await this.listModels(), reference); } async createModel(input: CreateModelDto): Promise { @@ -265,12 +282,12 @@ export class AgentyClient { async prepareSession(options: { agentRef?: string; - modelRef?: string; + modelInput?: string; newSession: boolean; reasoningEffort?: ReasoningEffort; }): Promise { const agent = await this.resolveAgent(options.agentRef); - const requestedModel = options.modelRef ? await this.resolveModel(options.modelRef) : undefined; + const requestedModel = options.modelInput ? await this.resolveModelInput(options.modelInput) : undefined; let session = options.newSession ? null : await this.getLastSessionByAgent(agent.code); if (!session) { const model = requestedModel ?? await this.resolveAgentModel(agent); @@ -292,9 +309,7 @@ export class AgentyClient { } if (session.currentModel) { - const model = await this.resolveModel( - `${session.currentModel.providerCode}/${session.currentModel.modelCode}`, - ); + const model = await this.getModel(session.currentModel); return { agent, model, session }; } @@ -305,7 +320,7 @@ export class AgentyClient { private async resolveAgentModel(agent: AgentDto): Promise { if (agent.defaultModel) { - return this.resolveModel(`${agent.defaultModel.providerCode}/${agent.defaultModel.modelCode}`); + return this.getModel(agent.defaultModel); } return this.getDefaultModel(); } @@ -322,8 +337,19 @@ function projectModel(provider: ModelProviderDto, model: CoreModelDto): ModelDto function normalizeProvider(provider: ModelProviderDto): ModelProviderDto { return { ...provider, + builtin: provider.builtin === true, + official: provider.official === true, + freeFormTool: provider.freeFormTool === true, + modelsCached: provider.modelsCached === true, models: Array.isArray(provider.models) - ? provider.models.filter((model): model is CoreModelDto => model !== null) + ? provider.models + .filter((model): model is CoreModelDto => model !== null) + .map((model) => ({ + ...model, + reasoningEfforts: Array.isArray(model.reasoningEfforts) + ? model.reasoningEfforts + : [...STANDARD_REASONING_EFFORTS], + })) : [], }; } diff --git a/packages/agenty-cli/src/api/modelReference.test.ts b/packages/agenty-cli/src/api/modelReference.test.ts new file mode 100644 index 0000000..ab9db23 --- /dev/null +++ b/packages/agenty-cli/src/api/modelReference.test.ts @@ -0,0 +1,65 @@ +import { describe, expect, test } from "bun:test"; + +import { + findModelByRef, + formatModelRef, + modelRefFromModel, + resolveModelInput, + sameModelRef, +} from "./modelReference"; +import type { ModelDto } from "./types"; + +function model(providerCode: string, code: string, name = code, providerName = providerCode): ModelDto { + return { + code, + providerCode, + providerName, + name, + contextWindow: 128_000, + maxOutputTokens: 8_192, + multiModal: false, + light: false, + reasoningEfforts: [], + isDefault: false, + }; +} + +describe("model references", () => { + test("formats and compares structured references without losing nested model slashes", () => { + const candidate = model("openrouter", "deepseek/deepseek-v4-pro"); + const ref = modelRefFromModel(candidate); + + expect(formatModelRef(ref)).toBe("openrouter/deepseek/deepseek-v4-pro"); + expect(findModelByRef([candidate], ref)).toBe(candidate); + expect(sameModelRef(ref, { providerCode: "openrouter", modelCode: "deepseek/deepseek-v4-pro" })).toBe(true); + }); + + test("gives an explicitly scoped reference precedence over another model's bare code", () => { + const deepSeek = model("deepseek", "deepseek-v4-pro", "DeepSeek V4 Pro", "DeepSeek"); + const openRouter = model( + "openrouter", + "deepseek/deepseek-v4-pro", + "DeepSeek: DeepSeek V4 Pro 0423", + "OpenRouter", + ); + + expect(resolveModelInput([deepSeek, openRouter], "deepseek/deepseek-v4-pro")).toBe(deepSeek); + expect(resolveModelInput([deepSeek, openRouter], "openrouter/deepseek/deepseek-v4-pro")).toBe(openRouter); + }); + + test("still rejects an unscoped code when providers share it", () => { + const first = model("first", "shared-model"); + const second = model("second", "shared-model"); + + expect(() => resolveModelInput([first, second], "shared-model")) + .toThrow("model reference is ambiguous: shared-model; use /"); + }); + + test("allows unique name aliases while keeping canonical references deterministic", () => { + const candidate = model("openai", "gpt-test", "GPT Test", "OpenAI"); + + expect(resolveModelInput([candidate], "gpt test")).toBe(candidate); + expect(resolveModelInput([candidate], "GPT Test")).toBe(candidate); + expect(resolveModelInput([candidate], "OpenAI/GPT Test")).toBe(candidate); + }); +}); diff --git a/packages/agenty-cli/src/api/modelReference.ts b/packages/agenty-cli/src/api/modelReference.ts new file mode 100644 index 0000000..fadfd3e --- /dev/null +++ b/packages/agenty-cli/src/api/modelReference.ts @@ -0,0 +1,72 @@ +import type { ModelDto, ModelRef } from "./types"; + +export type ModelIdentity = Pick; + +export function modelRefFromModel(model: ModelIdentity): ModelRef { + return { + providerCode: model.providerCode, + modelCode: model.code, + }; +} + +export function formatModelRef(ref: ModelRef): string { + return `${ref.providerCode}/${ref.modelCode}`; +} + +export function sameModelRef(left: ModelRef | undefined, right: ModelRef | undefined): boolean { + return !!left && !!right && left.providerCode === right.providerCode && left.modelCode === right.modelCode; +} + +export function findModelByRef(models: readonly ModelDto[], ref: ModelRef): ModelDto | undefined { + return models.find((model) => model.providerCode === ref.providerCode && model.code === ref.modelCode); +} + +/** + * Resolve text entered at a CLI/UI boundary. Canonical provider/model + * references take precedence over every shorthand so model IDs containing '/' + * cannot shadow an explicitly scoped reference. + */ +export function resolveModelInput(models: readonly ModelDto[], input: string): ModelDto { + const reference = input.trim(); + const lower = reference.toLowerCase(); + + const canonicalMatches = models.filter((model) => formatModelRef(modelRefFromModel(model)) === reference); + if (canonicalMatches.length === 1) { + return canonicalMatches[0]; + } + if (canonicalMatches.length > 1) { + throw ambiguousModelReference(reference); + } + + const codeMatches = models.filter((model) => model.code === reference); + if (codeMatches.length === 1) { + return codeMatches[0]; + } + if (codeMatches.length > 1) { + throw ambiguousModelReference(reference); + } + + const nameMatches = models.filter((model) => model.name.toLowerCase() === lower); + if (nameMatches.length === 1) { + return nameMatches[0]; + } + if (nameMatches.length > 1) { + throw ambiguousModelReference(reference); + } + + const providerNameMatches = models.filter((model) => + `${model.providerName}/${model.name}`.toLowerCase() === lower, + ); + if (providerNameMatches.length === 1) { + return providerNameMatches[0]; + } + if (providerNameMatches.length > 1) { + throw ambiguousModelReference(reference); + } + + throw new Error(`model not found: ${reference}`); +} + +function ambiguousModelReference(reference: string): Error { + return new Error(`model reference is ambiguous: ${reference}; use /`); +} diff --git a/packages/agenty-cli/src/api/types.ts b/packages/agenty-cli/src/api/types.ts index 7cd0d60..b9150d9 100644 --- a/packages/agenty-cli/src/api/types.ts +++ b/packages/agenty-cli/src/api/types.ts @@ -1,4 +1,11 @@ export type ReasoningEffort = "" | "off" | "low" | "medium" | "high" | "xhigh" | "max"; +export const STANDARD_REASONING_EFFORTS: readonly ReasoningEffort[] = [ + "low", + "medium", + "high", + "xhigh", + "max", +]; export type APIType = "openai" | "openai_completions" | "anthropic" | "gemini"; export interface ModelRef { @@ -40,11 +47,11 @@ export interface ModelDto { providerName: string; name: string; contextWindow: number; - /** Legacy projection; core now always uses 8192 for every model. */ + /** The exact model limit, or the custom-model fallback when omitted. */ maxOutputTokens: number; multiModal: boolean; light: boolean; - reasoningEffortMapping?: Record; + reasoningEfforts?: ReasoningEffort[]; isDefault: boolean; createdAt?: string; updatedAt?: string; @@ -52,16 +59,25 @@ export interface ModelDto { export interface CoreModelDto extends Omit {} +export interface AvailableModelDto { + code: string; + name: string; + contextWindow: number; + maxOutputTokens: number; + multiModal: boolean; + reasoningEfforts: ReasoningEffort[]; +} + export interface CreateModelDto { providerCode: string; modelCode: string; name: string; contextWindow?: number; - /** @deprecated Core ignores per-model output limits and uses 8192. */ + /** Defaults to the core fallback when omitted. */ maxOutputTokens?: number; multiModal?: boolean; light?: boolean; - reasoningEffortMapping?: Record; + reasoning?: boolean; isDefault?: boolean; } @@ -73,7 +89,14 @@ export interface ModelProviderDto { type: APIType; baseUrl: string; apiKey: string; + freeFormTool?: boolean; + builtin?: boolean; + official?: boolean; + modelsUrl?: string; + tokenCountUrl?: string; models: CoreModelDto[]; + /** True when core populated models from its discovery cache. */ + modelsCached?: boolean; metadata?: Record; createdAt: string; updatedAt: string; @@ -85,6 +108,7 @@ export interface CreateModelProviderDto { type: APIType; baseUrl?: string; apiKey?: string; + freeFormTool?: boolean; metadata?: Record; } diff --git a/packages/agenty-cli/src/cli/agent.ts b/packages/agenty-cli/src/cli/agent.ts index f440719..43a280b 100644 --- a/packages/agenty-cli/src/cli/agent.ts +++ b/packages/agenty-cli/src/cli/agent.ts @@ -1,4 +1,5 @@ import type { AgentyClient } from "@/api/client"; +import { formatModelRef } from "@/api/modelReference"; import type { UpdateAgentDto } from "@/api/types"; import { @@ -14,7 +15,7 @@ import { type ParsedArgs, render, requirePositionals, - resolveModel + resolveModelInput } from "./utils"; export async function handleAgent(client: AgentyClient, args: ParsedArgs): Promise { @@ -26,7 +27,7 @@ export async function handleAgent(client: AgentyClient, args: ParsedArgs): Promi render(args, result, () => result.data.length === 0 ? process.stdout.write("No agents.\n") : outputTable(["Agent Code", "Name", "Default", "Model"], result.data.map((agent) => [ - agent.code, agent.name, String(agent.isDefault), agent.defaultModel ? `${agent.defaultModel.providerCode}/${agent.defaultModel.modelCode}` : "", + agent.code, agent.name, String(agent.isDefault), agent.defaultModel ? formatModelRef(agent.defaultModel) : "", ]))); return; } @@ -35,13 +36,13 @@ export async function handleAgent(client: AgentyClient, args: ParsedArgs): Promi const agent = await client.resolveAgent(reference); render(args, agent, () => outputFields([ ["Agent Code", agent.code], ["Name", agent.name], ["Soul", agent.soul], ["Default", String(agent.isDefault)], - ["Model", agent.defaultModel ? `${agent.defaultModel.providerCode}/${agent.defaultModel.modelCode}` : ""], + ["Model", agent.defaultModel ? formatModelRef(agent.defaultModel) : ""], ])); return; } if (command === "add") { const [, , code] = requirePositionals(args, 3, "agent add [options]"); - const model = flag(args, "model") ? await resolveModel(client, flag(args, "model")!) : undefined; + const model = flag(args, "model") ? await resolveModelInput(client, flag(args, "model")!) : undefined; const created = await client.createAgent({ code, name: flag(args, "name")?.trim() || code, soul: flag(args, "soul") ?? "", isDefault: hasFlag(args, "default") ? parseBoolean(flag(args, "default"), "--default") : false, @@ -65,7 +66,7 @@ export async function handleAgent(client: AgentyClient, args: ParsedArgs): Promi update.isDefault = parseBoolean(flag(args, "default"), "--default"); } if (hasFlag(args, "model")) { - const model = await resolveModel(client, flag(args, "model")!); + const model = await resolveModelInput(client, flag(args, "model")!); update.defaultModel = { providerCode: model.providerCode, modelCode: model.code }; update.defaultContextWindow = model.contextWindow; } diff --git a/packages/agenty-cli/src/cli/init.ts b/packages/agenty-cli/src/cli/init.ts index b00f424..f0dc1e0 100644 --- a/packages/agenty-cli/src/cli/init.ts +++ b/packages/agenty-cli/src/cli/init.ts @@ -1,4 +1,5 @@ import type { AgentyClient } from "@/api/client"; +import { formatModelRef } from "@/api/modelReference"; import type { APIType } from "@/api/types"; import { @@ -19,34 +20,42 @@ export async function handleInit(client: AgentyClient, args: ParsedArgs): Promis const agentCode = flag(args, "agent")?.trim() || "default"; const contextWindow = positiveInteger(flag(args, "context-window") ?? "128000", "--context-window"); const apiKey = secret(args, "api-key", "api-key-env", "provider API key") ?? ""; + const providers = await client.listProviders(); + const existingProvider = providers.find((provider) => provider.code === providerCode); + const existingModel = existingProvider?.models.find((model) => model.code === modelCode); + const effectiveContextWindow = existingModel?.contextWindow ?? contextWindow; - await client.createProvider({ - code: providerCode, - name: flag(args, "provider-name")?.trim() || providerCode, - type: requireFlag(args, "type") as APIType, - baseUrl: flag(args, "base-url")?.trim() || "", - apiKey, - }); - await client.createModel({ - providerCode, - modelCode, - name: flag(args, "model-name")?.trim() || modelCode, - contextWindow, - isDefault: true, - }); + if (existingProvider?.builtin) { + await client.updateProvider(providerCode, { apiKey }); + } else { + await client.createProvider({ + code: providerCode, + name: flag(args, "provider-name")?.trim() || providerCode, + type: requireFlag(args, "type") as APIType, + baseUrl: flag(args, "base-url")?.trim() || "", + apiKey, + }); + await client.createModel({ + providerCode, + modelCode, + name: flag(args, "model-name")?.trim() || modelCode, + contextWindow, + isDefault: true, + }); + } await client.createAgent({ code: agentCode, name: flag(args, "agent-name")?.trim() || agentCode, soul: flag(args, "soul") ?? "", defaultModel: { providerCode, modelCode }, - defaultContextWindow: contextWindow, + defaultContextWindow: effectiveContextWindow, isDefault: true, }); const result = await client.completeInitialization({ agentCode, providerCode, modelCode }); render(args, result, () => outputFields([ ["Initialized", String(result.initialized)], ["Provider", providerCode], - ["Model", `${providerCode}/${modelCode}`], + ["Model", formatModelRef({ providerCode, modelCode })], ["Agent", agentCode], ])); } diff --git a/packages/agenty-cli/src/cli/model.ts b/packages/agenty-cli/src/cli/model.ts index 7c3114b..bb6df13 100644 --- a/packages/agenty-cli/src/cli/model.ts +++ b/packages/agenty-cli/src/cli/model.ts @@ -1,5 +1,5 @@ import type { AgentyClient } from "@/api/client"; -import type { ReasoningEffort, UpdateModelDto } from "@/api/types"; +import type { UpdateModelDto } from "@/api/types"; import { action, @@ -15,7 +15,7 @@ import { render, requireFlag, requirePositionals, - resolveModel, + resolveModelInput, resolveProvider } from "./utils"; @@ -35,13 +35,14 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi return; } if (command === "get") { - const [, , reference] = requirePositionals(args, 3, "model get "); - const model = await resolveModel(client, reference); + const [, , reference] = requirePositionals(args, 3, "model get /"); + const model = await resolveModelInput(client, reference); render(args, model, () => outputFields([ ["Model", displayModel(model)], ["Name", model.name], ["Default", String(model.isDefault)], ["Multimodal", String(model.multiModal)], ["Light", String(model.light)], ["Context window", String(model.contextWindow)], ["Max output tokens", String(model.maxOutputTokens)], + ["Reasoning", String((model.reasoningEfforts?.length ?? 0) > 0)], ])); return; } @@ -56,32 +57,34 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi multiModal: booleanFlag(args, "multi-modal"), light: booleanFlag(args, "light"), isDefault: booleanFlag(args, "default"), - reasoningEffortMapping: reasoningMapping(args), + reasoning: hasFlag(args, "reasoning") ? booleanFlag(args, "reasoning") : true, }); action(args, created, `Model added: ${displayModel(created)}`); return; } if (command === "update") { - const [, , reference] = requirePositionals(args, 3, "model update [options]"); - const current = await resolveModel(client, reference); + const [, , reference] = requirePositionals(args, 3, "model update / [options]"); + const current = await resolveModelInput(client, reference); const update: UpdateModelDto = { name: hasFlag(args, "name") ? requireFlag(args, "name") : current.name, contextWindow: hasFlag(args, "context-window") ? positiveInteger(requireFlag(args, "context-window"), "--context-window", true) : current.contextWindow, multiModal: hasFlag(args, "multi-modal") ? booleanFlag(args, "multi-modal") : current.multiModal, light: hasFlag(args, "light") ? booleanFlag(args, "light") : current.light, isDefault: hasFlag(args, "default") ? booleanFlag(args, "default") : current.isDefault, - reasoningEffortMapping: hasFlag(args, "reasoning-map") ? reasoningMapping(args) : current.reasoningEffortMapping, + reasoning: hasFlag(args, "reasoning") + ? booleanFlag(args, "reasoning") + : (current.reasoningEfforts?.length ?? 0) > 0, }; const updated = await client.updateModel(current.providerCode, current.code, update); action(args, updated, `Model updated: ${displayModel(updated)}`); return; } if (command === "remove") { - const [, , reference] = requirePositionals(args, 3, "model remove --yes"); + const [, , reference] = requirePositionals(args, 3, "model remove / --yes"); if (!hasFlag(args, "yes")) { throw new CliError("use --yes to remove a model non-interactively"); } - const current = await resolveModel(client, reference); + const current = await resolveModelInput(client, reference); await client.deleteModel(current.providerCode, current.code); action(args, { providerCode: current.providerCode, modelCode: current.code, deleted: true }, `Model removed: ${displayModel(current)}`); return; @@ -100,15 +103,3 @@ function positiveInteger(raw: string, label: string, allowZero = false): number } return value; } - -function reasoningMapping(args: ParsedArgs): Record | undefined { - const raw = flag(args, "reasoning-map")?.trim(); - if (!raw) { - return undefined; - } - try { - return JSON.parse(raw) as Record; - } catch { - throw new CliError("--reasoning-map must be a JSON object"); - } -} diff --git a/packages/agenty-cli/src/cli/provider.ts b/packages/agenty-cli/src/cli/provider.ts index 961eadb..5239afc 100644 --- a/packages/agenty-cli/src/cli/provider.ts +++ b/packages/agenty-cli/src/cli/provider.ts @@ -9,6 +9,7 @@ import { outputFields, outputTable, pageOptions, + parseBoolean, type ParsedArgs, render, requireFlag, @@ -36,18 +37,23 @@ export async function handleProvider(client: AgentyClient, args: ParsedArgs): Pr render(args, provider, () => outputFields([ ["Provider Code", provider.code], ["Name", provider.name], ["Type", provider.type], ["Base URL", provider.baseUrl], ["API Key", provider.apiKey ? "" : ""], + ["Free-form apply_patch", provider.freeFormTool === true ? "enabled" : "disabled"], ["Models", String(provider.models.length)], ])); return; } if (command === "add") { const [, , code] = requirePositionals(args, 3, "provider add --type [options]"); + const type = requireFlag(args, "type") as APIType; const created = await client.createProvider({ code, name: flag(args, "name")?.trim() || code, - type: requireFlag(args, "type") as APIType, + type, baseUrl: flag(args, "base-url")?.trim() || "", apiKey: secret(args, "api-key", "api-key-env", "provider API key") ?? "", + freeFormTool: type === "openai" && (hasFlag(args, "free-form-tool") + ? parseBoolean(flag(args, "free-form-tool"), "--free-form-tool") + : false), }); action(args, created, `Provider added: ${created.code}`); return; @@ -65,6 +71,9 @@ export async function handleProvider(client: AgentyClient, args: ParsedArgs): Pr if (hasFlag(args, "base-url")) { update.baseUrl = flag(args, "base-url") ?? ""; } + if (hasFlag(args, "free-form-tool")) { + update.freeFormTool = parseBoolean(flag(args, "free-form-tool"), "--free-form-tool"); + } const apiKey = secret(args, "api-key", "api-key-env", "provider API key"); if (apiKey !== undefined) { update.apiKey = apiKey; diff --git a/packages/agenty-cli/src/cli/utils.ts b/packages/agenty-cli/src/cli/utils.ts index f17eeb6..98c8a89 100644 --- a/packages/agenty-cli/src/cli/utils.ts +++ b/packages/agenty-cli/src/cli/utils.ts @@ -1,4 +1,5 @@ import { AgentyClient } from "@/api/client"; +import { formatModelRef, modelRefFromModel } from "@/api/modelReference"; import type { ModelDto, ModelProviderDto } from "@/api/types"; import { loadOptions } from "@/config"; import { startLocalCore } from "@/localCore"; @@ -125,25 +126,15 @@ export async function resolveProvider(client: AgentyClient, reference: string): } export function displayModel(model: ModelDto): string { - return `${model.providerCode}/${model.code}`; + return formatModelRef(modelRefFromModel(model)); } -export async function resolveModel(client: AgentyClient, reference: string): Promise { - const models = await listAll((page, pageSize) => client.listModelsPage(page, pageSize)); - const lower = reference.toLowerCase(); - - const matched = models.filter((model) => - model.code === reference || - model.name.toLowerCase() === lower || - displayModel(model).toLowerCase() === lower, - ); - if (matched.length === 0) { - throw new CliError(`model not found: ${reference}`); +export async function resolveModelInput(client: AgentyClient, reference: string): Promise { + try { + return await client.resolveModelInput(reference); + } catch (error) { + throw new CliError((error as Error).message); } - if (matched.length > 1) { - throw new CliError(`model reference is ambiguous: ${reference}; use provider/name or model code instead`); - } - return matched[0]; } export function configured(model: ModelDto): boolean { diff --git a/packages/agenty-cli/src/commands/registry.ts b/packages/agenty-cli/src/commands/registry.ts index b4957c5..22b68e4 100644 --- a/packages/agenty-cli/src/commands/registry.ts +++ b/packages/agenty-cli/src/commands/registry.ts @@ -1,4 +1,5 @@ import type { AgentyClient } from "../api/client"; +import { formatModelRef, modelRefFromModel } from "../api/modelReference"; export interface Command { name: string; @@ -17,12 +18,12 @@ export const commands: Command[] = [ { name: "/model", description: "Manage and switch chat models", - usage: "/model [provider/model]", - argHint: "provider/model", + usage: "/model [provider-code/model-code]", + argHint: "provider-code/model-code", completeArgs: async (client) => { const models = await client.listModels(); return models - .map((m) => `${m.providerCode}/${m.code}`); + .map((m) => formatModelRef(modelRefFromModel(m))); }, }, { diff --git a/packages/agenty-cli/src/components/AgentOverlay.test.ts b/packages/agenty-cli/src/components/AgentOverlay.test.ts deleted file mode 100644 index 1700a54..0000000 --- a/packages/agenty-cli/src/components/AgentOverlay.test.ts +++ /dev/null @@ -1,16 +0,0 @@ -import { describe, expect, test } from "bun:test"; - -import { parseModelRef } from "./AgentOverlay"; - -describe("parseModelRef", () => { - test("keeps slashes inside a model code", () => { - expect(parseModelRef("openai/org/model_name[v2]")) - .toEqual({ providerCode: "openai", modelCode: "org/model_name[v2]" }); - }); - - test("rejects references without both sides", () => { - expect(parseModelRef("openai")).toBeUndefined(); - expect(parseModelRef("/model")).toBeUndefined(); - expect(parseModelRef("openai/")).toBeUndefined(); - }); -}); diff --git a/packages/agenty-cli/src/components/AgentOverlay.tsx b/packages/agenty-cli/src/components/AgentOverlay.tsx index a633edc..5341cfb 100644 --- a/packages/agenty-cli/src/components/AgentOverlay.tsx +++ b/packages/agenty-cli/src/components/AgentOverlay.tsx @@ -1,40 +1,22 @@ import { useCallback, useEffect, useRef, useState } from "react"; -import type { AgentDto, ModelDto, ModelRef } from "../api/types"; +import { formatModelRef, modelRefFromModel } from "../api/modelReference"; +import type { AgentDto, ModelDto } from "../api/types"; import { useInput } from "../hooks/useInput"; import { useAppStore } from "../state/store"; import { useBottomDialogSize } from "./BottomDialog"; +import { ConfirmDialog } from "./ConfirmDialog"; import type { FormField, FormOption } from "./FormPanel"; import { FormPanel } from "./FormPanel"; -import { Box, Spinner, Text } from "./ui"; - -function trunc(s: string, width: number): string { - if (width <= 0) { - return ""; - } - if (s.length <= width) { - return s; - } - if (width === 1) { - return "…"; - } - return s.slice(0, width - 1) + "…"; -} - -function pad(s: string, width: number): string { - const clipped = trunc(s, width); - return clipped + " ".repeat(Math.max(width - clipped.length, 0)); -} - -export function parseModelRef(raw: string): ModelRef | undefined { - const separator = raw.indexOf("/"); - if (separator <= 0 || separator === raw.length - 1) { - return undefined; - } - const providerCode = raw.slice(0, separator); - const modelCode = raw.slice(separator + 1); - return { providerCode, modelCode }; -} +import { List, useListNavigation } from "./List"; +import { Panel } from "./Panel"; +import { + createTableLayout, + type TableColumn, + TableHeader, + TableRow, +} from "./Table"; +import { ActionBar, Box, Spinner, Text } from "./ui"; type Mode = | { kind: "list" } @@ -81,7 +63,7 @@ export function AgentOverlay() { setModelOptions( models.map((m) => ({ label: `${m.providerName} · ${m.name}`, - value: `${m.providerCode}/${m.code}`, + value: formatModelRef(modelRefFromModel(m)), })), ); } catch { @@ -115,15 +97,15 @@ export function AgentOverlay() { }); const buildFields = (target?: AgentDto): FormField[] => { - const modelRef = target?.defaultModel - ? `${target.defaultModel.providerCode}/${target.defaultModel.modelCode}` + const modelValue = target?.defaultModel + ? formatModelRef(target.defaultModel) : modelOptions[0]?.value ?? ""; return [ { key: "code", label: "Agent Code", kind: "text" as const, value: target?.code ?? "", placeholder: "my-agent", readOnly: !!target }, { key: "name", label: "Name", kind: "text" as const, value: target?.name ?? "", placeholder: "my-agent" }, { key: "soul", label: "Soul", kind: "text" as const, value: target?.soul ?? "", placeholder: "system prompt, leave blank for default" }, { key: "isDefault", label: "Default", kind: "boolean" as const, value: target ? (target.isDefault ? "true" : "false") : "false" }, - { key: "defaultModel", label: "Default model", kind: "select" as const, value: modelRef, options: modelOptions }, + { key: "defaultModel", label: "Default model", kind: "select" as const, value: modelValue, options: modelOptions }, ]; }; @@ -132,8 +114,8 @@ export function AgentOverlay() { return; } try { - const defaultModel = parseModelRef(values.defaultModel); - const selectedModel = models.find((model) => `${model.providerCode}/${model.code}` === values.defaultModel); + const selectedModel = modelForValue(models, values.defaultModel); + const defaultModel = selectedModel ? modelRefFromModel(selectedModel) : undefined; await client.createAgent({ code: values.code.trim(), name: values.name.trim(), @@ -155,8 +137,8 @@ export function AgentOverlay() { return; } try { - const defaultModel = parseModelRef(values.defaultModel); - const selectedModel = models.find((model) => `${model.providerCode}/${model.code}` === values.defaultModel); + const selectedModel = modelForValue(models, values.defaultModel); + const defaultModel = selectedModel ? modelRefFromModel(selectedModel) : undefined; await client.updateAgent(target.code, { name: values.name.trim(), soul: values.soul.trim(), @@ -243,10 +225,7 @@ export function AgentOverlay() { } return ( - - - Agents - + {agents === null ? ( ) : agents.length === 0 ? ( @@ -264,10 +243,14 @@ export function AgentOverlay() { onClose={close} /> )} - + ); } +function modelForValue(models: readonly ModelDto[], value: string): ModelDto | undefined { + return models.find((model) => formatModelRef(modelRefFromModel(model)) === value); +} + // ─── Agent list table ─────────────────────────────────────────────── function AgentList({ @@ -292,117 +275,100 @@ function AgentList({ onClose: () => void; }) { const dialogSize = useBottomDialogSize(); - const n = agents.length; - const compact = dialogSize.width < 44; - const flagsWidth = compact ? 0 : 24; - const nameWidth = compact - ? Math.max(dialogSize.width - 2, 8) - : Math.max(Math.min(dialogSize.width - flagsWidth - 4, 32), 12); - const maxVisible = Math.max(dialogSize.height - 6 - (compact ? 1 : 0), 1); - const maxVis = Math.min(maxVisible, n); - const half = Math.floor(maxVis / 2); - let start = cursor - half; - if (start < 0) { - start = 0; - } - if (start + maxVis > n) { - start = Math.max(n - maxVis, 0); - } - const visible = agents.slice(start, start + maxVis); + const maxVisible = Math.max(dialogSize.height - 5, 1); + const agentFlags = (agent: AgentDto): string => + `${agent.isDefault ? "[default] " : ""}${agent.code === currentAgentCode ? "← current" : ""}`.trim(); + const columns: Array> = [ + { + key: "name", + header: "Name", + value: (agent) => agent.name, + render: (agent, selected) => ( + + {agent.name} + + ), + }, + { + key: "flags", + header: "Flags", + value: agentFlags, + render: (agent, selected) => ( + + {agentFlags(agent)} + + ), + }, + ]; + const tableLayout = createTableLayout( + columns, + agents, + Math.max(dialogSize.width - 2, 0), + ); - useInput((input, key) => { - if (key.escape) { - onClose(); - return; - } - if (key.upArrow) { - onCursor(Math.max(cursor - 1, 0)); - return; - } - if (key.downArrow) { - onCursor(Math.min(cursor + 1, n - 1)); - return; - } - const lower = input.toLowerCase(); - const a = agents[cursor]; - if (key.return || lower === "s") { - onSwitch(a); - return; - } - if (lower === "a") { - onAdd(); - } else if (lower === "e") { - onEdit(a); - } else if (lower === "d") { - onDelete(a); - } + useListNavigation({ + items: agents, + cursor, + onCursor, + onActivate: onSwitch, + onClose, + onInput: (input, _key, _event, agent) => { + const lower = input.toLowerCase(); + if (lower === "s" && agent) { + onSwitch(agent); + } else if (lower === "a") { + onAdd(); + } else if (lower === "e" && agent) { + onEdit(agent); + } else if (lower === "d" && agent) { + onDelete(agent); + } + }, }); return ( - - - - {compact - ? ` ${pad("Name", nameWidth)}` - : ` ${pad("Name", nameWidth)} ${pad("Flags", flagsWidth)}`} - - - - {visible.map((a) => { - const i = agents.indexOf(a); - const selected = i === cursor; - const flags = - `${a.isDefault ? "[default] " : ""}${a.code === currentAgentCode ? "← current" : ""}`.trim(); - const name = pad(a.name, nameWidth); - return ( - onCursor(i)} - onMouseClick={() => { - onCursor(i); - onSwitch(a); - }} - > - - {selected ? "❯" : " "} - - - - {name} - - {compact ? null : } - {compact ? null : ( - - {pad(flags, flagsWidth)} - - )} - - ); - })} - - {compact ? ( - - - {trunc( - `${agents[cursor]?.isDefault ? "[default] " : ""}${agents[cursor]?.code === currentAgentCode ? "← current" : ""}`.trim() || "No flags", - dialogSize.width, - )} - - - ) : null} - - [Add] - onEdit(agents[cursor])}>[Edit] - onDelete(agents[cursor])}>[Delete] - - - - {compact - ? "↑↓ move · Enter switch · e edit · d del · Esc" - : "↑↓ navigate · Enter/s switch · e edit · d delete · Esc back"} - + { + const agent = agents[cursor]; + if (key === "add") { + onAdd(); + } else if (key === "edit" && agent) { + onEdit(agent); + } else if (key === "delete" && agent) { + onDelete(agent); + } + }} + /> + )} + hint="↑↓ navigate · Enter/s switch · e edit · d delete · Esc back" + > + + + - + agent.code} + onCursor={onCursor} + onActivate={onSwitch} + renderItem={(agent, { selected }) => ( + + )} + /> + ); } @@ -417,29 +383,12 @@ function DeleteConfirm({ onConfirm: () => void; onCancel: () => void; }) { - useInput((input, key) => { - if (key.escape) { - onCancel(); - return; - } - const lower = input.toLowerCase(); - if (lower === "y") { - onConfirm(); - } else if (lower === "n") { - onCancel(); - } - }); - return ( - - - Delete agent "{target.name}"? - - This also deletes all its sessions, messages and memories. - - [Delete] - [Cancel] - - + ); } diff --git a/packages/agenty-cli/src/components/CommandPalette.tsx b/packages/agenty-cli/src/components/CommandPalette.tsx index f2f6164..57c106f 100644 --- a/packages/agenty-cli/src/components/CommandPalette.tsx +++ b/packages/agenty-cli/src/components/CommandPalette.tsx @@ -4,7 +4,7 @@ import { useEffect, useRef, useState } from "react"; import { quoteArg } from "../commands/registry"; import type { Palette } from "../hooks/useCommandPalette"; import { useWindowSize } from "../hooks/useWindowSize"; -import { Box, Text } from "./ui"; +import { Box, Pressable, Text } from "./ui"; const MAX_ITEMS = 8; const HIGHLIGHT = "#4FA8FF"; @@ -86,22 +86,26 @@ export function CommandPalette({ palette, marginTop, onChoose }: CommandPaletteP const contentLen = cursor.length + c.name.length + 3 + c.description.length; return ( - onChoose(`${c.name}${c.argHint ? " " : ""}`)} + width="100%" + height={1} + onPress={() => onChoose(`${c.name}${c.argHint ? " " : ""}`)} > - - {cursor} + + + {cursor} + + + {c.name} + + + {" — "} + {c.description} + + {padSpaces(contentLen)} - - {c.name} - - - {" — "} - {c.description} - - {padSpaces(contentLen)} - + ); } const matchedPart = c.name.slice(0, matchPrefix.length); @@ -109,30 +113,34 @@ export function CommandPalette({ palette, marginTop, onChoose }: CommandPaletteP const contentLen = cursor.length + c.name.length + 3 + c.description.length; return ( - onChoose(`${c.name}${c.argHint ? " " : ""}`)} + width="100%" + height={1} + onPress={() => onChoose(`${c.name}${c.argHint ? " " : ""}`)} > - - {cursor} - - - {matchedPart} - - {unmatchedPart ? ( + - {unmatchedPart} + {cursor} + + + {matchedPart} + + {unmatchedPart ? ( + + {unmatchedPart} + + ) : null} + + {" — "} + {c.description} - ) : null} - - {" — "} - {c.description} + {padSpaces(contentLen)} - {padSpaces(contentLen)} - + ); })} @@ -181,17 +189,23 @@ export function CommandPalette({ palette, marginTop, onChoose }: CommandPaletteP const prefix = ` ${selected ? "❯" : " "} `; const contentLen = prefix.length + c.length; return ( - onChoose(`${command.name} ${quoteArg(c)}`)} + width="100%" + height={1} + onPress={() => onChoose(`${command.name} ${quoteArg(c)}`)} > - {prefix} - {c} - {padSpaces(contentLen)} - + + {prefix} + {c} + {padSpaces(contentLen)} + + ); }) )} diff --git a/packages/agenty-cli/src/components/CommonComponents.test.tsx b/packages/agenty-cli/src/components/CommonComponents.test.tsx new file mode 100644 index 0000000..07a87cf --- /dev/null +++ b/packages/agenty-cli/src/components/CommonComponents.test.tsx @@ -0,0 +1,339 @@ +import { BaseRenderable, BoxRenderable, InputRenderable, RGBA } from "@opentui/core"; +import { testRender } from "@opentui/react/test-utils"; +import { describe, expect, test } from "bun:test"; +import { act, useState } from "react"; + +import { BottomDialog } from "./BottomDialog"; +import { ConfirmDialog } from "./ConfirmDialog"; +import type { FormField } from "./FormPanel"; +import { FormPanel } from "./FormPanel"; +import { List } from "./List"; +import { Box, HOVER_BACKGROUND, Text } from "./ui"; + +function InteractiveList({ onActivate }: { onActivate: (value: string) => void }) { + const [cursor, setCursor] = useState(0); + return ( + item} + onCursor={setCursor} + onActivate={onActivate} + renderItem={(item, { selected }) => ( + + {item} + + )} + /> + ); +} + +function findBox( + renderable: BaseRenderable, + predicate: (box: BoxRenderable) => boolean, +): BoxRenderable | null { + if (renderable instanceof BoxRenderable && predicate(renderable)) { + return renderable; + } + for (const child of renderable.getChildren()) { + const match = findBox(child, predicate); + if (match) { + return match; + } + } + return null; +} + +function findInput(renderable: BaseRenderable): InputRenderable | null { + if (renderable instanceof InputRenderable) { + return renderable; + } + for (const child of renderable.getChildren()) { + const match = findInput(child); + if (match) { + return match; + } + } + return null; +} + +describe("common TUI components", () => { + test("highlights a hovered list row without selecting or activating it", async () => { + const activated: string[] = []; + const setup = await testRender( + activated.push(value)} />, + { width: 30, height: 4 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + await act(async () => { + await setup.mockMouse.moveTo(5, 1, { delayMs: 10 }); + }); + await act(async () => { + await setup.flush(); + await setup.waitForVisualIdle(); + }); + + let frame = setup.captureCharFrame(); + expect(frame.split("\n")[0]).toContain("❯ first"); + expect(frame.split("\n")[1]).not.toContain("❯"); + expect(activated).toEqual([]); + const hoverColor = RGBA.fromHex(HOVER_BACKGROUND); + expect(setup.captureSpans().lines[1]?.spans.some((span) => span.bg.equals(hoverColor))) + .toBe(true); + await act(async () => { + await setup.mockMouse.click(5, 1); + await setup.flush(); + }); + + frame = setup.captureCharFrame(); + expect(frame.split("\n")[1]).toContain("❯ second"); + expect(activated).toEqual(["second"]); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("edits a focused text row directly and confirms a select with a second Enter", async () => { + const fields: FormField[] = [ + { key: "name", label: "Name", kind: "text", value: "" }, + { + key: "type", + label: "Type", + kind: "select", + value: "second", + options: [ + { label: "First", value: "first" }, + { label: "Second", value: "second" }, + ], + }, + ]; + let saved: Record | undefined; + const setup = await testRender( + { + saved = values; + }} + onClose={() => undefined} + />, + { width: 50, height: 12 }, + ); + + try { + await act(async () => { + await setup.flush(); + await setup.mockInput.typeText("Direct edit"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + await setup.mockInput.pressKeys(["RETURN"], 10); + await setup.flush(); + }); + await act(async () => { + await setup.flush(); + }); + + expect(saved).toBeUndefined(); + expect(setup.captureCharFrame()).toContain("❯ Second"); + + await act(async () => { + setup.mockInput.pressArrow("up"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEnter(); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEnter(); + await setup.flush(); + }); + + expect(saved).toMatchObject({ + name: "Direct edit", + type: "first", + }); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("uses form shortcuts outside text editing without intercepting typed text", async () => { + const shortcuts: string[] = []; + const setup = await testRender( + { + if (input.toLowerCase() === "d") { + shortcuts.push("delete"); + return true; + } + return false; + }} + onAction={() => undefined} + onClose={() => undefined} + />, + { width: 50, height: 10 }, + ); + + try { + await act(async () => { + await setup.flush(); + await setup.mockInput.typeText("d"); + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + await setup.mockInput.typeText("model d"); + await setup.flush(); + }); + + expect(shortcuts).toEqual(["delete"]); + expect(setup.captureCharFrame()).toContain("model d"); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("skips fields marked as non-focusable", async () => { + const setup = await testRender( + undefined} + onClose={() => undefined} + />, + { width: 50, height: 10 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + + let frame = setup.captureCharFrame(); + const codeLine = frame.split("\n").find((line) => line.includes("openai")) ?? ""; + expect(codeLine).not.toContain("❯"); + expect(frame.split("\n").some((line) => line.includes("❯"))).toBe(true); + + await act(async () => { + setup.mockInput.pressArrow("up"); + await setup.flush(); + }); + + frame = setup.captureCharFrame(); + expect(frame.split("\n").some((line) => line.includes("❯"))).toBe(true); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("refocuses the active text input when clicked after blur", async () => { + const setup = await testRender( + undefined} + onClose={() => undefined} + />, + { width: 50, height: 10 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + + const input = findInput(setup.renderer.root); + expect(input).not.toBeNull(); + input?.blur(); + expect(input?.focused).toBe(false); + + await act(async () => { + if (input) { + await setup.mockMouse.click(input.x + 1, input.y); + } + await setup.flush(); + }); + + expect(input?.focused).toBe(true); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("centers confirmation dialogs inside the active panel", async () => { + const setup = await testRender( + + undefined} + onCancel={() => undefined} + /> + , + { width: 60, height: 22 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + + const dialog = findBox( + setup.renderer.root, + (box) => box.borderColor.equals(RGBA.fromHex("#ff0000")), + ); + expect(dialog).not.toBeNull(); + const container = dialog?.parent; + expect(container).toBeInstanceOf(BoxRenderable); + if (!(container instanceof BoxRenderable) || !dialog) { + return; + } + const left = dialog.x - container.x; + const right = container.width - left - dialog.width; + const top = dialog.y - container.y; + const bottom = container.height - top - dialog.height; + expect(Math.abs(left - right)).toBeLessThanOrEqual(1); + expect(Math.abs(top - bottom)).toBeLessThanOrEqual(1); + } finally { + act(() => setup.renderer.destroy()); + } + }); +}); diff --git a/packages/agenty-cli/src/components/ConfirmDialog.tsx b/packages/agenty-cli/src/components/ConfirmDialog.tsx new file mode 100644 index 0000000..863b472 --- /dev/null +++ b/packages/agenty-cli/src/components/ConfirmDialog.tsx @@ -0,0 +1,98 @@ +import { RGBA } from "@opentui/core"; +import type { ReactNode } from "react"; +import { useState } from "react"; + +import { useInput } from "../hooks/useInput"; +import { useBottomDialogSize } from "./BottomDialog"; +import { ActionBar, Box, Text } from "./ui"; + +const CONFIRM_DIALOG_Z_INDEX = 120; +const TERMINAL_BACKGROUND = RGBA.defaultBackground(); + +export interface ConfirmDialogProps { + title: string; + message: ReactNode; + confirmLabel?: string; + cancelLabel?: string; + onConfirm: () => void; + onCancel: () => void; + active?: boolean; +} + +export function ConfirmDialog({ + title, + message, + confirmLabel = "Delete", + cancelLabel = "Cancel", + onConfirm, + onCancel, + active = true, +}: ConfirmDialogProps) { + const dialogSize = useBottomDialogSize(); + const [cursor, setCursor] = useState(1); + const width = Math.max(Math.min(52, dialogSize.width), 1); + const height = Math.max(Math.min(7, dialogSize.height), 1); + + const activate = (index: number) => { + if (index === 0) { + onConfirm(); + } else { + onCancel(); + } + }; + + useInput((input, key) => { + if (key.escape || input.toLowerCase() === "n") { + onCancel(); + return; + } + if (input.toLowerCase() === "y") { + onConfirm(); + return; + } + if (key.leftArrow || key.rightArrow) { + setCursor((current) => current === 0 ? 1 : 0); + return; + } + if (key.return) { + activate(cursor); + } + }, { isActive: active }); + + return ( + + + {title} + + {message} + + activate(key === "confirm" ? 0 : 1)} + /> + + + ); +} diff --git a/packages/agenty-cli/src/components/FormPanel.tsx b/packages/agenty-cli/src/components/FormPanel.tsx index b491da9..122ff31 100644 --- a/packages/agenty-cli/src/components/FormPanel.tsx +++ b/packages/agenty-cli/src/components/FormPanel.tsx @@ -1,10 +1,12 @@ +import type { InputRenderable, KeyEvent } from "@opentui/core"; import { useCallback, useMemo, useRef, useState } from "react"; +import type { InputKey } from "../hooks/useInput"; import { useInput } from "../hooks/useInput"; import { useBottomDialogSize } from "./BottomDialog"; -import { Box, Text, TextInput } from "./ui"; - -// ─── types ────────────────────────────────────────────────────────── +import { Panel } from "./Panel"; +import { allocateColumnWidths, textWidth, truncateText } from "./Table"; +import { ActionBar, Box, Pressable, Text, TextInput } from "./ui"; export interface FormOption { label: string; @@ -20,6 +22,7 @@ export interface FormField { placeholder?: string; secret?: boolean; readOnly?: boolean; + focusable?: boolean; visible?: boolean; } @@ -32,38 +35,30 @@ export interface FormPanelProps { title: string; fields: FormField[]; actions?: FormAction[]; + active?: boolean; + error?: string | null; + hint?: string; + shortcutHint?: string; onChange?: (key: string, allValues: Record) => void; + onShortcut?: ( + input: string, + key: InputKey, + event: KeyEvent, + values: Record, + ) => boolean; onAction: (key: string, values: Record) => void; onClose: () => void; } -// ─── display helpers ───────────────────────────────────────────────── - -const KEY_WIDTH = 22; - -function pad(s: string, w: number): string { - if (w <= 0) { - return ""; - } - if (s.length <= w) { - return s + " ".repeat(w - s.length); - } - if (w === 1) { - return "\u2026"; +function maskValue(value: string): string { + if (!value) { + return "—"; } - return s.slice(0, w - 1) + "\u2026"; -} - -function maskValue(v: string): string { - if (!v) { - return "\u2014"; - } - return "\u2022".repeat(Math.min(v.length, 20)); + return "•".repeat(Math.min(value.length, 20)); } function selectLabel(options: FormOption[], value: string): string { - const found = options.find((o) => o.value === value); - return found ? found.label : value; + return options.find((option) => option.value === value)?.label ?? value; } function parseMulti(value: string): Set { @@ -73,20 +68,17 @@ function parseMulti(value: string): Set { return new Set(parsed.filter((item): item is string => typeof item === "string")); } } catch { - // fall through + return new Set(); } return new Set(); } -function serializeMulti(set: Set): string { - return JSON.stringify(Array.from(set)); +function serializeMulti(values: Set): string { + return JSON.stringify(Array.from(values)); } -// ─── component ─────────────────────────────────────────────────────── - -type FieldState = +type ChoiceState = | { kind: "idle" } - | { kind: "editing"; visibleIndex: number; text: string } | { kind: "selecting"; visibleIndex: number; selection: number } | { kind: "multi-selecting"; @@ -99,662 +91,470 @@ export function FormPanel({ title, fields, actions, + active = true, + error, + hint: hintOverride, + shortcutHint, onChange, + onShortcut, onAction, onClose, }: FormPanelProps) { const dialogSize = useBottomDialogSize(); - const [values, setValues] = useState>(() => { - const init: Record = {}; - for (const f of fields) { - init[f.key] = f.value; - } - return init; - }); - - const actDefs: FormAction[] = actions ?? [ - { key: "save", label: "Save" }, - { key: "cancel", label: "Cancel" }, - ]; - - // visible subset const visibleFields = useMemo( - () => fields.filter((f) => f.visible !== false), + () => fields.filter((field) => field.visible !== false), [fields], ); + const actionDefs = useMemo( + () => actions ?? [ + { key: "save", label: "Save" }, + { key: "cancel", label: "Cancel" }, + ], + [actions], + ); const actionStart = visibleFields.length; - const actionEnd = actionStart + actDefs.length - 1; - const keyWidth = Math.min(KEY_WIDTH, Math.max(Math.floor(dialogSize.width / 3), 12)); + const actionEnd = actionStart + actionDefs.length - 1; + const navigationIndexes = useMemo(() => [ + ...visibleFields.flatMap((field, index) => field.focusable === false ? [] : [index]), + ...actionDefs.map((_action, index) => actionStart + index), + ], [actionDefs, actionStart, visibleFields]); + const [values, setValues] = useState>(() => + Object.fromEntries(fields.map((field) => [field.key, field.value])), + ); + const [cursor, setCursor] = useState(navigationIndexes[0] ?? 0); + const [choice, setChoice] = useState({ kind: "idle" }); + const textInputRef = useRef(null); + + const formColumnBudget = Math.max(dialogSize.width - 3, 0); + const labelContentWidth = Math.max( + ...visibleFields.map((field) => textWidth(`${field.label}:`)), + 0, + ); + const [keyWidth = 0] = allocateColumnWidths( + formColumnBudget, + [labelContentWidth, formColumnBudget], + ); const maxExpandedOptions = Math.max( 2, Math.min(6, dialogSize.height - visibleFields.length - 4), ); - - const [cursor, setCursor] = useState(0); - const [fstate, setFstate] = useState({ kind: "idle" }); - - // refs for stale closure safety const valuesRef = useRef(values); valuesRef.current = values; const cursorRef = useRef(cursor); cursorRef.current = cursor; - const fstateRef = useRef(fstate); - fstateRef.current = fstate; - const visibleRef = useRef(visibleFields); - visibleRef.current = visibleFields; - const actDefsRef = useRef(actDefs); - actDefsRef.current = actDefs; + const choiceRef = useRef(choice); + choiceRef.current = choice; + const visibleFieldsRef = useRef(visibleFields); + visibleFieldsRef.current = visibleFields; + const actionDefsRef = useRef(actionDefs); + actionDefsRef.current = actionDefs; + + const updateValue = useCallback((key: string, value: string) => { + const next = { ...valuesRef.current, [key]: value }; + valuesRef.current = next; + setValues(next); + if (onChange) { + queueMicrotask(() => onChange(key, next)); + } + }, [onChange]); + + const moveCursor = useCallback((direction: -1 | 1, from = cursorRef.current) => { + const currentPosition = navigationIndexes.indexOf(from); + const fallbackPosition = direction > 0 ? -1 : navigationIndexes.length; + const nextPosition = Math.min( + Math.max(currentPosition < 0 ? fallbackPosition + direction : currentPosition + direction, 0), + Math.max(navigationIndexes.length - 1, 0), + ); + const next = navigationIndexes[nextPosition]; + if (next === undefined) { + return; + } + textInputRef.current?.blur(); + setChoice({ kind: "idle" }); + setCursor(next); + }, [navigationIndexes]); + + const openChoice = useCallback((visibleIndex: number) => { + const field = visibleFieldsRef.current[visibleIndex]; + if (!field || field.readOnly) { + return; + } + const options = field.options ?? []; + if (options.length === 0) { + return; + } - const doSetValues = useCallback( - (updater: (prev: Record) => Record) => { - setValues((prev) => { - const next = updater(prev); - return next; + const current = valuesRef.current[field.key] ?? field.value; + if (field.kind === "select") { + const selected = options.findIndex((option) => option.value === current); + setChoice({ + kind: "selecting", + visibleIndex, + selection: selected >= 0 ? selected : 0, }); - }, - [], - ); - - const notifyChange = useCallback( - (key: string) => { - if (!onChange) { - return; - } - // read latest values via ref inside a timeout to avoid setState-in-setState - queueMicrotask(() => { - onChange(key, valuesRef.current); + } else if (field.kind === "multiselect") { + const chosen = parseMulti(current); + const firstChosen = options.findIndex((option) => chosen.has(option.value)); + setChoice({ + kind: "multi-selecting", + visibleIndex, + selection: firstChosen >= 0 ? firstChosen : 0, + chosen, }); - }, - [onChange], - ); - - const updateValue = useCallback( - (key: string, v: string) => { - doSetValues((prev) => ({ ...prev, [key]: v })); - notifyChange(key); - }, - [doSetValues, notifyChange], - ); - - const commitEdit = useCallback( - (text: string) => { - const st = fstateRef.current; - if (st.kind !== "editing") { - return; - } - const f = visibleRef.current[st.visibleIndex]; - if (!f) { - return; - } - doSetValues((prev) => ({ ...prev, [f.key]: text })); - notifyChange(f.key); - setFstate({ kind: "idle" }); - }, - [doSetValues, notifyChange], - ); - - const cancelEdit = useCallback(() => { - setFstate({ kind: "idle" }); + } }, []); - const commitSelect = useCallback( - (selection: number) => { - const st = fstateRef.current; - if (st.kind !== "selecting") { - return; + const commitSelect = useCallback((selection: number) => { + const state = choiceRef.current; + if (state.kind !== "selecting") { + return; + } + const field = visibleFieldsRef.current[state.visibleIndex]; + const option = field?.options?.[selection]; + if (field && option) { + updateValue(field.key, option.value); + } + setChoice({ kind: "idle" }); + }, [updateValue]); + + const toggleMultiSelect = useCallback((selection: number) => { + setChoice((state) => { + if (state.kind !== "multi-selecting") { + return state; } - const f = visibleRef.current[st.visibleIndex]; - if (!f) { - return; + const option = visibleFieldsRef.current[state.visibleIndex]?.options?.[selection]; + if (!option) { + return state; } - const opts = f.options ?? []; - if (selection >= 0 && selection < opts.length) { - doSetValues((prev) => ({ ...prev, [f.key]: opts[selection].value })); - notifyChange(f.key); + const chosen = new Set(state.chosen); + if (chosen.has(option.value)) { + chosen.delete(option.value); + } else { + chosen.add(option.value); } - setFstate({ kind: "idle" }); - }, - [doSetValues, notifyChange], - ); - - const cancelSelect = useCallback(() => { - setFstate({ kind: "idle" }); + return { ...state, selection, chosen }; + }); }, []); const commitMultiSelect = useCallback(() => { - const st = fstateRef.current; - if (st.kind !== "multi-selecting") { - setFstate({ kind: "idle" }); + const state = choiceRef.current; + if (state.kind !== "multi-selecting") { return; } - const f = visibleRef.current[st.visibleIndex]; - if (!f) { - setFstate({ kind: "idle" }); - return; + const field = visibleFieldsRef.current[state.visibleIndex]; + if (field) { + updateValue(field.key, serializeMulti(state.chosen)); } - doSetValues((prev) => ({ ...prev, [f.key]: serializeMulti(st.chosen) })); - notifyChange(f.key); - setFstate({ kind: "idle" }); - }, [doSetValues, notifyChange]); - - const cancelMultiSelect = useCallback(() => { - setFstate({ kind: "idle" }); - }, []); + setChoice({ kind: "idle" }); + }, [updateValue]); - const activateField = useCallback( - (visibleIndex: number) => { - const field = visibleRef.current[visibleIndex]; - if (!field || field.readOnly) { - return; - } - setCursor(visibleIndex); - const current = valuesRef.current[field.key] ?? field.value; - if (field.kind === "boolean") { - updateValue(field.key, current === "true" ? "false" : "true"); - return; - } - if (field.kind === "select") { - const options = field.options ?? []; - if (options.length === 0) { - return; - } - const selected = options.findIndex((option) => option.value === current); - setFstate({ - kind: "selecting", - visibleIndex, - selection: selected >= 0 ? selected : 0, - }); - return; - } - if (field.kind === "multiselect") { - const options = field.options ?? []; - if (options.length === 0) { - return; - } - const chosen = parseMulti(current); - const firstChosen = options.findIndex((option) => chosen.has(option.value)); - setFstate({ - kind: "multi-selecting", - visibleIndex, - selection: firstChosen >= 0 ? firstChosen : 0, - chosen, - }); - return; - } - setFstate({ kind: "editing", visibleIndex, text: current }); - }, - [updateValue], - ); - - useInput((input, key, event) => { - const st = fstateRef.current; - - // ── editing text field ── - if (st.kind === "editing") { - if (key.escape) { - event.preventDefault(); - cancelEdit(); - return; - } - // TextInput handles Enter to submit + const runAction = useCallback((actionIndex: number) => { + const action = actionDefsRef.current[actionIndex]; + if (!action) { return; } + if (action.key === "cancel") { + onClose(); + } else { + onAction(action.key, valuesRef.current); + } + }, [onAction, onClose]); - // ── selecting option ── - if (st.kind === "selecting") { - const f = visibleRef.current[st.visibleIndex]; - const opts = f?.options ?? []; + useInput((input, key, event) => { + const state = choiceRef.current; + if (state.kind === "selecting") { + const options = visibleFieldsRef.current[state.visibleIndex]?.options ?? []; if (key.escape) { - cancelSelect(); - return; - } - if (key.upArrow) { - setFstate((s) => - s.kind === "selecting" - ? { ...s, selection: s.selection > 0 ? s.selection - 1 : 0 } - : s, - ); - return; - } - if (key.downArrow) { - setFstate((s) => - s.kind === "selecting" - ? { - ...s, - selection: - s.selection < opts.length - 1 - ? s.selection + 1 - : s.selection, - } - : s, - ); - return; - } - if (key.leftArrow || key.return) { - commitSelect(st.selection); - return; + setChoice({ kind: "idle" }); + } else if (key.upArrow) { + setChoice({ ...state, selection: Math.max(state.selection - 1, 0) }); + } else if (key.downArrow) { + setChoice({ + ...state, + selection: Math.min(state.selection + 1, Math.max(options.length - 1, 0)), + }); + } else if (key.return) { + commitSelect(state.selection); } return; } - // ── multi-selecting options ── - if (st.kind === "multi-selecting") { - const f = visibleRef.current[st.visibleIndex]; - const opts = f?.options ?? []; + if (state.kind === "multi-selecting") { + const options = visibleFieldsRef.current[state.visibleIndex]?.options ?? []; if (key.escape) { - cancelMultiSelect(); - return; - } - if (key.upArrow) { - setFstate((s) => - s.kind === "multi-selecting" - ? { ...s, selection: s.selection > 0 ? s.selection - 1 : 0 } - : s, - ); - return; - } - if (key.downArrow) { - setFstate((s) => - s.kind === "multi-selecting" - ? { - ...s, - selection: - s.selection < opts.length - 1 - ? s.selection + 1 - : s.selection, - } - : s, - ); - return; - } - if (input === " ") { - setFstate((s) => { - if (s.kind !== "multi-selecting") { - return s; - } - const opt = opts[s.selection]; - if (!opt) { - return s; - } - const next = new Set(s.chosen); - if (next.has(opt.value)) { - next.delete(opt.value); - } else { - next.add(opt.value); - } - return { ...s, chosen: next }; + setChoice({ kind: "idle" }); + } else if (key.upArrow) { + setChoice({ ...state, selection: Math.max(state.selection - 1, 0) }); + } else if (key.downArrow) { + setChoice({ + ...state, + selection: Math.min(state.selection + 1, Math.max(options.length - 1, 0)), }); - return; - } - if (key.leftArrow || key.return) { + } else if (input === " ") { + toggleMultiSelect(state.selection); + } else if (key.return) { commitMultiSelect(); - return; } return; } - // ── idle navigation ── - const c = cursorRef.current; - + const current = cursorRef.current; + const field = visibleFieldsRef.current[current]; + const editingText = field?.kind === "text" && !field.readOnly; + if (!editingText && onShortcut?.(input, key, event, valuesRef.current)) { + return; + } if (key.escape) { onClose(); return; } - - // navigate among visible fields + actions if (key.upArrow) { - setCursor((prev) => Math.max(prev - 1, 0)); + event.preventDefault(); + moveCursor(-1); return; } - if (key.downArrow) { - setCursor((prev) => Math.min(prev + 1, actionEnd)); + if (key.downArrow || key.tab) { + event.preventDefault(); + moveCursor(1); return; } - // cursor on action row - if (c >= actionStart && c <= actionEnd) { + if (current >= actionStart && current <= actionEnd) { if (key.leftArrow) { - setCursor((prev) => (prev > actionStart ? prev - 1 : actionEnd)); - return; - } - if (key.rightArrow) { - setCursor((prev) => (prev < actionEnd ? prev + 1 : actionStart)); - return; + moveCursor(-1); + } else if (key.rightArrow) { + moveCursor(1); + } else if (key.return) { + runAction(current - actionStart); } - if (key.return) { - const act = actDefsRef.current[c - actionStart]; - if (act.key === "cancel") { - onClose(); - return; - } - onAction(act.key, valuesRef.current); - return; - } - return; - } - - // cursor on a visible field - const vf = visibleRef.current[c]; - if (!vf) { return; } - if (vf.readOnly) { + if (!field || field.focusable === false || field.readOnly || editingText) { return; } - - if (vf.kind === "boolean") { - if (key.leftArrow || key.rightArrow) { - activateField(c); + if (field.kind === "boolean") { + if (key.leftArrow || key.rightArrow || key.return || input === " ") { + const value = valuesRef.current[field.key] ?? field.value; + updateValue(field.key, value === "true" ? "false" : "true"); } - return; + } else if ((field.kind === "select" || field.kind === "multiselect") && key.return) { + openChoice(current); } - - if (vf.kind === "select") { - if (key.rightArrow || key.return) { - activateField(c); - } - return; - } - - if (vf.kind === "multiselect") { - if (key.rightArrow || key.return) { - activateField(c); - } - return; - } - - if (vf.kind === "text") { - if (key.return) { - activateField(c); - } - return; - } - }); - - // ── render ──────────────────────────────────────────────────────── + }, { isActive: active }); + + const hint = hintOverride ?? (dialogSize.width < 60 + ? "↑↓ move · Enter choose · Esc back" + : "↑↓ navigate · type to edit · Enter open/choose · Space toggle · Esc back"); + const choiceField = choice.kind === "idle" + ? undefined + : visibleFields[choice.visibleIndex]; + const choiceOptions = choiceField?.options ?? []; + const choiceSelection = choice.kind === "idle" ? 0 : choice.selection; + const choiceOptionStart = Math.max( + 0, + Math.min( + choiceSelection - Math.floor(maxExpandedOptions / 2), + Math.max(choiceOptions.length - maxExpandedOptions, 0), + ), + ); + const choiceVisibleOptions = choiceOptions.slice( + choiceOptionStart, + choiceOptionStart + maxExpandedOptions, + ); return ( - - - {title} - - - - {visibleFields.map((f, vi) => { - const isActive = cursor === vi; - const st = fstate; - const isEditing = st.kind === "editing" && st.visibleIndex === vi; - const isSelecting = - st.kind === "selecting" && st.visibleIndex === vi; - const isMultiSelecting = - st.kind === "multi-selecting" && st.visibleIndex === vi; - const curVal = values[f.key] ?? f.value; - const expandedSelection = isSelecting - ? st.kind === "selecting" - ? st.selection - : 0 - : isMultiSelecting && st.kind === "multi-selecting" - ? st.selection - : 0; - const optionCount = f.options?.length ?? 0; - const optionStart = Math.max( - 0, - Math.min( - expandedSelection - Math.floor(maxExpandedOptions / 2), - Math.max(optionCount - maxExpandedOptions, 0), - ), - ); - const visibleOptions = - f.options?.slice(optionStart, optionStart + maxExpandedOptions) ?? []; + = actionStart ? actionDefs[cursor - actionStart]?.key : undefined} + gap={3} + onAction={(key) => { + const index = actionDefs.findIndex((action) => action.key === key); + if (index >= 0) { + setCursor(actionStart + index); + runAction(index); + } + }} + /> + )} + hint={shortcutHint ? `${hint} · ${shortcutHint}` : hint} + > + + {visibleFields.map((field, visibleIndex) => { + const selected = cursor === visibleIndex; + const value = values[field.key] ?? field.value; + const options = field.options ?? []; + const editingText = active && selected && field.kind === "text" && !field.readOnly; return ( - { - if (fstate.kind === "idle") { - setCursor(vi); + { + if (field.focusable === false) { + return; } - }} - onMouseClick={() => { - if (!isEditing) { - activateField(vi); + setCursor(visibleIndex); + setChoice({ kind: "idle" }); + if (field.kind === "boolean" && !field.readOnly) { + updateValue(field.key, value === "true" ? "false" : "true"); } }} > - {/* label + value row */} - - - - {isActive && !isEditing ? "\u276f" : " "} - - - - - {pad(f.label + ":", keyWidth)} - - - - {isEditing ? ( - - updateValue(f.key, v)} - onSubmit={(v) => commitEdit(v)} - placeholder={f.placeholder ?? ""} - /> - + + + {selected ? "❯" : " "} + + + + + {truncateText(`${field.label}:`, keyWidth)} + + + + + {editingText ? ( + updateValue(field.key, next)} + onSubmit={() => moveCursor(1, visibleIndex)} + placeholder={field.placeholder ?? ""} + focus={active} + onKeyDown={(event) => { + if (event.name === "up") { + event.preventDefault(); + event.stopPropagation(); + moveCursor(-1, visibleIndex); + } else if (event.name === "down" || event.name === "tab") { + event.preventDefault(); + event.stopPropagation(); + moveCursor(1, visibleIndex); + } else if (event.name === "escape") { + event.preventDefault(); + event.stopPropagation(); + onClose(); + } + }} + /> ) : ( - - {f.kind === "boolean" - ? renderBoolean(isActive, curVal) - : f.kind === "select" - ? selectLabel(f.options ?? [], curVal) - : f.kind === "multiselect" - ? renderMultiValue(curVal) - : f.secret - ? maskValue(curVal) - : curVal || ( - - {"\u2014"} - - )} + + {field.kind === "boolean" + ? renderBoolean(selected, value) + : field.kind === "select" + ? selectLabel(options, value) + : field.kind === "multiselect" + ? renderMultiValue(value) + : field.secret + ? maskValue(value) + : value || } )} - - {/* expanded select options */} - {isSelecting && - f.kind === "select" && - f.options ? ( - - {visibleOptions.map((opt, localIndex) => { - const oi = optionStart + localIndex; - const sel = - st.kind === "selecting" && - st.selection === oi; - return ( - { - setFstate((state) => - state.kind === "selecting" - ? { ...state, selection: oi } - : state, - ); - }} - onMouseClick={() => { - commitSelect(oi); - }} - > - - {sel - ? "\u276f " - : " "} - - - {opt.label} - - - ); - })} - - ) : null} - - {/* expanded multiselect options */} - {isMultiSelecting && - f.kind === "multiselect" && - f.options ? ( - - {visibleOptions.map((opt, localIndex) => { - const oi = optionStart + localIndex; - const sel = - st.kind === "multi-selecting" && - st.selection === oi; - const checked = - st.kind === "multi-selecting" && - st.chosen.has(opt.value); - return ( - { - setFstate((state) => - state.kind === "multi-selecting" - ? { ...state, selection: oi } - : state, - ); - }} - onMouseClick={() => { - setFstate((state) => { - if (state.kind !== "multi-selecting") { - return state; - } - const next = new Set(state.chosen); - if (next.has(opt.value)) { - next.delete(opt.value); - } else { - next.add(opt.value); - } - return { ...state, selection: oi, chosen: next }; - }); - }} - > - - {sel ? "❯ " : " "} - - - {checked ? "✓ " : "☐ "} - {opt.label} - - - ); - })} - - ) : null} - + ); })} - - - {/* action row */} - - {actDefs.map((act, ai) => { - const ac = actionStart + ai; - const active = cursor === ac; - return ( + {choice.kind === "idle" ? null : ( + + { - setCursor(ac); - if (act.key === "cancel") { - onClose(); - } else { - onAction(act.key, valuesRef.current); - } - }} + flexDirection="column" + flexGrow={1} + flexBasis={0} + height={maxExpandedOptions} + backgroundColor="#101417" > - - {active ? "\u276f " : " "} - {act.label} - + {Array.from({ length: maxExpandedOptions }, (_, localIndex) => { + const option = choiceVisibleOptions[localIndex]; + const index = choiceOptionStart + localIndex; + const activeOption = option !== undefined && choice.selection === index; + const checked = option !== undefined && choice.kind === "multi-selecting" && + choice.chosen.has(option.value); + return ( + { + if (!option) { + return; + } + if (choice.kind === "selecting") { + commitSelect(index); + } else { + toggleMultiSelect(index); + } + }} + > + + {option && activeOption ? "❯ " : " "} + + + {choice.kind === "multi-selecting" + ? `${checked ? "✓" : "☐"} ${option?.label ?? ""}` + : option?.label ?? ""} + + + ); + })} - ); - })} + + )} - - {/* hints */} - - - {dialogSize.width < 60 - ? "\u2191\u2193 move · \u2190\u2192 change · Enter select · Esc back" - : "\u2191\u2193 navigate · \u2190\u2192 toggle · Enter edit/choose · Space toggle · Esc back"} - - - + ); } -// ─── boolean render helper ─────────────────────────────────────────── - -function renderBoolean(active: boolean, v: string): React.ReactNode { - const isTrue = v === "true"; +function renderBoolean(selected: boolean, value: string): React.ReactNode { + const enabled = value === "true"; return ( - <> - - {isTrue ? "\u25c9 true" : "\u25cb false"} - - + + {enabled ? "◉ true" : "○ false"} + ); } function renderMultiValue(value: string): React.ReactNode { const chosen = parseMulti(value); if (chosen.size === 0) { - return {"\u2014"}; + return ; } return {`${chosen.size} selected`}; } diff --git a/packages/agenty-cli/src/components/List.tsx b/packages/agenty-cli/src/components/List.tsx new file mode 100644 index 0000000..2942852 --- /dev/null +++ b/packages/agenty-cli/src/components/List.tsx @@ -0,0 +1,185 @@ +import type { KeyEvent } from "@opentui/core"; +import type { ReactNode } from "react"; + +import type { InputKey } from "../hooks/useInput"; +import { useInput } from "../hooks/useInput"; +import { + createTableLayout, + type TableColumn, + TableHeader, + TableRow, +} from "./Table"; +import { Box, Pressable, Text } from "./ui"; + +export interface ListRenderState { + index: number; + selected: boolean; +} + +export interface ListProps { + items: T[]; + cursor: number; + visibleCount: number; + getKey: (item: T, index: number) => string; + renderItem: (item: T, state: ListRenderState) => ReactNode; + onCursor: (index: number) => void; + onActivate?: (item: T, index: number) => void; + active?: boolean; + emptyHint?: ReactNode; +} + +export function listWindow( + itemCount: number, + cursor: number, + visibleCount: number, +): { start: number; end: number } { + const count = Math.max(Math.min(visibleCount, itemCount), 0); + const boundedCursor = Math.min(Math.max(cursor, 0), Math.max(itemCount - 1, 0)); + const half = Math.floor(count / 2); + let start = Math.max(boundedCursor - half, 0); + if (start + count > itemCount) { + start = Math.max(itemCount - count, 0); + } + return { start, end: start + count }; +} + +export function List({ + items, + cursor, + visibleCount, + getKey, + renderItem, + onCursor, + onActivate, + active = true, + emptyHint, +}: ListProps) { + if (items.length === 0) { + return ( + + {emptyHint ?? No items.} + + ); + } + + const { start, end } = listWindow(items.length, cursor, visibleCount); + return ( + + {items.slice(start, end).map((item, localIndex) => { + const index = start + localIndex; + const selected = active && index === cursor; + return ( + { + onCursor(index); + onActivate?.(item, index); + }} + > + + + {selected ? "❯" : " "} + + + {renderItem(item, { index, selected })} + + ); + })} + + ); +} + +export interface ListNavigationOptions { + items: T[]; + cursor: number; + onCursor: (index: number) => void; + onActivate?: (item: T, index: number) => void; + onClose?: () => void; + onInput?: ( + input: string, + key: InputKey, + event: KeyEvent, + item: T | undefined, + ) => void; + active?: boolean; +} + +export function useListNavigation({ + items, + cursor, + onCursor, + onActivate, + onClose, + onInput, + active = true, +}: ListNavigationOptions) { + useInput((input, key, event) => { + if (key.escape && onClose) { + onClose(); + return; + } + if (key.upArrow) { + event.preventDefault(); + onCursor(Math.max(cursor - 1, 0)); + return; + } + if (key.downArrow) { + event.preventDefault(); + onCursor(Math.min(cursor + 1, Math.max(items.length - 1, 0))); + return; + } + + const item = items[cursor]; + if (key.return && item && onActivate) { + onActivate(item, cursor); + return; + } + onInput?.(input, key, event, item); + }, { isActive: active }); +} + +export interface KeyValueRow { + key: string; + value: string; +} + +export function KeyValueList({ + rows, + availableWidth, +}: { + rows: KeyValueRow[]; + availableWidth: number; +}) { + const columns: Array> = [ + { + key: "key", + header: "Key", + value: (row) => row.key, + render: (row) => {row.key}, + }, + { + key: "value", + header: "Value", + value: (row) => row.value, + render: (row) => {row.value}, + }, + ]; + const tableLayout = createTableLayout(columns, rows, availableWidth, 1); + return ( + + + {rows.map((row) => ( + + ))} + + ); +} diff --git a/packages/agenty-cli/src/components/ModelOverlay.test.ts b/packages/agenty-cli/src/components/ModelOverlay.test.ts index 10ed6ef..005f1aa 100644 --- a/packages/agenty-cli/src/components/ModelOverlay.test.ts +++ b/packages/agenty-cli/src/components/ModelOverlay.test.ts @@ -3,11 +3,14 @@ import { describe, expect, test } from "bun:test"; import type { CoreModelDto, ModelProviderDto } from "../api/types"; import { isBuiltinProvider } from "../consts/providerPresets"; import { + configuredProviders, + filterModelsByCodeLike, initialModelIndex, initialProviderIndex, - modelInputFromValues, - modelUpdateFromValues, - parseReasoningMapping, + isConfiguredProvider, + isModelSearchInput, + normalizeModelSearchInput, + sortModelsByCode, } from "./ModelOverlay"; function model(code: string, isDefault = false): CoreModelDto { @@ -22,13 +25,19 @@ function model(code: string, isDefault = false): CoreModelDto { }; } -function provider(code: string, models: CoreModelDto[] = []): ModelProviderDto { +function provider( + code: string, + models: CoreModelDto[] = [], + apiKey = "test-key", +): ModelProviderDto { return { code, name: code, type: "openai_completions", baseUrl: "https://example.invalid/v1", - apiKey: "test-key", + apiKey, + builtin: code === "openai", + official: code === "openai", models, createdAt: "", updatedAt: "", @@ -41,8 +50,8 @@ describe("model overlay behavior", () => { expect(initialProviderIndex(providers, "openai")).toBe(1); expect(initialProviderIndex(providers, "missing")).toBe(0); - expect(isBuiltinProvider("openai")).toBe(true); - expect(isBuiltinProvider("custom")).toBe(false); + expect(isBuiltinProvider(provider("openai"))).toBe(true); + expect(isBuiltinProvider(provider("custom"))).toBe(false); }); test("starts on the current model, then provider default, then first model", () => { @@ -54,47 +63,65 @@ describe("model overlay behavior", () => { expect(initialModelIndex(provider("empty"))).toBe(0); }); - test("parses and validates reasoning mappings", () => { - expect(parseReasoningMapping("{\"fast\":\"low\",\"deep\":\"high\"}")).toEqual({ - fast: "low", - deep: "high", - }); - expect(parseReasoningMapping("{}")).toEqual({}); - expect(parseReasoningMapping(" ")).toBeUndefined(); - expect(() => parseReasoningMapping("[]")).toThrow("JSON object"); - expect(() => parseReasoningMapping("{\"fast\":\"unsupported\"}")).toThrow("Invalid reasoning effort"); + test("sorts models by code without mutating the provider response", () => { + const models = [model("zeta"), model("Alpha"), model("beta")]; + + expect(sortModelsByCode(models).map((candidate) => candidate.code)).toEqual([ + "Alpha", + "beta", + "zeta", + ]); + expect(models.map((candidate) => candidate.code)).toEqual(["zeta", "Alpha", "beta"]); + expect(configuredProviders([provider("custom", models)])[0]?.models.map( + (candidate) => candidate.code, + )).toEqual(["Alpha", "beta", "zeta"]); + }); + + test("filters sorted model codes with a case-insensitive contains match", () => { + const models = sortModelsByCode([ + model("gpt-5"), + model("GPT-4o"), + model("openrouter:google/gemini-2.5-pro"), + { ...model("claude-3"), name: "gpt display name" }, + ]); + + expect(filterModelsByCodeLike(models, "pt-").map((candidate) => candidate.code)).toEqual([ + "GPT-4o", + "gpt-5", + ]); + expect(filterModelsByCodeLike(models, "AUdE").map((candidate) => candidate.code)).toEqual([ + "claude-3", + ]); + expect(filterModelsByCodeLike(models, ":GOOGLE/GEMINI-2.5").map( + (candidate) => candidate.code, + )).toEqual(["openrouter:google/gemini-2.5-pro"]); + expect(filterModelsByCodeLike(models, "display")).toEqual([]); + expect(filterModelsByCodeLike(models, "")).toBe(models); }); - test("builds create and update payloads without changing model codes", () => { - const values = { - code: "org/model_name[v2]", - name: "Model v2", - contextWindow: "64000", - multiModal: "true", - light: "false", - isDefault: "true", - reasoningMapping: "{\"deep\":\"xhigh\"}", - }; - const created = modelInputFromValues("custom", values); - expect(created).toMatchObject({ - providerCode: "custom", - modelCode: "org/model_name[v2]", - contextWindow: 64_000, - multiModal: true, - isDefault: true, - reasoningEffortMapping: { deep: "xhigh" }, - }); - - const updated = modelUpdateFromValues(model("old-id"), { - ...values, - name: "Updated model", - code: "new-id", - }); - expect(updated).toMatchObject({ - name: "Updated model", - contextWindow: 64_000, - multiModal: true, - }); + test("only accepts model code characters as automatic search input", () => { + const namespacedCode = "openrouter:google/gemini-2.5-pro"; + + expect(isModelSearchInput("Az09-_ /\\:.".replace(" ", ""))).toBe(true); + expect(isModelSearchInput(namespacedCode)).toBe(true); + expect(normalizeModelSearchInput(namespacedCode)).toBe(namespacedCode); + expect(normalizeModelSearchInput("gpt 4.1")).toBe("gpt4.1"); + expect(isModelSearchInput("model name")).toBe(false); + expect(isModelSearchInput("")).toBe(false); + }); + + test("only treats providers with a non-blank API key as configured", () => { + const providers = [ + provider("configured"), + provider("empty", [], ""), + provider("blank", [], " "), + ]; + + expect(isConfiguredProvider(providers[0]!)).toBe(true); + expect(isConfiguredProvider(providers[1]!)).toBe(false); + expect(configuredProviders(providers).map((candidate) => candidate.code)).toEqual([ + "configured", + ]); }); }); diff --git a/packages/agenty-cli/src/components/ModelOverlay.tsx b/packages/agenty-cli/src/components/ModelOverlay.tsx index a71b52a..a5eab28 100644 --- a/packages/agenty-cli/src/components/ModelOverlay.tsx +++ b/packages/agenty-cli/src/components/ModelOverlay.tsx @@ -1,34 +1,61 @@ import { useCallback, useEffect, useRef, useState } from "react"; -import type { - CoreModelDto, - CreateModelDto, - ModelProviderDto, - ReasoningEffort, - UpdateModelDto, -} from "../api/types"; -import { isBuiltinProvider } from "../consts/providerPresets"; +import type { CoreModelDto, ModelProviderDto, ModelRef } from "../api/types"; import { useInput } from "../hooks/useInput"; import { useAppStore } from "../state/store"; import { useBottomDialogSize } from "./BottomDialog"; -import type { FormField } from "./FormPanel"; -import { FormPanel } from "./FormPanel"; -import { Box, Spinner, Text } from "./ui"; +import { List, useListNavigation } from "./List"; +import { Panel } from "./Panel"; +import { + createTableLayout, + type TableColumn, + TableHeader, + TableRow, +} from "./Table"; +import { Box, Pressable, Spinner, Text, TextInput } from "./ui"; + +export function isConfiguredProvider(provider: Pick): boolean { + return provider.apiKey.trim() !== ""; +} + +export function configuredProviders(providers: ModelProviderDto[]): ModelProviderDto[] { + return providers + .filter(isConfiguredProvider) + .map((provider) => ({ + ...provider, + models: sortModelsByCode(provider.models), + })); +} + +export function sortModelsByCode(models: CoreModelDto[]): CoreModelDto[] { + return [...models].sort((left, right) => { + const insensitiveOrder = left.code.localeCompare(right.code, "en", { + sensitivity: "base", + }); + return insensitiveOrder !== 0 + ? insensitiveOrder + : left.code.localeCompare(right.code, "en"); + }); +} -const REASONING_EFFORTS: readonly ReasoningEffort[] = [ - "off", - "low", - "medium", - "high", - "xhigh", - "max", -]; +export function filterModelsByCodeLike( + models: CoreModelDto[], + query: string, +): CoreModelDto[] { + const normalizedQuery = query.trim().toLocaleLowerCase(); + if (normalizedQuery === "") { + return models; + } + return models.filter((model) => model.code.toLocaleLowerCase().includes(normalizedQuery)); +} + +export function normalizeModelSearchInput(input: string): string { + return input.replace(/[^A-Za-z0-9_./\\:-]/g, ""); +} -type Mode = - | { kind: "list" } - | { kind: "create"; provider: ModelProviderDto } - | { kind: "edit"; provider: ModelProviderDto; target: CoreModelDto } - | { kind: "confirm-delete"; provider: ModelProviderDto; target: CoreModelDto }; +export function isModelSearchInput(input: string): boolean { + return input !== "" && normalizeModelSearchInput(input) === input; +} export function initialProviderIndex( providers: ModelProviderDto[], @@ -61,161 +88,33 @@ export function initialModelIndex( return defaultIndex >= 0 ? defaultIndex : 0; } -export function parseReasoningMapping(raw: string): Record | undefined { - const trimmed = raw.trim(); - if (!trimmed) { - return undefined; - } - - let parsed: unknown; - try { - parsed = JSON.parse(trimmed); - } catch { - throw new Error("Reasoning mapping must be a JSON object."); - } - if (!parsed || typeof parsed !== "object" || Array.isArray(parsed)) { - throw new Error("Reasoning mapping must be a JSON object."); - } - - const mapping: Record = {}; - for (const [nativeEffort, agentyEffort] of Object.entries(parsed)) { - if (!nativeEffort.trim()) { - throw new Error("Reasoning mapping keys cannot be empty."); - } - if (typeof agentyEffort !== "string" || !REASONING_EFFORTS.includes(agentyEffort as ReasoningEffort)) { - throw new Error(`Invalid reasoning effort for ${nativeEffort}.`); - } - mapping[nativeEffort] = agentyEffort as ReasoningEffort; - } - return mapping; -} - -function serializeReasoningMapping(mapping?: Record): string { - return mapping && Object.keys(mapping).length > 0 ? JSON.stringify(mapping) : ""; -} - -function parsePositiveInteger(raw: string, label: string): number { - const value = Number(raw.trim()); - if (!Number.isSafeInteger(value) || value <= 0) { - throw new Error(`${label} must be a positive integer.`); - } - return value; -} - -export function modelInputFromValues( - providerCode: string, - values: Record, -): CreateModelDto { - const modelCode = values.code.trim(); - const name = values.name.trim(); - if (!modelCode) { - throw new Error("Model Code is required."); - } - if (!name) { - throw new Error("Model name is required."); - } - - return { - providerCode, - modelCode, - name, - contextWindow: parsePositiveInteger(values.contextWindow, "Context window"), - multiModal: values.multiModal === "true", - light: values.light === "true", - isDefault: values.isDefault === "true", - reasoningEffortMapping: parseReasoningMapping(values.reasoningMapping), - }; -} - -export function modelUpdateFromValues( - target: CoreModelDto, - values: Record, -): UpdateModelDto { - const input = modelInputFromValues(target.code, values); - return { - name: input.name, - contextWindow: input.contextWindow, - multiModal: input.multiModal, - light: input.light, - isDefault: input.isDefault, - reasoningEffortMapping: input.reasoningEffortMapping, - }; -} - -function modelFields(target?: CoreModelDto): FormField[] { - return [ - { - key: "code", - label: "Model Code", - kind: "text", - value: target?.code ?? "", - placeholder: "model-code or org/model-code", - readOnly: !!target, - }, - { - key: "name", - label: "Model name", - kind: "text", - value: target?.name ?? "", - placeholder: "Model name", - }, - { - key: "contextWindow", - label: "Context window", - kind: "text", - value: target ? String(target.contextWindow) : "128000", - placeholder: "128000", - }, - { - key: "multiModal", - label: "Multimodal", - kind: "boolean", - value: target?.multiModal ? "true" : "false", - }, - { - key: "light", - label: "Light", - kind: "boolean", - value: target?.light ? "true" : "false", - }, - { - key: "isDefault", - label: "Default", - kind: "boolean", - value: target?.isDefault ? "true" : "false", - }, - { - key: "reasoningMapping", - label: "Reasoning map", - kind: "text", - value: serializeReasoningMapping(target?.reasoningEffortMapping), - placeholder: "{\"native\":\"medium\"}", - }, - ]; -} - export function ModelOverlay() { const client = useAppStore((state) => state.client); const sessionModel = useAppStore((state) => state.session?.currentModel); const setToast = useAppStore((state) => state.setToast); const setOverlay = useAppStore((state) => state.setOverlay); const switchModel = useAppStore((state) => state.switchModel); - const [providers, setProviders] = useState(null); const [providerCursor, setProviderCursor] = useState(0); const [modelCursor, setModelCursor] = useState(0); - const [mode, setMode] = useState({ kind: "list" }); - const modeRef = useRef(mode); - modeRef.current = mode; + const [loadingProviderCode, setLoadingProviderCode] = useState(null); const selectedProviderCodeRef = useRef(sessionModel?.providerCode); const selectedModelCodeRef = useRef(sessionModel?.modelCode); + const requestIdRef = useRef(0); - const reload = useCallback(async () => { + const reload = useCallback(async (providerCode?: string) => { if (!client) { return; } + const requestId = ++requestIdRef.current; + if (providerCode) { + setLoadingProviderCode(providerCode); + } try { - const list = await client.listProviders(); + const list = configuredProviders(await client.listProviders(providerCode)); + if (requestId !== requestIdRef.current) { + return; + } const nextProviderIndex = initialProviderIndex( list, selectedProviderCodeRef.current, @@ -225,165 +124,71 @@ export function ModelOverlay() { nextProvider, selectedModelCodeRef.current, ); - const nextModel = nextProvider?.models[nextModelIndex]; selectedProviderCodeRef.current = nextProvider?.code; - selectedModelCodeRef.current = nextModel?.code; + selectedModelCodeRef.current = nextProvider?.models[nextModelIndex]?.code; setProviders(list); setProviderCursor(nextProviderIndex); setModelCursor(nextModelIndex); } catch (error) { + if (requestId !== requestIdRef.current) { + return; + } setToast(`failed to load models: ${(error as Error).message}`, true); setProviders([]); + } finally { + if (requestId === requestIdRef.current) { + setLoadingProviderCode(null); + } } }, [client, setToast]); useEffect(() => { - void reload(); + void reload(selectedProviderCodeRef.current); }, [reload]); const close = () => setOverlay(null); - - useInput((input, key) => { - if (modeRef.current.kind !== "list") { - return; - } - if (providers !== null && providers.length > 0) { - return; - } - if (key.escape) { - close(); - return; - } - if (input.toLowerCase() === "a") { - setToast("No provider is available for a new model.", true); - } - }); - const selectProvider = (index: number) => { const provider = providers?.[index]; if (!provider) { return; } - selectedProviderCodeRef.current = provider.code; const nextModelIndex = initialModelIndex(provider); + selectedProviderCodeRef.current = provider.code; selectedModelCodeRef.current = provider.models[nextModelIndex]?.code; setProviderCursor(index); setModelCursor(nextModelIndex); + void reload(provider.code); }; - const selectModelCursor = (index: number) => { - const provider = providers?.[providerCursor]; - const model = provider?.models[index]; + const model = providers?.[providerCursor]?.models[index]; if (model) { selectedModelCodeRef.current = model.code; } setModelCursor(index); }; - const handleSwitch = async (model: CoreModelDto) => { - if (!client || !providers) { - return; - } - const provider = providers[providerCursor]; + const provider = providers?.[providerCursor]; if (!provider) { return; } - const projected = { + await switchModel({ ...model, providerCode: provider.code, providerName: provider.name, - }; - await switchModel(projected); - }; - - const handleCreate = async (provider: ModelProviderDto, values: Record) => { - if (!client) { - return; - } - try { - const input = modelInputFromValues(provider.code, values); - const created = await client.createModel(input); - selectedModelCodeRef.current = created.code; - await reload(); - setMode({ kind: "list" }); - setToast(`Model added: ${provider.name} · ${created.name}`); - } catch (error) { - setToast(`add model failed: ${(error as Error).message}`, true); - } + }); }; - const handleUpdate = async ( - provider: ModelProviderDto, - target: CoreModelDto, - values: Record, - ) => { - if (!client) { - return; - } - try { - const update = modelUpdateFromValues(target, values); - const updated = await client.updateModel(provider.code, target.code, update); - selectedModelCodeRef.current = updated.code; - await reload(); - setMode({ kind: "list" }); - setToast(`Model updated: ${provider.name} · ${updated.name}`); - } catch (error) { - setToast(`update model failed: ${(error as Error).message}`, true); - } - }; - - const handleDelete = async (provider: ModelProviderDto, target: CoreModelDto) => { - if (!client) { + useInput((_input, key) => { + if (providers !== null && providers.length > 0) { return; } - try { - await client.deleteModel(provider.code, target.code); - selectedModelCodeRef.current = undefined; - await reload(); - setMode({ kind: "list" }); - setToast(`Model deleted: ${provider.name} · ${target.name}`); - } catch (error) { - setToast(`delete model failed: ${(error as Error).message}`, true); + if (key.escape) { + close(); } - }; - - if (mode.kind === "create") { - return ( - void handleCreate(mode.provider, values)} - onClose={() => setMode({ kind: "list" })} - /> - ); - } - - if (mode.kind === "edit") { - return ( - void handleUpdate(mode.provider, mode.target, values)} - onClose={() => setMode({ kind: "list" })} - /> - ); - } - - if (mode.kind === "confirm-delete") { - return ( - void handleDelete(mode.provider, mode.target)} - onCancel={() => setMode({ kind: "list" })} - /> - ); - } + }); return ( - - - Models - + {providers === null ? ( ) : providers.length === 0 ? ( @@ -394,45 +199,21 @@ export function ModelOverlay() { providerCursor={providerCursor} modelCursor={modelCursor} currentModel={sessionModel} + loadingProviderCode={loadingProviderCode} onProviderCursor={selectProvider} onModelCursor={selectModelCursor} onSwitch={(model) => void handleSwitch(model)} - onAdd={(provider) => setMode({ kind: "create", provider })} - onEdit={(provider, target) => setMode({ kind: "edit", provider, target })} - onDelete={(provider, target) => setMode({ kind: "confirm-delete", provider, target })} onClose={close} /> )} - + ); } -function ModelForm({ - title, - fields, - onSave, - onClose, -}: { - title: string; - fields: FormField[]; - onSave: (values: Record) => void; - onClose: () => void; -}) { - return ( - { - if (action === "save") { - onSave(values); - } else { - onClose(); - } - }} - onClose={onClose} - /> - ); +function modelLabel(model: CoreModelDto): string { + return model.name && model.code + ? `${model.name} · ${model.code}` + : model.code; } function ModelList({ @@ -440,246 +221,197 @@ function ModelList({ providerCursor, modelCursor, currentModel, + loadingProviderCode, onProviderCursor, onModelCursor, onSwitch, - onAdd, - onEdit, - onDelete, onClose, }: { providers: ModelProviderDto[]; providerCursor: number; modelCursor: number; - currentModel?: { providerCode: string; modelCode: string }; + currentModel?: ModelRef; + loadingProviderCode: string | null; onProviderCursor: (index: number) => void; onModelCursor: (index: number) => void; onSwitch: (model: CoreModelDto) => void; - onAdd: (provider: ModelProviderDto) => void; - onEdit: (provider: ModelProviderDto, model: CoreModelDto) => void; - onDelete: (provider: ModelProviderDto, model: CoreModelDto) => void; onClose: () => void; }) { const dialogSize = useBottomDialogSize(); + const [searchQuery, setSearchQuery] = useState(""); + const [searchFocused, setSearchFocused] = useState(false); const provider = providers[providerCursor]; const models = provider?.models ?? []; - const manageable = provider ? !isBuiltinProvider(provider) : false; - const compact = dialogSize.width < 66; - const providerWidth = compact - ? Math.max(dialogSize.width - 16, 16) - : Math.max(Math.min(Math.floor(dialogSize.width * 0.34), 28), 18); - const modelWidth = compact - ? 0 - : Math.max(dialogSize.width - providerWidth - 22, 18); - const maxVisible = Math.max(dialogSize.height - 8, 1); - const maxVis = Math.min(maxVisible, models.length); - const half = Math.floor(maxVis / 2); - let start = modelCursor - half; - if (start < 0) { - start = 0; - } - if (start + maxVis > models.length) { - start = Math.max(models.length - maxVis, 0); - } - const visible = models.slice(start, start + maxVis); - - useInput((input, key) => { - if (key.escape) { - onClose(); - return; - } - if (key.leftArrow || key.rightArrow) { - const direction = key.leftArrow ? -1 : 1; - const next = (providerCursor + direction + providers.length) % providers.length; - onProviderCursor(next); - return; - } - if (key.upArrow) { - onModelCursor(Math.max(modelCursor - 1, 0)); - return; - } - if (key.downArrow) { - onModelCursor(Math.min(modelCursor + 1, Math.max(models.length - 1, 0))); - return; - } - - const model = models[modelCursor]; - const lower = input.toLowerCase(); - if ((key.return || lower === "s") && model) { + const filteredModels = filterModelsByCodeLike(models, searchQuery); + const selectedModel = models[modelCursor]; + const filteredCursor = selectedModel + ? Math.max( + filteredModels.findIndex((model) => model.code === selectedModel.code), + 0, + ) + : 0; + const selectFilteredModel = () => { + const model = filteredModels[filteredCursor]; + if (model) { onSwitch(model); - return; } - if (!provider || !manageable) { - return; - } - if (lower === "a") { - onAdd(provider); - } else if (lower === "e" && model) { - onEdit(provider, model); - } else if (lower === "d" && model) { - onDelete(provider, model); + }; + const selectFilteredCursor = (index: number) => { + const model = filteredModels[index]; + if (model) { + onModelCursor(models.indexOf(model)); } - }); - - return ( - - - Provider: - { - const next = (providerCursor - 1 + providers.length) % providers.length; - onProviderCursor(next); - }}> - ‹ - - {provider?.name ?? "—"} - { - const next = (providerCursor + 1) % providers.length; - onProviderCursor(next); - }}> - › - - - {provider ? ` · ${manageable ? "custom" : "built-in"}` : ""} + }; + const moveProvider = (direction: -1 | 1) => { + onProviderCursor( + (providerCursor + direction + providers.length) % providers.length, + ); + }; + const maxVisible = Math.max(dialogSize.height - 8, 1); + const modelState = (model: CoreModelDto): string => { + const current = currentModel?.providerCode === provider?.code && + currentModel.modelCode === model.code; + return current ? "current" : model.isDefault ? "default" : ""; + }; + const columns: Array> = [ + { + key: "model", + header: "Model", + value: modelLabel, + render: (model, selected) => ( + + {modelLabel(model)} - - - - {provider - ? ` ${compact ? "Model" : `${pad("Model", modelWidth)} ${pad("Context", 10)}`} ${models.length} model${models.length === 1 ? "" : "s"}` - : "No provider selected"} + ), + }, + { + key: "context", + header: "Context", + value: (model) => model.contextWindow.toLocaleString(), + render: (model, selected) => ( + + {model.contextWindow.toLocaleString()} - - {models.length === 0 ? ( - - - {manageable - ? "No models. Press `a` to add one." - : "No models available for this built-in provider."} - - - ) : ( - - {visible.map((model) => { - const index = models.indexOf(model); - const selected = index === modelCursor; - const current = currentModel?.providerCode === provider?.code && - currentModel.modelCode === model.code; - const label = model.name && model.code - ? `${model.name} · ${model.code}` - : model.code; - return ( - onModelCursor(index)} - onMouseClick={() => onSwitch(model)} - > - - {selected ? "❯" : " "} - - - - - {label} - - - {compact ? null : } - {compact ? null : ( - - {pad(model.contextWindow.toLocaleString(), 10)} - - )} - - {current ? " · current" : model.isDefault ? " · default" : ""} - - - ); - })} - - )} - - {manageable ? ( - <> - provider && onAdd(provider)}>Add - { - const model = models[modelCursor]; - if (provider && model) { - onEdit(provider, model); - } - }}>Edit - { - const model = models[modelCursor]; - if (provider && model) { - onDelete(provider, model); - } - }}>Delete - - ) : ( - Built-in provider · switching only - )} - - - - {manageable - ? "←→ provider · ↑↓ model · Enter/s switch · a add · e edit · d delete · Esc back" - : "←→ provider · ↑↓ model · Enter/s switch · Esc back"} + ), + }, + { + key: "state", + header: "State", + value: modelState, + render: (model) => ( + + {modelState(model)} - - + ), + }, + ]; + const tableLayout = createTableLayout( + columns, + models, + Math.max(dialogSize.width - 2, 0), ); -} -function DeleteConfirm({ - target, - onConfirm, - onCancel, -}: { - target: CoreModelDto; - onConfirm: () => void; - onCancel: () => void; -}) { - useInput((input, key) => { - if (key.escape) { - onCancel(); - return; - } - const lower = input.toLowerCase(); - if (lower === "y") { - onConfirm(); - } else if (lower === "n") { - onCancel(); - } + useListNavigation({ + items: filteredModels, + cursor: filteredCursor, + onCursor: selectFilteredCursor, + onActivate: onSwitch, + onClose, + onInput: (input, key, event) => { + if (key.leftArrow || key.rightArrow) { + event.preventDefault(); + moveProvider(key.leftArrow ? -1 : 1); + return; + } + if (isModelSearchInput(input)) { + event.preventDefault(); + event.stopPropagation(); + setSearchQuery((query) => query + input); + setSearchFocused(true); + } + }, + active: !searchFocused, }); + const hint = "←→ provider · ↑↓ model · type code to search · Enter select · Esc back"; + return ( - - - Delete model "{target.name}"? - - This cannot be undone. - - [Delete] - [Cancel] + + + + Provider: + + + moveProvider(-1)}> + + + {provider?.name ?? "—"} + moveProvider(1)}> + + - + + + Search: + + + + setSearchQuery(normalizeModelSearchInput(next))} + onSubmit={selectFilteredModel} + placeholder="filter by model code" + focus={searchFocused} + onMouseDown={() => setSearchFocused(true)} + onKeyDown={(event) => { + if (event.name === "escape") { + event.preventDefault(); + event.stopPropagation(); + onClose(); + } else if (event.name === "up" || event.name === "down") { + event.preventDefault(); + event.stopPropagation(); + const direction = event.name === "up" ? -1 : 1; + const nextIndex = Math.min( + Math.max(filteredCursor + direction, 0), + Math.max(filteredModels.length - 1, 0), + ); + selectFilteredCursor(nextIndex); + } else if (event.name === "left" || event.name === "right") { + event.preventDefault(); + event.stopPropagation(); + moveProvider(event.name === "left" ? -1 : 1); + } + }} + /> + + + + + + + model.code} + onCursor={selectFilteredCursor} + onActivate={onSwitch} + emptyHint={loadingProviderCode === provider?.code + ? + : searchQuery.trim() === "" + ? No models available for this provider. + : No model code contains “{searchQuery}”.} + renderItem={(model, { selected }) => { + return ( + + ); + }} + /> + ); } - -function trunc(value: string, width: number): string { - if (width <= 0) { - return ""; - } - if (value.length <= width) { - return value; - } - if (width === 1) { - return "…"; - } - return value.slice(0, width - 1) + "…"; -} - -function pad(value: string, width: number): string { - const clipped = trunc(value, width); - return clipped + " ".repeat(Math.max(width - clipped.length, 0)); -} diff --git a/packages/agenty-cli/src/components/Panel.tsx b/packages/agenty-cli/src/components/Panel.tsx new file mode 100644 index 0000000..8c06cb1 --- /dev/null +++ b/packages/agenty-cli/src/components/Panel.tsx @@ -0,0 +1,44 @@ +import type { ReactNode } from "react"; + +import { Box, Text } from "./ui"; + +export interface PanelProps { + title?: ReactNode; + description?: ReactNode; + error?: string | null; + children: ReactNode; + footer?: ReactNode; + hint?: ReactNode; + gap?: number; +} + +export function Panel({ + title, + description, + error, + children, + footer, + hint, + gap = 0, +}: PanelProps) { + return ( + + {title ? ( + + {title} + {description ? {description} : null} + + ) : null} + {error ? {error} : null} + + {children} + + {footer ? {footer} : null} + {hint ? ( + + {hint} + + ) : null} + + ); +} diff --git a/packages/agenty-cli/src/components/ProviderOverlay.test.ts b/packages/agenty-cli/src/components/ProviderOverlay.test.ts new file mode 100644 index 0000000..7a8c1a8 --- /dev/null +++ b/packages/agenty-cli/src/components/ProviderOverlay.test.ts @@ -0,0 +1,55 @@ +import { describe, expect, test } from "bun:test"; + +import type { ModelProviderDto } from "../api/types"; +import { + buildBuiltinProviderUpdate, + buildProviderFields, +} from "./ProviderOverlay"; + +function builtinProvider(): ModelProviderDto { + return { + code: "openai", + name: "OpenAI", + type: "openai", + baseUrl: "https://api.openai.com/v1", + apiKey: "existing-key", + builtin: true, + official: true, + freeFormTool: true, + models: [], + createdAt: "", + updatedAt: "", + }; +} + +describe("provider overlay builtin configuration", () => { + test("keeps only the API key focusable for a built-in provider", () => { + const fields = buildProviderFields(builtinProvider(), "configure"); + + expect(fields.filter((field) => field.focusable !== false).map((field) => field.key)).toEqual([ + "apiKey", + ]); + expect(fields.filter((field) => field.key !== "apiKey").every((field) => field.readOnly)).toBe(true); + const apiKeyField = fields.find((field) => field.key === "apiKey"); + expect(apiKeyField?.readOnly).toBeUndefined(); + expect(apiKeyField?.value).toBe(""); + }); + + test("builds an API-key-only update and ignores blank input", () => { + expect(buildBuiltinProviderUpdate({ apiKey: " next-key " })).toEqual({ + apiKey: "next-key", + }); + expect(buildBuiltinProviderUpdate({ apiKey: " " })).toBeNull(); + }); + + test("exposes the free-form setting for custom Responses providers", () => { + const fields = buildProviderFields({ + ...builtinProvider(), + builtin: false, + }, "edit"); + const freeFormField = fields.find((field) => field.key === "freeFormTool"); + + expect(freeFormField?.value).toBe("true"); + expect(freeFormField?.readOnly).toBe(false); + }); +}); diff --git a/packages/agenty-cli/src/components/ProviderOverlay.tsx b/packages/agenty-cli/src/components/ProviderOverlay.tsx index 9b949f4..e3ec134 100644 --- a/packages/agenty-cli/src/components/ProviderOverlay.tsx +++ b/packages/agenty-cli/src/components/ProviderOverlay.tsx @@ -1,70 +1,228 @@ -import { useCallback, useEffect, useRef, useState } from "react"; +import { useCallback, useEffect, useMemo, useRef, useState } from "react"; -import type { APIType, ModelProviderDto, UpdateModelProviderDto } from "../api/types"; +import type { + APIType, + CoreModelDto, + CreateModelDto, + ModelProviderDto, + UpdateModelProviderDto, +} from "../api/types"; import { providerDefaultBaseURLs, providerTypes } from "../consts/providerTypes"; import { useInput } from "../hooks/useInput"; import { useAppStore } from "../state/store"; import { useBottomDialogSize } from "./BottomDialog"; +import { ConfirmDialog } from "./ConfirmDialog"; import type { FormField } from "./FormPanel"; import { FormPanel } from "./FormPanel"; +import { List, useListNavigation } from "./List"; +import { Panel } from "./Panel"; +import { + buildProviderRows, + defaultExpandedProviderCodes, + type ProviderListRow, + type ProviderRowEnterAction, + providerRowEnterAction, + providerRowLabel, + providerRowState, + shouldExpandProvider, +} from "./providerRows"; +import { + createTableLayout, + type TableColumn, + TableHeader, + TableRow, +} from "./Table"; import { Box, Spinner, Text } from "./ui"; -const PROVIDER_TYPE_OPTIONS = providerTypes.map((t) => ({ label: t, value: t })); +const PROVIDER_TYPE_OPTIONS = providerTypes.map((type) => ({ label: type, value: type })); -function buildCreateFields(formType: string): FormField[] { +function buildCreateProviderFields(formType: string): FormField[] { const baseUrl = providerDefaultBaseURLs[formType] ?? ""; return [ - { key: "code", label: "Provider Code", kind: "text" as const, value: "", placeholder: "my-provider" }, - { key: "name", label: "Name", kind: "text" as const, value: "", placeholder: "my-provider" }, - { key: "type", label: "Type", kind: "select" as const, value: formType, options: PROVIDER_TYPE_OPTIONS }, - { key: "baseUrl", label: "Base URL", kind: "text" as const, value: baseUrl }, - { key: "apiKey", label: "API Key", kind: "text" as const, value: "", placeholder: "sk-...", secret: true }, + { key: "code", label: "Provider Code", kind: "text", value: "", placeholder: "my-provider" }, + { key: "name", label: "Name", kind: "text", value: "", placeholder: "my-provider" }, + { key: "type", label: "Type", kind: "select", value: formType, options: PROVIDER_TYPE_OPTIONS }, + { key: "baseUrl", label: "Base URL", kind: "text", value: baseUrl }, + { + key: "freeFormTool", + label: "Free-form apply_patch", + kind: "boolean", + value: "false", + visible: formType === "openai", + }, + { key: "apiKey", label: "API Key", kind: "text", value: "", placeholder: "sk-...", secret: true }, ]; } -function buildEditFields(target: ModelProviderDto): FormField[] { +type ProviderFormMode = "edit" | "configure"; + +export function buildProviderFields( + target: ModelProviderDto, + mode: ProviderFormMode, +): FormField[] { + const configuringBuiltin = mode === "configure"; + return [ + { + key: "code", + label: "Provider Code", + kind: "text", + value: target.code, + readOnly: true, + focusable: !configuringBuiltin, + }, + { + key: "name", + label: "Name", + kind: "text", + value: target.name, + readOnly: configuringBuiltin, + focusable: !configuringBuiltin, + }, + { + key: "type", + label: "Type", + kind: "select", + value: target.type, + options: PROVIDER_TYPE_OPTIONS, + readOnly: configuringBuiltin, + focusable: !configuringBuiltin, + }, + { + key: "baseUrl", + label: "Base URL", + kind: "text", + value: target.baseUrl, + readOnly: configuringBuiltin, + focusable: !configuringBuiltin, + }, + { + key: "apiKey", + label: "API Key", + kind: "text", + value: "", + placeholder: "leave blank to keep", + secret: true, + }, + { + key: "freeFormTool", + label: "Free-form apply_patch", + kind: "boolean", + value: String(target.freeFormTool === true), + readOnly: configuringBuiltin || target.type !== "openai", + focusable: !configuringBuiltin && target.type === "openai", + visible: target.type === "openai", + }, + ]; +} + +export function buildBuiltinProviderUpdate( + values: Record, +): UpdateModelProviderDto | null { + const apiKey = values.apiKey?.trim() ?? ""; + return apiKey ? { apiKey } : null; +} + +function buildCreateModelFields(): FormField[] { + return [ + { key: "code", label: "Model Code", kind: "text", value: "", placeholder: "model-code or org/model-code" }, + { key: "name", label: "Model name", kind: "text", value: "", placeholder: "Model name" }, + { key: "contextWindow", label: "Context window", kind: "text", value: "128000", placeholder: "128000" }, + { key: "maxOutputTokens", label: "Max output tokens", kind: "text", value: "8192", placeholder: "8192" }, + { key: "multiModal", label: "Multimodal", kind: "boolean", value: "false" }, + { key: "light", label: "Light", kind: "boolean", value: "false" }, + { key: "reasoning", label: "Reasoning", kind: "boolean", value: "true" }, + ]; +} + +function buildModelFields( + model: CoreModelDto, + readOnly: boolean, +): FormField[] { return [ - { key: "code", label: "Provider Code", kind: "text" as const, value: target.code, readOnly: true }, - { key: "name", label: "Name", kind: "text" as const, value: target.name, placeholder: target.name }, - { key: "type", label: "Type", kind: "select" as const, value: target.type, options: PROVIDER_TYPE_OPTIONS }, - { key: "baseUrl", label: "Base URL", kind: "text" as const, value: target.baseUrl }, - { key: "apiKey", label: "API Key", kind: "text" as const, value: "", placeholder: "leave blank to keep", secret: true }, + { key: "code", label: "Model Code", kind: "text", value: model.code, readOnly: true }, + { key: "name", label: "Model name", kind: "text", value: model.name, readOnly }, + { + key: "contextWindow", + label: "Context window", + kind: "text", + value: String(model.contextWindow), + readOnly, + }, + { + key: "maxOutputTokens", + label: "Max output tokens", + kind: "text", + value: String(model.maxOutputTokens), + readOnly, + }, + { key: "multiModal", label: "Multimodal", kind: "boolean", value: String(model.multiModal), readOnly }, + { key: "light", label: "Light", kind: "boolean", value: String(model.light), readOnly }, + { + key: "reasoning", + label: "Reasoning", + kind: "boolean", + value: String((model.reasoningEfforts ?? []).length > 0), + readOnly, + }, ]; } -function trunc(s: string, width: number): string { - if (width <= 0) { - return ""; +function parseModelValues(values: Record): CreateModelDto | string { + const modelCode = values.code?.trim() ?? ""; + const name = values.name?.trim() ?? ""; + const contextWindow = Number(values.contextWindow); + const maxOutputTokens = Number(values.maxOutputTokens); + if (!modelCode) { + return "Model Code is required."; } - if (s.length <= width) { - return s; + if (!name) { + return "Model name is required."; } - if (width === 1) { - return "\u2026"; + if (!Number.isSafeInteger(contextWindow) || contextWindow <= 0) { + return "Context window must be a positive integer."; } - return s.slice(0, width - 1) + "\u2026"; -} - -function pad(s: string, width: number): string { - const clipped = trunc(s, width); - return clipped + " ".repeat(Math.max(width - clipped.length, 0)); + if (!Number.isSafeInteger(maxOutputTokens) || maxOutputTokens <= 0) { + return "Max output tokens must be a positive integer."; + } + return { + providerCode: "", + modelCode, + name, + contextWindow, + maxOutputTokens, + multiModal: values.multiModal === "true", + light: values.light === "true", + reasoning: values.reasoning === "true", + }; } type Mode = | { kind: "list" } - | { kind: "create" } - | { kind: "edit"; target: ModelProviderDto } - | { kind: "confirm-delete"; target: ModelProviderDto }; + | { kind: "create-provider" } + | { kind: "edit-provider"; target: ModelProviderDto } + | { kind: "configure-provider"; target: ModelProviderDto } + | { kind: "confirm-delete-provider"; target: ModelProviderDto } + | { kind: "create-model"; provider: ModelProviderDto } + | { kind: "edit-model"; provider: ModelProviderDto; target: CoreModelDto } + | { kind: "view-model"; provider: ModelProviderDto; target: CoreModelDto } + | { kind: "confirm-delete-model"; provider: ModelProviderDto; target: CoreModelDto }; -export function ProviderOverlay() { - const client = useAppStore((s) => s.client); - const setToast = useAppStore((s) => s.setToast); - const setOverlay = useAppStore((s) => s.setOverlay); +function errorMessage(cause: unknown): string { + return cause instanceof Error ? cause.message : String(cause); +} +export function ProviderOverlay() { + const client = useAppStore((state) => state.client); + const setToast = useAppStore((state) => state.setToast); + const setOverlay = useAppStore((state) => state.setOverlay); const [providers, setProviders] = useState(null); - const [cursor, setCursor] = useState(0); const [mode, setMode] = useState({ kind: "list" }); + const [selectedKey, setSelectedKey] = useState(null); + const [expandedProviderCodes, setExpandedProviderCodes] = useState>(new Set()); + const [formType, setFormType] = useState(providerTypes[0]); const modeRef = useRef(mode); + const expansionInitializedRef = useRef(false); + const previousProviderCodesRef = useRef>(new Set()); modeRef.current = mode; const reload = useCallback(async () => { @@ -74,9 +232,31 @@ export function ProviderOverlay() { try { const list = await client.listProviders(); setProviders(list); - setCursor((c) => Math.min(c, Math.max(list.length - 1, 0))); - } catch (e) { - setToast(`failed to load providers: ${(e as Error).message}`, true); + const providerCodes = new Set(list.map((provider) => provider.code)); + setExpandedProviderCodes((current) => { + if (!expansionInitializedRef.current) { + expansionInitializedRef.current = true; + previousProviderCodesRef.current = providerCodes; + return defaultExpandedProviderCodes(list); + } + const next = new Set( + list + .filter((provider) => current.has(provider.code)) + .map((provider) => provider.code), + ); + for (const provider of list) { + if ( + !previousProviderCodesRef.current.has(provider.code) && + shouldExpandProvider(provider) + ) { + next.add(provider.code); + } + } + previousProviderCodesRef.current = providerCodes; + return next; + }); + } catch (cause: unknown) { + setToast(`failed to load providers: ${errorMessage(cause)}`, true); setProviders([]); } }, [client, setToast]); @@ -85,322 +265,441 @@ export function ProviderOverlay() { void reload(); }, [reload]); - // track provider type for auto baseUrl - const [formType, setFormType] = useState(providerTypes[0]); + useInput((_input, key) => { + if (modeRef.current.kind === "list" && providers === null && key.escape) { + setOverlay(null); + } + }); const close = () => setOverlay(null); + const returnToList = () => setMode({ kind: "list" }); - // Top-level input handler — ensures empty / loading list states respond to keyboard. - useInput((input, key) => { - if (modeRef.current.kind !== "list") { - return; - } - if (providers !== null && providers.length > 0) { - return; + const openAction = (action: ProviderRowEnterAction) => { + switch (action.kind) { + case "configure-provider": + setMode({ kind: "configure-provider", target: action.provider }); + return; + case "edit-provider": + setMode({ kind: "edit-provider", target: action.provider }); + return; + case "view-model": + setMode({ kind: "view-model", provider: action.provider, target: action.model }); + return; + case "edit-model": + setMode({ kind: "edit-model", provider: action.provider, target: action.model }); + return; + case "add-provider": + setFormType(providerTypes[0]); + setMode({ kind: "create-provider" }); + return; + case "add-model": + setMode({ kind: "create-model", provider: action.provider }); + return; + case "none": + return; } - if (key.escape) { - close(); - return; - } - if (input === "a") { - setMode({ kind: "create" }); - return; - } - }); + }; - const handleCreate = async (values: Record) => { + const toggleProvider = (providerCode: string, expanded: boolean) => { + setExpandedProviderCodes((current) => { + const next = new Set(current); + if (expanded) { + next.add(providerCode); + } else { + next.delete(providerCode); + } + return next; + }); + }; + + const handleCreateProvider = async (values: Record) => { if (!client) { return; } try { + const type = values.type as APIType; await client.createProvider({ code: values.code.trim(), name: values.name.trim(), - type: values.type as APIType, + type, baseUrl: values.baseUrl.trim(), - apiKey: values.apiKey, + apiKey: values.apiKey.trim(), + freeFormTool: type === "openai" && values.freeFormTool === "true", }); setToast(`Provider created: ${values.name.trim()}`); await reload(); - } catch (e) { - setToast(`create failed: ${(e as Error).message}`, true); + returnToList(); + } catch (cause: unknown) { + setToast(`create provider failed: ${errorMessage(cause)}`, true); } - setMode({ kind: "list" }); }; - const handleEdit = async (target: ModelProviderDto, values: Record) => { - if (!client) { + const handleEditProvider = async (target: ModelProviderDto, values: Record) => { + if (!client || target.builtin === true) { + returnToList(); return; } try { - const dto: UpdateModelProviderDto = { + const type = values.type as APIType; + const update: UpdateModelProviderDto = { name: values.name.trim(), - type: values.type as APIType, + type, baseUrl: values.baseUrl.trim(), + freeFormTool: type === "openai" && values.freeFormTool === "true", }; - if (values.apiKey && values.apiKey.trim() !== "") { - dto.apiKey = values.apiKey; + if (values.apiKey.trim()) { + update.apiKey = values.apiKey.trim(); } - await client.updateProvider(target.code, dto); + await client.updateProvider(target.code, update); setToast(`Provider updated: ${values.name.trim()}`); await reload(); - } catch (e) { - setToast(`update failed: ${(e as Error).message}`, true); + returnToList(); + } catch (cause: unknown) { + setToast(`update provider failed: ${errorMessage(cause)}`, true); } - setMode({ kind: "list" }); }; - const handleDelete = async (target: ModelProviderDto) => { - if (!client) { + const handleConfigureProvider = async ( + target: ModelProviderDto, + values: Record, + ) => { + if (!client || target.builtin !== true) { + returnToList(); + return; + } + const update = buildBuiltinProviderUpdate(values); + if (!update) { + returnToList(); + return; + } + try { + await client.updateProvider(target.code, update); + setToast(`Provider API key updated: ${target.name}`); + await reload(); + returnToList(); + } catch (cause: unknown) { + setToast(`update provider API key failed: ${errorMessage(cause)}`, true); + } + }; + + const handleDeleteProvider = async (target: ModelProviderDto) => { + if (!client || target.builtin === true) { + returnToList(); return; } try { await client.deleteProvider(target.code); setToast(`Provider deleted: ${target.name}`); await reload(); - } catch (e) { - setToast(`delete failed: ${(e as Error).message}`, true); + returnToList(); + } catch (cause: unknown) { + setToast(`delete provider failed: ${errorMessage(cause)}`, true); + } + }; + + const handleCreateModel = async (provider: ModelProviderDto, values: Record) => { + if (!client || provider.builtin === true) { + returnToList(); + return; + } + const parsed = parseModelValues(values); + if (typeof parsed === "string") { + setToast(parsed, true); + return; + } + try { + await client.createModel({ ...parsed, providerCode: provider.code }); + setToast(`Model created: ${parsed.name}`); + await reload(); + returnToList(); + } catch (cause: unknown) { + setToast(`create model failed: ${errorMessage(cause)}`, true); } - setMode({ kind: "list" }); }; - if (mode.kind === "create") { + const handleEditModel = async ( + provider: ModelProviderDto, + target: CoreModelDto, + values: Record, + ) => { + if (!client || provider.builtin === true) { + returnToList(); + return; + } + const parsed = parseModelValues(values); + if (typeof parsed === "string") { + setToast(parsed, true); + return; + } + try { + await client.updateModel(provider.code, target.code, { + name: parsed.name, + contextWindow: parsed.contextWindow, + maxOutputTokens: parsed.maxOutputTokens, + multiModal: parsed.multiModal, + light: parsed.light, + reasoning: parsed.reasoning, + }); + setToast(`Model updated: ${parsed.name}`); + await reload(); + returnToList(); + } catch (cause: unknown) { + setToast(`update model failed: ${errorMessage(cause)}`, true); + } + }; + + const handleDeleteModel = async (provider: ModelProviderDto, target: CoreModelDto) => { + if (!client || provider.builtin === true) { + returnToList(); + return; + } + try { + await client.deleteModel(provider.code, target.code); + setToast(`Model deleted: ${target.name || target.code}`); + await reload(); + returnToList(); + } catch (cause: unknown) { + setToast(`delete model failed: ${errorMessage(cause)}`, true); + } + }; + + if (mode.kind === "create-provider") { return ( { + fields={buildCreateProviderFields(formType)} + onChange={(key, values) => { if (key === "type") { - setFormType(allValues.type); + setFormType(values.type); } }} onAction={(action, values) => { if (action === "save") { - handleCreate(values); + void handleCreateProvider(values); } else { - setMode({ kind: "list" }); + returnToList(); } }} - onClose={() => setMode({ kind: "list" })} + onClose={returnToList} /> ); } - if (mode.kind === "edit") { + if (mode.kind === "edit-provider" || mode.kind === "configure-provider") { + const configuringBuiltin = mode.kind === "configure-provider"; + const target = mode.target; return ( { + if (!configuringBuiltin && input.toLowerCase() === "d") { + setMode({ kind: "confirm-delete-provider", target }); + return true; + } + return false; + }} + onAction={(action, values) => { + if (action !== "save") { + returnToList(); + } else if (configuringBuiltin) { + void handleConfigureProvider(target, values); + } else { + void handleEditProvider(target, values); + } + }} + onClose={returnToList} + /> + ); + } + + if (mode.kind === "create-model") { + const provider = mode.provider; + return ( + { if (action === "save") { - void handleEdit(mode.target, values); + void handleCreateModel(provider, values); + } else { + returnToList(); + } + }} + onClose={returnToList} + /> + ); + } + + if (mode.kind === "edit-model" || mode.kind === "view-model") { + const readOnly = mode.kind === "view-model"; + const { provider, target } = mode; + return ( + { + if (!readOnly && input.toLowerCase() === "d") { + setMode({ kind: "confirm-delete-model", provider, target }); + return true; + } + return false; + }} + onAction={(action, values) => { + if (readOnly || action !== "save") { + returnToList(); } else { - setMode({ kind: "list" }); + void handleEditModel(provider, target, values); } }} - onClose={() => setMode({ kind: "list" })} + onClose={returnToList} + /> + ); + } + + if (mode.kind === "confirm-delete-provider") { + return ( + void handleDeleteProvider(mode.target)} + onCancel={returnToList} /> ); } - if (mode.kind === "confirm-delete") { + if (mode.kind === "confirm-delete-model") { return ( - void handleDelete(mode.target)} - onCancel={() => setMode({ kind: "list" })} + void handleDeleteModel(mode.provider, mode.target)} + onCancel={returnToList} /> ); } return ( - - - Providers - + {providers === null ? ( - ) : providers.length === 0 ? ( - No providers. Press `a` to add one. ) : ( setMode({ kind: "edit", target: t })} - onAdd={() => setMode({ kind: "create" })} - onDelete={(t) => { - setMode({ kind: "confirm-delete", target: t }); - }} + expandedProviderCodes={expandedProviderCodes} + selectedKey={selectedKey} + onSelectedKey={setSelectedKey} + onToggleProvider={toggleProvider} + onActivate={openAction} onClose={close} /> )} - + ); } -// ─── Provider list table ─────────────────────────────────────────── - function ProviderList({ providers, - cursor, - onCursor, - onSelect, - onAdd, - onDelete, + expandedProviderCodes, + selectedKey, + onSelectedKey, + onToggleProvider, + onActivate, onClose, }: { providers: ModelProviderDto[]; - cursor: number; - onCursor: (i: number) => void; - onSelect: (p: ModelProviderDto) => void; - onAdd: () => void; - onDelete: (p: ModelProviderDto) => void; + expandedProviderCodes: ReadonlySet; + selectedKey: string | null; + onSelectedKey: (key: string | null) => void; + onToggleProvider: (code: string, expanded: boolean) => void; + onActivate: (action: ProviderRowEnterAction) => void; onClose: () => void; }) { const dialogSize = useBottomDialogSize(); - const n = providers.length; - const compact = dialogSize.width < 46; - const nameWidth = compact - ? Math.max(dialogSize.width - 2, 8) - : Math.max(Math.min(Math.floor(dialogSize.width * 0.28), 30), 14); - const urlWidth = compact - ? 0 - : Math.max(dialogSize.width - nameWidth - 4, 8); - const maxVisible = Math.max(dialogSize.height - 7 - (compact ? 1 : 0), 1); - const maxVis = Math.min(maxVisible, n); - const half = Math.floor(maxVis / 2); - let start = cursor - half; - if (start < 0) { - start = 0; - } - if (start + maxVis > n) { - start = Math.max(n - maxVis, 0); - } - const visible = providers.slice(start, start + maxVis); - - useInput((input, key) => { - if (key.escape) { - onClose(); - return; - } - if (key.upArrow) { - onCursor(Math.max(cursor - 1, 0)); - return; - } - if (key.downArrow) { - onCursor(Math.min(cursor + 1, n - 1)); - return; - } - if (key.return) { - onSelect(providers[cursor]); - return; - } - const lower = input.toLowerCase(); - if (lower === "a") { - onAdd(); - } else if (lower === "d") { - onDelete(providers[cursor]); - } - }); - - return ( - - - - {compact - ? ` ${pad("Name", nameWidth)}` - : ` ${pad("Name", nameWidth)} ${pad("Base URL", urlWidth)}`} + const rows = useMemo( + () => buildProviderRows(providers, expandedProviderCodes), + [providers, expandedProviderCodes], + ); + const cursor = Math.max(rows.findIndex((row) => row.key === selectedKey), 0); + const columns: Array> = [ + { + key: "resource", + header: "Provider / model", + value: providerRowLabel, + render: (row, selected) => ( + + {providerRowLabel(row)} - - - {visible.map((p) => { - const i = providers.indexOf(p); - const selected = i === cursor; - const name = pad(p.name, nameWidth); - const url = pad(p.baseUrl, urlWidth); - return ( - onCursor(i)} - onMouseClick={() => { - onCursor(i); - onSelect(p); - }} - > - - {selected ? "\u276f" : " "} - - - - {name} - - {compact ? null : } - {compact ? null : ( - - {url} - - )} - - ); - })} - - {compact ? ( - - - {trunc(`Base URL: ${providers[cursor]?.baseUrl ?? "—"}`, dialogSize.width)} - - - ) : null} - - Add - onDelete(providers[cursor])}>Delete - - - - {compact - ? "\u2191\u2193 move · Enter edit · d del · Esc close" - : "\u2191\u2193 navigate · Enter edit · d delete · Esc back"} + ), + }, + { + key: "state", + header: "State", + value: providerRowState, + render: (row) => ( + + {providerRowState(row)} - - + ), + }, + ]; + const tableLayout = createTableLayout( + columns, + rows, + Math.max(dialogSize.width - 2, 0), ); -} - -// ─── Delete confirm ───────────────────────────────────────────────── -function DeleteConfirm({ - target, - onConfirm, - onCancel, -}: { - target: ModelProviderDto; - onConfirm: () => void; - onCancel: () => void; -}) { - useInput((input, key) => { - if (key.escape) { - onCancel(); - return; - } - const lower = input.toLowerCase(); - if (lower === "y") { - onConfirm(); - } else if (lower === "n") { - onCancel(); - } + useListNavigation({ + items: rows, + cursor, + onCursor: (index) => onSelectedKey(rows[index]?.key ?? null), + onActivate: (row) => onActivate(providerRowEnterAction(row)), + onClose, + onInput: (_input, key, _event, row) => { + if (row?.kind !== "provider") { + return; + } + if (key.leftArrow) { + onToggleProvider(row.provider.code, false); + } else if (key.rightArrow) { + onToggleProvider(row.provider.code, true); + } + }, }); return ( - - - Delete provider "{target.name}"? - - This also deletes all its models. - - [Delete] - [Cancel] + + + + + row.key} + onCursor={(index) => onSelectedKey(rows[index]?.key ?? null)} + onActivate={(row) => onActivate(providerRowEnterAction(row))} + renderItem={(row, { selected }) => ( + + )} + /> ); } diff --git a/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx b/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx new file mode 100644 index 0000000..dee9338 --- /dev/null +++ b/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx @@ -0,0 +1,222 @@ +import { TextAttributes } from "@opentui/core"; +import { testRender } from "@opentui/react/test-utils"; +import { describe, expect, test } from "bun:test"; +import { act } from "react"; + +import { useWindowSize } from "../hooks/useWindowSize"; +import { BottomDialog } from "./BottomDialog"; +import { FormPanel } from "./FormPanel"; +import { KeyValueList } from "./List"; +import { Table, type TableColumn } from "./Table"; + +interface ModelRow { + model: string; + context: string; + state: string; +} + +const MODEL_ROWS: ModelRow[] = [ + { + model: "DeepSeek V4 Flash · deepseek-v4-flash", + context: "1,000,000", + state: "current", + }, + { + model: "DeepSeek V4 Pro · deepseek-v4-pro", + context: "1,000,000", + state: "", + }, +]; + +const MODEL_COLUMNS: Array> = [ + { key: "model", header: "Model", value: (row) => row.model }, + { key: "context", header: "Context", value: (row) => row.context }, + { key: "state", header: "State", value: (row) => row.state }, +]; + +interface ProviderRow { + name: string; + state: string; +} + +const PROVIDER_ROWS: ProviderRow[] = [ + { name: "OpenAI", state: "configured" }, + { name: "OpenAI (Legacy API)", state: "unconfigured" }, + { name: "Anthropic", state: "custom" }, +]; + +const PROVIDER_COLUMNS: Array> = [ + { key: "name", header: "Provider / model", value: (row) => row.name }, + { key: "state", header: "State", value: (row) => row.state }, +]; + +function ResponsiveTable() { + const { columns } = useWindowSize(); + return ( + + ); +} + +function ResponsiveProviderTable() { + const { columns } = useWindowSize(); + return ( +
+ ); +} + +function ResponsiveForm() { + const { columns, rows } = useWindowSize(); + return ( + + undefined} + onClose={() => undefined} + /> + + ); +} + +describe("responsive table and form layout", () => { + test("distributes three columns across the full width and reflows when resized", async () => { + const setup = await testRender(, { width: 180, height: 6 }); + + try { + await act(async () => { + await setup.flush(); + }); + + let lines = setup.captureCharFrame().split("\n"); + const header = lines[0]; + expect(header).toContain("MODEL"); + expect(header.indexOf("CONTEXT")).toBeGreaterThanOrEqual(55); + expect(header.indexOf("STATE")).toBeGreaterThanOrEqual(115); + expect(lines[1].indexOf("DeepSeek")).toBe(0); + expect(lines[1].indexOf("1,000,000")).toBeGreaterThanOrEqual(55); + expect(lines[1].indexOf("current")).toBeGreaterThanOrEqual(115); + + const headerSpans = setup.captureSpans().lines[0]?.spans ?? []; + for (const label of ["MODEL", "CONTEXT", "STATE"]) { + const span = headerSpans.find(({ text }) => text.includes(label)); + expect(span).toBeDefined(); + expect((span?.attributes ?? 0) & TextAttributes.BOLD).toBeTruthy(); + } + + await act(async () => { + setup.resize(54, 6); + await setup.flush(); + }); + await act(async () => { + await setup.flush(); + }); + + lines = setup.captureCharFrame().split("\n"); + expect(lines[1].indexOf("DeepSeek")).toBe(0); + expect(lines[1]).toContain("1,000,000"); + expect(lines[1].length).toBeLessThanOrEqual(54); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("distributes two provider columns across the full width", async () => { + const setup = await testRender(, { + width: 180, + height: 6, + }); + + try { + await act(async () => { + await setup.flush(); + }); + + const lines = setup.captureCharFrame().split("\n"); + expect(lines[0]).toContain("PROVIDER / MODEL"); + expect(lines[0].indexOf("STATE")).toBeGreaterThanOrEqual(85); + expect(lines[1].indexOf("OpenAI")).toBe(0); + expect(lines[1].indexOf("configured")).toBeGreaterThanOrEqual(85); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("shows uppercase bold headers for key-value tables", async () => { + const setup = await testRender( + , + { width: 60, height: 4 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + + expect(setup.captureCharFrame().split("\n")[0]).toContain("KEY"); + expect(setup.captureCharFrame().split("\n")[0]).toContain("VALUE"); + const spans = setup.captureSpans().lines[0]?.spans ?? []; + const headerSpans = spans.filter(({ text }) => /KEY|VALUE/.test(text)); + expect(headerSpans.length).toBeGreaterThan(0); + for (const span of headerSpans) { + expect(span.attributes & TextAttributes.BOLD).toBeTruthy(); + } + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("keeps form labels and values near the left edge on wide and resized screens", async () => { + const setup = await testRender(, { width: 180, height: 16 }); + + try { + await act(async () => { + await setup.flush(); + }); + + let line = setup.captureCharFrame().split("\n") + .find((candidate) => candidate.includes("Provider Code:")) ?? ""; + expect(line.indexOf("Provider Code:")).toBeLessThan(30); + expect(line.indexOf("deepseek")).toBeLessThan(45); + + await act(async () => { + setup.resize(72, 16); + await setup.flush(); + }); + + line = setup.captureCharFrame().split("\n") + .find((candidate) => candidate.includes("Provider Code:")) ?? ""; + expect(line.indexOf("Provider Code:")).toBeLessThan(30); + expect(line.indexOf("deepseek")).toBeLessThan(45); + } finally { + act(() => setup.renderer.destroy()); + } + }); +}); diff --git a/packages/agenty-cli/src/components/SelectOverlay.tsx b/packages/agenty-cli/src/components/SelectOverlay.tsx index 5fc04bc..7f182b1 100644 --- a/packages/agenty-cli/src/components/SelectOverlay.tsx +++ b/packages/agenty-cli/src/components/SelectOverlay.tsx @@ -1,7 +1,8 @@ import { useEffect, useRef, useState } from "react"; -import { useInput } from "../hooks/useInput"; -import { Box, Select, Spinner, Text } from "./ui"; +import { List, useListNavigation } from "./List"; +import { Panel } from "./Panel"; +import { Box, Spinner, Text } from "./ui"; export interface SelectEntry { label: string; @@ -29,6 +30,7 @@ export function SelectOverlay({ }: SelectOverlayProps) { const [entries, setEntries] = useState[] | null>(null); const [error, setError] = useState(null); + const [cursor, setCursor] = useState(0); const loadRef = useRef(load); loadRef.current = load; @@ -51,15 +53,15 @@ export function SelectOverlay({ }; }, []); - useInput((_input, key) => { - if (key.escape) { - onClose(); - } + useListNavigation({ + items: entries ?? [], + cursor, + onCursor: setCursor, + onActivate: (entry) => onSelect(entry.data), + onClose, + active: entries !== null && entries.length > 0, }); - const options = - entries?.map((e, i) => ({ label: e.label, value: String(i) })) ?? []; - return ( ({ paddingX={dialog ? 0 : 2} paddingY={dialog ? 0 : 1} > - - - {title} - - · Esc to cancel - - {error ? ( - Failed: {error} - ) : entries === null ? ( - - ) : entries.length === 0 ? ( - {emptyHint ?? "No items"} - ) : ( - ({ - name: option.label, - description: "", - value: option.value, - }))} - showDescription={false} - showScrollIndicator={options.length > visibleOptionCount} - selectedBackgroundColor="#24383f" - selectedTextColor="#c7f5ff" - onSelect={(_index, option) => { - if (option) { - onChange(String(option.value)); - } - }} - /> - ); -} diff --git a/packages/agenty-cli/src/components/ui/Text.tsx b/packages/agenty-cli/src/components/ui/Text.tsx index 4f1b1f6..11dcd7c 100644 --- a/packages/agenty-cli/src/components/ui/Text.tsx +++ b/packages/agenty-cli/src/components/ui/Text.tsx @@ -2,6 +2,7 @@ import { createTextAttributes, type MouseEvent } from "@opentui/core"; import { createContext, type ReactNode, useContext, useRef } from "react"; const TextNestingContext = createContext(false); +export const TextBackgroundContext = createContext(undefined); export type TextProps = { children?: ReactNode; @@ -38,6 +39,8 @@ export function Text({ ...layout }: TextProps) { const nested = useContext(TextNestingContext); + const inheritedBackground = useContext(TextBackgroundContext); + const resolvedBackground = backgroundColor ?? inheritedBackground; const clickStart = useRef<{ x: number; y: number } | null>(null); const attributes = createTextAttributes({ bold, @@ -47,7 +50,7 @@ export function Text({ }); if (nested) { return ( - + {children} ); @@ -56,7 +59,7 @@ export function Text({ void; + onKeyDown?: (event: KeyEvent) => void; }; export const TextInput = forwardRef( @@ -19,6 +21,8 @@ export const TextInput = forwardRef( placeholder, focus = true, keepFocus = false, + onMouseDown, + onKeyDown, }, ref, ) => { @@ -86,6 +90,11 @@ export const TextInput = forwardRef( focusedTextColor="#ffffff" cursorColor="#00e5ff" onInput={onChange} + onMouseDown={() => { + inputRef.current?.focus(); + onMouseDown?.(); + }} + onKeyDown={onKeyDown} onSubmit={handleSubmit} /> ); diff --git a/packages/agenty-cli/src/components/ui/index.ts b/packages/agenty-cli/src/components/ui/index.ts index a0cc8e6..c17d51c 100644 --- a/packages/agenty-cli/src/components/ui/index.ts +++ b/packages/agenty-cli/src/components/ui/index.ts @@ -1,6 +1,13 @@ export { Box, type BoxProps } from "./Box"; export { GradientText } from "./GradientText"; -export { Select, type SelectProps } from "./Select"; +export { + ActionBar, + type ActionBarItem, + type ActionBarProps, + HOVER_BACKGROUND, + Pressable, + type PressableProps, +} from "./Interactive"; export { Spinner } from "./Spinner"; export { Text, type TextProps } from "./Text"; export { TextInput, type TextInputProps } from "./TextInput"; diff --git a/packages/agenty-cli/src/components/wizardRows.test.ts b/packages/agenty-cli/src/components/wizardRows.test.ts new file mode 100644 index 0000000..39c21b1 --- /dev/null +++ b/packages/agenty-cli/src/components/wizardRows.test.ts @@ -0,0 +1,119 @@ +import { describe, expect, test } from "bun:test"; + +import type { ModelProviderDto } from "../api/types"; +import { + createBuiltinDraft, + createCustomDraft, + createModelDraft, +} from "../consts/providerPresets"; +import { + buildWizardModelRows, + buildWizardProviderRows, + wizardModelEnterAction, +} from "./wizardRows"; + +function provider(code: string, builtin: boolean): ModelProviderDto { + return { + code, + name: code === "openai" ? "OpenAI" : "Custom", + type: builtin ? "openai" : "openai_completions", + baseUrl: "https://example.invalid/v1", + apiKey: builtin ? "" : "test-key", + builtin, + official: builtin, + models: [], + createdAt: "", + updatedAt: "", + }; +} + +describe("wizard hierarchical rows", () => { + test("keeps Add more provider as the last selectable provider row", () => { + const builtinProvider = provider("openai", true); + const customProvider = provider("custom", false); + const rows = buildWizardProviderRows( + [createCustomDraft("provider:custom", customProvider)], + [builtinProvider], + ); + + expect(rows.map((row) => row.kind)).toEqual(["builtin", "custom", "add"]); + expect(rows.at(-1)).toEqual({ + kind: "add", + key: "add-provider", + label: "Add more provider...", + }); + }); + + test("places each provider before all of its model rows", () => { + const builtin = createBuiltinDraft(provider("openai", true)); + const custom = createCustomDraft("provider:custom", provider("custom", false)); + const first = createModelDraft(custom, "custom:first", { + code: "first", + name: "First", + contextWindow: 128_000, + maxOutputTokens: 8_192, + multiModal: false, + light: false, + isDefault: true, + }); + const second = createModelDraft(custom, "custom:second", { + ...first, + code: "second", + name: "Second", + isDefault: false, + }); + + const rows = buildWizardModelRows([builtin, custom], [first, second]); + + expect(rows.map((row) => row.kind)).toEqual([ + "provider", + "provider", + "model", + "model", + "add-model", + ]); + expect(rows.map((row) => row.key)).toEqual([ + builtin.id, + custom.id, + first.id, + second.id, + `${custom.id}:add-model`, + ]); + expect(rows.at(-1)).toMatchObject({ + kind: "add-model", + provider: custom, + label: "Add more models...", + }); + }); + + test("routes Enter to select built-ins, edit custom models, and add model rows", () => { + const builtin = createBuiltinDraft(provider("openai", true)); + const custom = createCustomDraft("provider:custom", provider("custom", false)); + const builtinModel = createModelDraft(builtin, "openai:model", { + code: "builtin", + name: "Built-in", + contextWindow: 128_000, + maxOutputTokens: 8_192, + multiModal: false, + light: false, + isDefault: true, + }); + const customModel = createModelDraft(custom, "custom:model", { + ...builtinModel, + code: "custom", + name: "Custom", + isDefault: false, + }); + const rows = buildWizardModelRows( + [builtin, custom], + [builtinModel, customModel], + ); + + expect(wizardModelEnterAction(rows.find((row) => row.key === builtinModel.id))) + .toEqual({ kind: "select", model: builtinModel }); + expect(wizardModelEnterAction(rows.find((row) => row.key === customModel.id))) + .toEqual({ kind: "edit", model: customModel }); + expect(wizardModelEnterAction(rows.find((row) => row.kind === "add-model"))) + .toEqual({ kind: "add", provider: custom }); + }); +}); diff --git a/packages/agenty-cli/src/components/wizardRows.ts b/packages/agenty-cli/src/components/wizardRows.ts new file mode 100644 index 0000000..656683b --- /dev/null +++ b/packages/agenty-cli/src/components/wizardRows.ts @@ -0,0 +1,127 @@ +import type { APIType, ModelProviderDto } from "../api/types"; +import { + compatibleProviderTypes, + type ModelDraft, + type ProviderDraft, +} from "../consts/providerPresets"; + +export type WizardProviderRow = + | { + kind: "builtin" | "custom"; + key: string; + code: string; + draft?: ProviderDraft; + label: string; + description: string; + } + | { + kind: "add"; + key: "add-provider"; + label: "Add more provider..."; + }; + +export type WizardModelRow = + | { + kind: "provider"; + key: string; + provider: ProviderDraft; + modelCount: number; + } + | { + kind: "model"; + key: string; + provider: ProviderDraft; + model: ModelDraft; + } + | { + kind: "add-model"; + key: string; + provider: ProviderDraft; + label: "Add more models..."; + }; + +export type WizardModelEnterAction = + | { kind: "select"; model: ModelDraft } + | { kind: "edit"; model: ModelDraft } + | { kind: "add"; provider: ProviderDraft } + | { kind: "none" }; + +function providerTypeLabel(type: APIType): string { + return compatibleProviderTypes.find((option) => option.value === type)?.label ?? type; +} + +export function buildWizardProviderRows( + drafts: ProviderDraft[], + builtinProviders: ModelProviderDto[], +): WizardProviderRow[] { + const rows: WizardProviderRow[] = builtinProviders.map((provider) => ({ + kind: "builtin", + key: `builtin:${provider.code}`, + code: provider.code, + draft: drafts.find((draft) => draft.code === provider.code), + label: provider.name, + description: providerTypeLabel(provider.type), + })); + rows.push( + ...drafts + .filter((draft) => draft.source === "custom") + .map((draft) => ({ + kind: "custom" as const, + key: draft.id, + code: draft.code, + draft, + label: draft.name || draft.code || "Custom provider", + description: providerTypeLabel(draft.type), + })), + { + kind: "add", + key: "add-provider", + label: "Add more provider...", + }, + ); + return rows; +} + +export function buildWizardModelRows( + providers: ProviderDraft[], + models: ModelDraft[], +): WizardModelRow[] { + const rows: WizardModelRow[] = []; + for (const provider of providers) { + const providerModels = models.filter((model) => model.providerId === provider.id); + rows.push({ + kind: "provider", + key: provider.id, + provider, + modelCount: providerModels.length, + }); + rows.push(...providerModels.map((model) => ({ + kind: "model" as const, + key: model.id, + provider, + model, + }))); + if (!provider.builtin) { + rows.push({ + kind: "add-model", + key: `${provider.id}:add-model`, + provider, + label: "Add more models...", + }); + } + } + return rows; +} + +export function wizardModelEnterAction(row: WizardModelRow | undefined): WizardModelEnterAction { + if (!row || row.kind === "provider") { + return { kind: "none" }; + } + if (row.kind === "add-model") { + return { kind: "add", provider: row.provider }; + } + if (row.model.isBuiltin) { + return { kind: "select", model: row.model }; + } + return { kind: "edit", model: row.model }; +} diff --git a/packages/agenty-cli/src/components/wizardSetup.test.ts b/packages/agenty-cli/src/components/wizardSetup.test.ts index 4b6cf59..f381c19 100644 --- a/packages/agenty-cli/src/components/wizardSetup.test.ts +++ b/packages/agenty-cli/src/components/wizardSetup.test.ts @@ -2,29 +2,70 @@ import { describe, expect, test } from "bun:test"; import type { AgentDto, ModelProviderDto } from "../api/types"; import { - createPresetDraft, + createBuiltinDraft, + createCustomDraft, + createModelDraft, type ModelDraft, - modelDraftForProvider, + modelDraftsForProvider, type ProviderDraft, - providerPresets, validateModelDraft, validateProviderDraft, } from "../consts/providerPresets"; import { persistWizardSetup, - selectedModelCode, + selectedModelId, + validateWizardDrafts, type WizardSetupClient, } from "./wizardSetup"; function createDraft(): ProviderDraft { return { - ...createPresetDraft(providerPresets[0]), + ...createCustomDraft("custom:0", { + ...createProviderResource(), + code: "custom", + name: "Custom", + type: "openai_completions", + builtin: false, + official: false, + }), apiKey: "test-key", }; } function createModel(draft: ProviderDraft): ModelDraft { - return modelDraftForProvider(draft, providerPresets[0]); + const provider = { + ...createProviderResource(), + code: draft.code, + name: draft.name, + type: draft.type, + builtin: false, + official: false, + }; + return createModelDraft(draft, `${draft.id}:model`, provider.models[0]); +} + +function createProviderResource(): ModelProviderDto { + return { + code: "openai", + name: "OpenAI", + type: "openai", + baseUrl: "https://api.openai.com/v1", + apiKey: "", + builtin: true, + official: true, + models: [{ + code: "gpt-5-mini", + name: "GPT-5 mini", + contextWindow: 400_000, + maxOutputTokens: 128_000, + multiModal: true, + light: true, + reasoningEfforts: ["low", "medium", "high", "xhigh"], + isDefault: true, + }], + createdAt: "", + updatedAt: "", + }; } function createAgent(code: string, isDefault: boolean): AgentDto { @@ -63,10 +104,18 @@ function createProvider(draft: ProviderDraft, model: ModelDraft = createModel(dr function fakeClient( providers: ModelProviderDto[] = [], agents: AgentDto[] = [], -): WizardSetupClient & { calls: string[] } { +): WizardSetupClient & { + calls: string[]; + createdModels: Array<{ modelCode: string; isDefault?: boolean }>; + deletedModels: Array<{ providerCode: string; modelCode: string }>; +} { const calls: string[] = []; + const createdModels: Array<{ modelCode: string; isDefault?: boolean }> = []; + const deletedModels: Array<{ providerCode: string; modelCode: string }> = []; return { calls, + createdModels, + deletedModels, listProviders: async () => { calls.push("provider.list"); return providers; @@ -83,10 +132,15 @@ function fakeClient( calls.push("provider.update"); return providers[0] ?? createProvider(createDraft()); }, - createModel: async () => { + createModel: async (input) => { calls.push("provider.addModel"); + createdModels.push(input); return undefined; }, + deleteModel: async (providerCode, modelCode) => { + calls.push("provider.removeModel"); + deletedModels.push({ providerCode, modelCode }); + }, createAgent: async () => { calls.push("agent.create"); return createAgent("default", true); @@ -103,22 +157,15 @@ function fakeClient( } describe("first-run provider setup", () => { - test("exposes the three supported preset providers", () => { - expect(providerPresets.map((preset) => preset.key)).toEqual([ - "openai", - "anthropic", - "google", - ]); - expect(providerPresets.map((preset) => preset.type)).toEqual([ - "openai", - "anthropic", - "gemini", - ]); - expect(providerPresets.every((preset) => preset.model.code.length > 0)).toBe(true); + test("restores built-in provider metadata from the core projection", () => { + const provider = createProviderResource(); + const draft = createBuiltinDraft(provider); + expect(draft.source).toBe("builtin"); + expect(draft.code).toBe("openai"); + expect(modelDraftsForProvider(draft, provider).map((model) => model.code)).toEqual(["gpt-5-mini"]); }); test("restores an existing provider and its preferred model", () => { - const preset = providerPresets[0]; const draft = createDraft(); const existingModel = createModel(draft); const existing: ModelProviderDto = { @@ -132,11 +179,11 @@ describe("first-run provider setup", () => { isDefault: true, }], }; - const restoredProvider = createPresetDraft(preset, existing); - const restoredModel = modelDraftForProvider(restoredProvider, preset, existing); + const restoredProvider = createBuiltinDraft(createProviderResource(), existing); + const restoredModel = modelDraftsForProvider(restoredProvider, existing)[0]; - expect(restoredProvider.name).toBe("OpenAI gateway"); - expect(restoredProvider.baseUrl).toBe("https://gateway.example/v1"); + expect(restoredProvider.name).toBe("OpenAI"); + expect(restoredProvider.baseUrl).toBe("https://api.openai.com/v1"); expect(restoredModel.code).toBe("gateway-model"); }); @@ -155,7 +202,7 @@ describe("first-run provider setup", () => { const model = createModel(draft); const client = fakeClient(); - await persistWizardSetup(client, [draft], [model], selectedModelCode(model)); + await persistWizardSetup(client, [draft], [model], selectedModelId(model)); expect(client.calls).toEqual([ "provider.list", @@ -172,15 +219,72 @@ describe("first-run provider setup", () => { const model = createModel(draft); const client = fakeClient([createProvider(draft, model)], [createAgent("default", true)]); - await persistWizardSetup(client, [draft], [model], selectedModelCode(model)); + await persistWizardSetup(client, [draft], [model], selectedModelId(model)); + + expect(client.calls).toEqual([ + "provider.list", + "agent.list", + "provider.update", + "provider.addModel", + "agent.update", + "initialize.complete", + ]); + }); + + test("persists multiple custom models and removes models deleted in the wizard", async () => { + const draft = createDraft(); + const first = { ...createModel(draft), id: `${draft.id}:first`, code: "first", name: "First" }; + const second = { ...createModel(draft), id: `${draft.id}:second`, code: "second", name: "Second" }; + const existing = createProvider(draft, first); + existing.models.push({ ...existing.models[0], code: "removed", name: "Removed" }); + const client = fakeClient([existing], [createAgent("default", true)]); + + await persistWizardSetup(client, [draft], [first, second], selectedModelId(second)); expect(client.calls).toEqual([ "provider.list", "agent.list", "provider.update", + "provider.removeModel", + "provider.addModel", "provider.addModel", "agent.update", "initialize.complete", ]); + expect(client.deletedModels).toEqual([{ providerCode: "custom", modelCode: "removed" }]); + expect(client.createdModels.map(({ modelCode, isDefault }) => ({ modelCode, isDefault }))).toEqual([ + { modelCode: "first", isDefault: false }, + { modelCode: "second", isDefault: true }, + ]); + }); + + test("allows multiple models per provider and rejects duplicate model codes", () => { + const draft = createDraft(); + const first = { ...createModel(draft), id: `${draft.id}:first` }; + const second = { ...createModel(draft), id: `${draft.id}:second`, code: "second" }; + + expect(validateWizardDrafts([draft], [first, second], selectedModelId(second))).toBeNull(); + expect(validateWizardDrafts( + [draft], + [first, { ...second, code: first.code }], + selectedModelId(second), + )).toContain("already configured"); + }); + + test("updates only the API key when a built-in provider is selected", async () => { + const provider = { ...createProviderResource(), apiKey: "test-key" }; + const draft = createBuiltinDraft(provider); + const models = modelDraftsForProvider(draft, provider); + const client = fakeClient([provider]); + + await persistWizardSetup(client, [draft], models, selectedModelId(models[0])); + + expect(client.calls).toEqual([ + "provider.list", + "agent.list", + "provider.update", + "agent.create", + "initialize.complete", + ]); }); }); diff --git a/packages/agenty-cli/src/components/wizardSetup.ts b/packages/agenty-cli/src/components/wizardSetup.ts index a2b3474..7529caf 100644 --- a/packages/agenty-cli/src/components/wizardSetup.ts +++ b/packages/agenty-cli/src/components/wizardSetup.ts @@ -1,5 +1,6 @@ import type { AgentDto, + CreateModelDto, CreateModelProviderDto, ModelProviderDto, ReasoningEffort, @@ -22,14 +23,8 @@ export interface WizardSetupClient { listAgents(): Promise; createProvider(input: CreateModelProviderDto): Promise; updateProvider(code: string, input: UpdateModelProviderDto): Promise; - createModel(input: { - providerCode: string; - modelCode: string; - name: string; - contextWindow?: number; - reasoningEffortMapping?: Record; - isDefault?: boolean; - }): Promise; + createModel(input: CreateModelDto): Promise; + deleteModel(providerCode: string, modelCode: string): Promise; createAgent(input: { code: string; name: string; @@ -47,8 +42,8 @@ export interface WizardSetupClient { }): Promise<{ initialized: boolean }>; } -export function selectedModelCode(model: ModelDraft): string { - return `${model.providerId}:${model.code.trim()}`; +export function selectedModelId(model: ModelDraft): string { + return model.id; } export function validateWizardDrafts( @@ -59,12 +54,12 @@ export function validateWizardDrafts( if (drafts.length === 0) { return "Configure at least one provider to continue."; } - if (models.length !== drafts.length) { - return "Configure one model for each provider to continue."; + if (models.length === 0) { + return "Configure or select at least one model to continue."; } const providerCodes = new Set(); - const providerIds = new Set(drafts.map((draft) => draft.id)); + const providersById = new Map(drafts.map((draft) => [draft.id, draft])); for (const draft of drafts) { const providerError = validateProviderDraft(draft); if (providerError) { @@ -77,23 +72,30 @@ export function validateWizardDrafts( providerCodes.add(code); } - const modelCodes = new Set(); + const modelCodesByProvider = new Map>(); for (const model of models) { - if (!providerIds.has(model.providerId)) { + const provider = providersById.get(model.providerId); + if (!provider) { return "A model is attached to an unknown provider."; } + if (model.providerCode.trim() !== provider.code.trim()) { + return `Model provider does not match ${provider.name.trim() || provider.code.trim()}.`; + } const modelError = validateModelDraft(model); if (modelError) { return modelError; } - const modelCode = selectedModelCode(model); + + const modelCodes = modelCodesByProvider.get(model.providerId) ?? new Set(); + const modelCode = model.code.trim(); if (modelCodes.has(modelCode)) { return `Model Code already configured: ${model.code.trim()}`; } modelCodes.add(modelCode); + modelCodesByProvider.set(model.providerId, modelCodes); } - if (!modelCodes.has(selectedId)) { + if (!models.some((model) => selectedModelId(model) === selectedId)) { return "Select a default model to continue."; } return null; @@ -112,55 +114,68 @@ export async function persistWizardSetup( const existingProviders = (await client.listProviders()) ?? []; const existingAgents = (await client.listAgents()) ?? []; - const modelsByProvider = new Map(models.map((model) => [model.providerId, model])); - let selectedModel: ModelDraft | undefined; + const modelsByProvider = new Map(); + for (const model of models) { + const providerModels = modelsByProvider.get(model.providerId) ?? []; + providerModels.push(model); + modelsByProvider.set(model.providerId, providerModels); + } + const selectedModel = models.find((model) => selectedModelId(model) === selectedId); + if (!selectedModel) { + throw new Error("Select a default model to continue."); + } for (const draft of drafts) { - const model = modelsByProvider.get(draft.id); - if (!model) { - throw new Error(`No model configured for provider ${draft.name || draft.code}.`); - } - const providerCode = draft.code.trim(); - const providerInput: CreateModelProviderDto = { - code: providerCode, - name: draft.name.trim(), - type: draft.type, - baseUrl: draft.baseUrl.trim(), - apiKey: draft.apiKey.trim(), - }; const existing = existingProviders.find((provider) => provider.code === providerCode); - if (existing) { + if (draft.builtin) { + await client.updateProvider(providerCode, { apiKey: draft.apiKey.trim() }); + } else if (existing) { await client.updateProvider(providerCode, { - name: providerInput.name, - type: providerInput.type, - baseUrl: providerInput.baseUrl, - apiKey: providerInput.apiKey, + name: draft.name.trim(), + type: draft.type, + baseUrl: draft.baseUrl.trim(), + apiKey: draft.apiKey.trim(), + freeFormTool: draft.type === "openai" && draft.freeFormTool, }); } else { + const providerInput: CreateModelProviderDto = { + code: providerCode, + name: draft.name.trim(), + type: draft.type, + baseUrl: draft.baseUrl.trim(), + apiKey: draft.apiKey.trim(), + freeFormTool: draft.type === "openai" && draft.freeFormTool, + }; await client.createProvider(providerInput); } - const modelCode = model.code.trim(); - const isSelected = selectedModelCode(model) === selectedId; - if (isSelected) { - selectedModel = model; - } - await client.createModel({ - providerCode, - modelCode, - name: model.name.trim(), - contextWindow: model.contextWindow, - reasoningEffortMapping: model.reasoningEffortMapping, - isDefault: isSelected, - }); - } + if (!draft.builtin) { + const providerModels = modelsByProvider.get(draft.id) ?? []; + const desiredModelCodes = new Set(providerModels.map((model) => model.code.trim())); + for (const model of existing?.models ?? []) { + if (!desiredModelCodes.has(model.code)) { + await client.deleteModel(providerCode, model.code); + } + } - if (!selectedModel) { - throw new Error("Select a default model to continue."); + for (const model of providerModels) { + await client.createModel({ + providerCode, + modelCode: model.code.trim(), + name: model.name.trim(), + contextWindow: model.contextWindow, + maxOutputTokens: model.maxOutputTokens, + multiModal: model.multiModal, + light: model.light, + reasoning: model.reasoningEfforts.length > 0, + isDefault: selectedModelId(model) === selectedId, + }); + } + } } - const selectedProvider = drafts.find((draft) => draft.id === selectedModel?.providerId); + const selectedProvider = drafts.find((draft) => draft.id === selectedModel.providerId); if (!selectedProvider) { throw new Error("Selected model provider is missing."); } diff --git a/packages/agenty-cli/src/config.ts b/packages/agenty-cli/src/config.ts index 917c294..7231b64 100644 --- a/packages/agenty-cli/src/config.ts +++ b/packages/agenty-cli/src/config.ts @@ -2,7 +2,7 @@ export type ThinkingFlag = "off" | "on" | string; export interface CliOptions { agentRef?: string; - modelRef?: string; + modelInput?: string; thinking: ThinkingFlag; dataDir?: string; newSession: boolean; @@ -33,7 +33,7 @@ export function loadOptions(): CliOptions { const flags = parseArgs(process.argv.slice(2)); return { agentRef: typeof flags.agent === "string" ? flags.agent : undefined, - modelRef: typeof flags.model === "string" ? flags.model : undefined, + modelInput: typeof flags.model === "string" ? flags.model : undefined, thinking: typeof flags.thinking === "string" ? flags.thinking : "off", dataDir: typeof flags["data-dir"] === "string" ? flags["data-dir"] : undefined, newSession: flags["new-session"] === true, diff --git a/packages/agenty-cli/src/consts/providerPresets.ts b/packages/agenty-cli/src/consts/providerPresets.ts index f12c376..ac376c7 100644 --- a/packages/agenty-cli/src/consts/providerPresets.ts +++ b/packages/agenty-cli/src/consts/providerPresets.ts @@ -1,89 +1,5 @@ import type { APIType, CoreModelDto, ModelProviderDto, ReasoningEffort } from "../api/types"; - -export interface ModelPreset { - code: string; - name: string; - contextWindow: number; - reasoningEffortMapping?: Record; -} - -export interface ProviderPreset { - key: string; - label: string; - description: string; - code: string; - name: string; - type: APIType; - baseUrl: string; - model: ModelPreset; -} - -export const providerPresets: readonly ProviderPreset[] = [ - { - key: "openai", - label: "OpenAI", - description: "Responses API", - code: "openai", - name: "OpenAI", - type: "openai", - baseUrl: "https://api.openai.com/v1", - model: { - code: "gpt-5-mini", - name: "GPT-5 mini", - contextWindow: 128_000, - reasoningEffortMapping: { - low: "low", - medium: "medium", - high: "high", - xhigh: "xhigh", - }, - }, - }, - { - key: "anthropic", - label: "Anthropic", - description: "Messages API", - code: "anthropic", - name: "Anthropic", - type: "anthropic", - baseUrl: "https://api.anthropic.com", - model: { - code: "claude-haiku-4-5", - name: "Claude Haiku 4.5", - contextWindow: 200_000, - reasoningEffortMapping: { - low: "low", - medium: "medium", - high: "high", - max: "max", - }, - }, - }, - { - key: "google", - label: "Google", - description: "Gemini API", - code: "google", - name: "Google", - type: "gemini", - baseUrl: "https://generativelanguage.googleapis.com/v1beta", - model: { - code: "gemini-2.5-flash", - name: "Gemini 2.5 Flash", - contextWindow: 128_000, - reasoningEffortMapping: { - low: "low", - medium: "medium", - high: "high", - }, - }, - }, -]; - -export function isBuiltinProvider(provider: Pick | string): boolean { - const code = typeof provider === "string" ? provider : provider.code; - return providerPresets.some((preset) => preset.code === code); -} +import { STANDARD_REASONING_EFFORTS } from "../api/types"; export const compatibleProviderTypes: readonly { label: string; value: APIType }[] = [ { label: "OpenAI Responses API", value: "openai" }, @@ -92,15 +8,21 @@ export const compatibleProviderTypes: readonly { label: string; value: APIType } { label: "Google Gemini API", value: "gemini" }, ]; +export function isBuiltinProvider(provider: Pick | string): boolean { + return typeof provider === "string" ? false : provider.builtin === true; +} + export interface ProviderDraft { id: string; - source: "preset" | "custom"; - presetKey?: string; + source: "builtin" | "custom"; + originalCode?: string; code: string; name: string; type: APIType; baseUrl: string; apiKey: string; + freeFormTool: boolean; + builtin: boolean; } export interface ModelDraft { @@ -108,77 +30,87 @@ export interface ModelDraft { providerId: string; providerCode: string; providerName: string; + originalCode?: string; code: string; name: string; contextWindow: number; - reasoningEffortMapping?: Record; + maxOutputTokens: number; + multiModal: boolean; + light: boolean; + reasoningEfforts: ReasoningEffort[]; + isDefault: boolean; + isBuiltin: boolean; } -function preferredModel(existing?: ModelProviderDto): CoreModelDto | undefined { - const models = existing?.models ?? []; - return models.find((model) => model.isDefault) ?? models[0]; -} - -export function createPresetDraft( - preset: ProviderPreset, - existing?: ModelProviderDto, -): ProviderDraft { +export function createBuiltinDraft(provider: ModelProviderDto, existing?: ModelProviderDto): ProviderDraft { + const source = existing ?? provider; return { - id: `preset:${preset.key}`, - source: "preset", - presetKey: preset.key, - code: existing?.code ?? preset.code, - name: existing?.name ?? preset.name, - type: existing?.type ?? preset.type, - baseUrl: existing?.baseUrl ?? preset.baseUrl, - apiKey: existing?.apiKey ?? "", + id: `builtin:${provider.code}`, + source: "builtin", + originalCode: provider.code, + code: provider.code, + name: provider.name, + type: provider.type, + baseUrl: provider.baseUrl, + apiKey: source.apiKey ?? "", + freeFormTool: provider.freeFormTool === true, + builtin: true, }; } -export function createCustomDraft( - id: string, - existing?: ModelProviderDto, -): ProviderDraft { +export function createCustomDraft(id: string, existing?: ModelProviderDto): ProviderDraft { return { id, source: "custom", + originalCode: existing?.code, code: existing?.code ?? "", name: existing?.name ?? "", type: existing?.type ?? "openai_completions", baseUrl: existing?.baseUrl ?? "", apiKey: existing?.apiKey ?? "", + freeFormTool: existing?.freeFormTool === true, + builtin: false, }; } -export function draftForProvider( - provider: ModelProviderDto, - preset?: ProviderPreset, -): ProviderDraft { - if (preset) { - return createPresetDraft(preset, provider); - } - return createCustomDraft(`provider:${provider.code}`, provider); +export function draftForProvider(provider: ModelProviderDto): ProviderDraft { + return provider.builtin === true + ? createBuiltinDraft(provider) + : createCustomDraft(`provider:${provider.code}`, provider); } -export function modelDraftForProvider( +export function createModelDraft( provider: ProviderDraft, - preset?: ProviderPreset, - existing?: ModelProviderDto, + id: string, + existing?: CoreModelDto, ): ModelDraft { - const model = preferredModel(existing); - const fallback = preset?.model; return { - id: `${provider.id}:model`, + id, providerId: provider.id, providerCode: provider.code, providerName: provider.name, - code: model?.code ?? fallback?.code ?? "", - name: model?.name ?? fallback?.name ?? "", - contextWindow: model?.contextWindow ?? fallback?.contextWindow ?? 128_000, - reasoningEffortMapping: model?.reasoningEffortMapping ?? fallback?.reasoningEffortMapping, + originalCode: existing?.code, + code: existing?.code ?? "", + name: existing?.name ?? "", + contextWindow: existing?.contextWindow ?? 128_000, + maxOutputTokens: existing?.maxOutputTokens ?? 8_192, + multiModal: existing?.multiModal ?? false, + light: existing?.light ?? false, + reasoningEfforts: existing?.reasoningEfforts ?? [...STANDARD_REASONING_EFFORTS], + isDefault: existing?.isDefault ?? false, + isBuiltin: provider.builtin === true, }; } +export function modelDraftsForProvider( + provider: ProviderDraft, + existing?: ModelProviderDto, +): ModelDraft[] { + return (existing?.models ?? []).map((model) => + createModelDraft(provider, `${provider.id}:model:${model.code}`, model), + ); +} + export function validateProviderDraft(draft: ProviderDraft): string | null { if (!draft.code.trim()) { return "Provider code is required."; @@ -205,5 +137,8 @@ export function validateModelDraft(draft: ModelDraft): string | null { if (!Number.isSafeInteger(draft.contextWindow) || draft.contextWindow <= 0) { return `Context window for ${draft.name.trim()} must be a positive integer.`; } + if (!Number.isSafeInteger(draft.maxOutputTokens) || draft.maxOutputTokens <= 0) { + return `Max output tokens for ${draft.name.trim()} must be a positive integer.`; + } return null; } diff --git a/packages/agenty-cli/src/state/store.test.ts b/packages/agenty-cli/src/state/store.test.ts index 0f5665f..70ebc79 100644 --- a/packages/agenty-cli/src/state/store.test.ts +++ b/packages/agenty-cli/src/state/store.test.ts @@ -253,7 +253,7 @@ describe("chat tool event projection", () => { async getSession() { return persisted; }, - async resolveModel() { + async getModel() { return { code: "model", providerCode: "provider", @@ -339,7 +339,7 @@ describe("chat tool event projection", () => { async getSession() { return persisted; }, - async resolveModel() { + async getModel() { return { code: "model", providerCode: "provider", diff --git a/packages/agenty-cli/src/state/store.ts b/packages/agenty-cli/src/state/store.ts index beb23d7..6b31bd7 100644 --- a/packages/agenty-cli/src/state/store.ts +++ b/packages/agenty-cli/src/state/store.ts @@ -461,7 +461,7 @@ export const useAppStore = create((set, get) => { const parsed = parseThinking(options.thinking); const prepared = await client.prepareSession({ agentRef: options.agentRef, - modelRef: options.modelRef, + modelInput: options.modelInput, newSession: options.newSession, reasoningEffort: reasoningEffort(parsed.thinking, parsed.thinkingLevel), }); @@ -681,7 +681,7 @@ export const useAppStore = create((set, get) => { try { const full = await client.getSession(session.id); const model = full.currentModel - ? await client.resolveModel(`${full.currentModel.providerCode}/${full.currentModel.modelCode}`) + ? await client.getModel(full.currentModel) : get().model; set({ session: full, model, history: buildHistory(full), current: null, tokenConsumed: actualContextSize(full), overlay: null }); } catch (error) { @@ -696,7 +696,7 @@ export const useAppStore = create((set, get) => { } try { const model = agent.defaultModel - ? await client.resolveModel(`${agent.defaultModel.providerCode}/${agent.defaultModel.modelCode}`) + ? await client.getModel(agent.defaultModel) : await client.getDefaultModel(); const session = await client.getLastSessionByAgent(agent.code) ?? await client.createSession(agent.code, model); set({ agent, model, session, history: buildHistory(session), current: null, tokenConsumed: actualContextSize(session), overlay: null }); From facbcece5859f87c1d560400bff62c927e3f9808 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Mon, 24 Aug 2026 18:31:28 +0800 Subject: [PATCH 03/12] feat: replace file edits with transactional apply_patch Signed-off-by: masteryyh --- .../agenty-cli/src/components/MessageItem.tsx | 23 +- .../agenty-cli/src/components/MessageList.tsx | 12 +- .../src/components/toolDisplay.test.ts | 16 + .../agenty-cli/src/components/toolDisplay.ts | 73 +- .../pkg/agentloop/builtin/apply_patch.go | 292 +--- .../pkg/agentloop/builtin/apply_patch_test.go | 181 +-- .../agenty-core/pkg/agentloop/builtin/file.go | 266 ---- .../pkg/agentloop/builtin/file_test.go | 237 +-- .../pkg/agentloop/builtin/register.go | 3 - .../pkg/agentloop/builtin/register_test.go | 3 - .../agenty-core/pkg/agentloop/compaction.go | 2 +- packages/agenty-core/pkg/agentloop/engine.go | 46 +- .../agenty-core/pkg/agentloop/engine_test.go | 70 +- .../agenty-core/pkg/domain/agent/agent.go | 20 +- .../pkg/domain/agent/agent_test.go | 30 +- packages/agenty-core/pkg/utils/apply_diff.go | 393 ----- .../agenty-core/pkg/utils/apply_diff_test.go | 258 ---- .../agenty-core/test/e2e/execution_test.go | 4 +- .../agenty-core/test/e2e/test_helpers_test.go | 11 +- packages/patch-applier/.gitignore | 2 + packages/patch-applier/Cargo.lock | 114 ++ packages/patch-applier/Cargo.toml | 19 + packages/patch-applier/README.md | 32 + packages/patch-applier/package.json | 11 + packages/patch-applier/src/lib.rs | 1297 +++++++++++++++++ packages/patch-applier/src/main.rs | 61 + packages/patch-applier/tests/cli.rs | 83 ++ 27 files changed, 1911 insertions(+), 1648 deletions(-) delete mode 100644 packages/agenty-core/pkg/utils/apply_diff.go delete mode 100644 packages/agenty-core/pkg/utils/apply_diff_test.go create mode 100644 packages/patch-applier/.gitignore create mode 100644 packages/patch-applier/Cargo.lock create mode 100644 packages/patch-applier/Cargo.toml create mode 100644 packages/patch-applier/README.md create mode 100644 packages/patch-applier/package.json create mode 100644 packages/patch-applier/src/lib.rs create mode 100644 packages/patch-applier/src/main.rs create mode 100644 packages/patch-applier/tests/cli.rs diff --git a/packages/agenty-cli/src/components/MessageItem.tsx b/packages/agenty-cli/src/components/MessageItem.tsx index d99cadd..a09f7c6 100644 --- a/packages/agenty-cli/src/components/MessageItem.tsx +++ b/packages/agenty-cli/src/components/MessageItem.tsx @@ -9,7 +9,7 @@ import { type ShellOutputStream, type ToolDisplay, } from "./toolDisplay"; -import { Box, Text } from "./ui"; +import { Box, Pressable, Text } from "./ui"; const USER_MESSAGE_BACKGROUNDS: Record = { dark: "#2a3f47", @@ -155,11 +155,12 @@ function ToolCallLine({ const hasShellDetails = (display.shellCommands?.length ?? 0) > 0; const hasDetails = display.detailLines.length > 0 || hasShellDetails; return ( - @@ -187,7 +188,7 @@ function ToolCallLine({ marginLeft={4} /> ) : null} - + ); } @@ -256,7 +257,7 @@ function Rail({ onMouseClick?: () => void; }) { return ( - {children} - + ); } @@ -286,11 +288,12 @@ export const MessageItem = memo(({ }) => { if (item.type === "reasoning") { return ( - onToggleReasoning?.(item.id)} + disabled={!onToggleReasoning} + onPress={() => onToggleReasoning?.(item.id)} > {item.done @@ -304,7 +307,7 @@ export const MessageItem = memo(({ ) : null} - + ); } diff --git a/packages/agenty-cli/src/components/MessageList.tsx b/packages/agenty-cli/src/components/MessageList.tsx index 94db16e..836fd09 100644 --- a/packages/agenty-cli/src/components/MessageList.tsx +++ b/packages/agenty-cli/src/components/MessageList.tsx @@ -9,7 +9,7 @@ import { type ReactNode, useEffect, useMemo, useRef, useState } from "react"; import { useInput } from "../hooks/useInput"; import type { UIMessage } from "../state/store"; import { MessageItem, type MessageRenderItem } from "./MessageItem"; -import { Box, Text } from "./ui"; +import { Box, Pressable, Text } from "./ui"; const HINT_BACKGROUND = "#24383f"; const HINT_FOREGROUND = "#c7f5ff"; @@ -427,13 +427,9 @@ export function MessageList({ {showHint ? ( - - {` ${hintLabel} `} - + + {` ${hintLabel} `} + ) : null} diff --git a/packages/agenty-cli/src/components/toolDisplay.test.ts b/packages/agenty-cli/src/components/toolDisplay.test.ts index 3ab2fb7..7211a4b 100644 --- a/packages/agenty-cli/src/components/toolDisplay.test.ts +++ b/packages/agenty-cli/src/components/toolDisplay.test.ts @@ -84,6 +84,22 @@ describe("tool display", () => { summaryLines: ["update src/main.go"], }); + const completedDisplay = buildToolDisplay(toolCall( + "apply_patch", + { patch: "*** Begin Patch\n*** Add File: notes.txt\n+hello\n*** End Patch" }, + JSON.stringify({ + success: true, + files: [{ + path: "notes.txt", + diff: "--- /dev/null\n+++ b/notes.txt\n@@ -0,0 +1 @@\n+hello", + addedLines: 1, + removedLines: 0, + }], + }), + )); + expect(completedDisplay.summaryLines).toEqual(["notes.txt · +1 -0"]); + expect(completedDisplay.detailLines).toContain("+hello"); + const customDisplay = buildToolDisplay(toolCall("apply_patch", { patch: [ "*** Begin Patch", diff --git a/packages/agenty-cli/src/components/toolDisplay.ts b/packages/agenty-cli/src/components/toolDisplay.ts index 21c84b2..bc7f818 100644 --- a/packages/agenty-cli/src/components/toolDisplay.ts +++ b/packages/agenty-cli/src/components/toolDisplay.ts @@ -30,9 +30,6 @@ export interface ToolDisplay { const TOOL_LABELS: Record = { read_file: "Read file", - write_file: "Write file", - patch_file: "Edit file", - delete_file: "Delete file", grep: "Search text", glob: "Find files", ls: "List directory", @@ -123,11 +120,6 @@ function formatRange(input: JsonRecord | undefined, result: JsonRecord | undefin return start !== undefined ? `from line ${start}` : `through line ${end}`; } -function formatBytes(value: unknown): string { - const bytes = numberValue(value); - return bytes === undefined ? "" : `${bytes.toLocaleString()} bytes`; -} - function formatCount(value: unknown, singular: string, plural = `${singular}s`): string { const count = numberValue(value); if (count === undefined) { @@ -243,47 +235,6 @@ function readFileDisplay(input: JsonRecord | undefined, result: ToolResult | und }; } -function writeFileDisplay(input: JsonRecord | undefined, result: ToolResult | undefined): ToolDisplay { - const output = resultRecord(result); - const path = formatPath(input?.path ?? output?.path); - const created = booleanValue(output?.created); - const bytes = formatBytes(output?.bytesWritten); - const action = created === undefined ? "file" : created ? "created" : "updated"; - const summary = [path, action, bytes].filter(Boolean).join(" · "); - return { - label: TOOL_LABELS.write_file, - status: toolStatus(result), - summaryLines: [summary], - detailLines: [], - }; -} - -function patchFileDisplay(input: JsonRecord | undefined, result: ToolResult | undefined): ToolDisplay { - const output = resultRecord(result); - const path = formatPath(input?.path ?? output?.path); - const replacements = formatCount(output?.replacements, "replacement"); - const bytes = formatBytes(output?.bytesWritten); - return { - label: TOOL_LABELS.patch_file, - status: toolStatus(result), - summaryLines: [[path, replacements, bytes].filter(Boolean).join(" · ")], - detailLines: [], - }; -} - -function deleteFileDisplay(input: JsonRecord | undefined, result: ToolResult | undefined): ToolDisplay { - const output = resultRecord(result); - const path = formatPath(input?.path ?? output?.path); - const deleted = booleanValue(output?.deleted); - const action = deleted === undefined ? "file" : deleted ? "deleted" : "not deleted"; - return { - label: TOOL_LABELS.delete_file, - status: toolStatus(result), - summaryLines: [`${path} · ${action}`], - detailLines: [], - }; -} - function grepDisplay(input: JsonRecord | undefined, result: ToolResult | undefined): ToolDisplay { const output = resultRecord(result); const pattern = truncate(stringValue(input?.pattern, "pattern"), 72); @@ -460,15 +411,25 @@ function applyPatchDisplay(input: JsonRecord | undefined, result: ToolResult | u if (summaryLines.length === 0) { summaryLines.push("patch"); } - const visibleSummary = summaryLines.slice(0, MAX_SUMMARY_LINES); - if (summaryLines.length > visibleSummary.length) { - visibleSummary.push(`… ${summaryLines.length - visibleSummary.length} more operations`); + const output = resultRecord(result); + const files = Array.isArray(output?.files) ? output.files.filter(isRecord) : []; + const resultSummary = files.map((file) => { + const path = formatPath(file.path); + const added = numberValue(file.addedLines) ?? 0; + const removed = numberValue(file.removedLines) ?? 0; + return `${path} · +${added} -${removed}`; + }); + const visibleSummary = (resultSummary.length > 0 ? resultSummary : summaryLines).slice(0, MAX_SUMMARY_LINES); + const totalSummaryLines = resultSummary.length > 0 ? resultSummary.length : summaryLines.length; + if (totalSummaryLines > visibleSummary.length) { + visibleSummary.push(`… ${totalSummaryLines - visibleSummary.length} more files`); } + const diffDetails = files.flatMap((file) => splitLines(stringValue(file.diff))); return { label: TOOL_LABELS.apply_patch, status: toolStatus(result), summaryLines: visibleSummary, - detailLines: formatResultPreview(result), + detailLines: diffDetails.length > 0 ? diffDetails : formatResultPreview(result), }; } @@ -521,12 +482,6 @@ export function buildToolDisplay(toolCall: UIToolCall, expanded = true): ToolDis switch (toolCall.name) { case "read_file": return readFileDisplay(input, toolCall.result); - case "write_file": - return writeFileDisplay(input, toolCall.result); - case "patch_file": - return patchFileDisplay(input, toolCall.result); - case "delete_file": - return deleteFileDisplay(input, toolCall.result); case "grep": return grepDisplay(input, toolCall.result); case "glob": diff --git a/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go b/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go index ea13a14..5529c99 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go +++ b/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go @@ -1,25 +1,17 @@ package builtin import ( + "bytes" "context" "errors" "fmt" - "os" - "path/filepath" + "os/exec" "strings" + json "github.com/bytedance/sonic" + "github.com/masteryyh/agenty-core/pkg/agentloop" "github.com/masteryyh/agenty-core/pkg/domain/conversation" - "github.com/masteryyh/agenty-core/pkg/utils" -) - -const ( - patchBeginMarker = "*** Begin Patch" - patchEndMarker = "*** End Patch" - patchUpdateFile = "*** Update File:" - patchDeleteFile = "*** Delete File:" - patchAddFile = "*** Add File:" - patchMoveTo = "*** Move to:" ) type applyPatchTool struct { @@ -27,39 +19,32 @@ type applyPatchTool struct { } type applyPatchArguments struct { - Operation *conversation.ApplyPatchOperation `json:"operation,omitempty"` - Patch string `json:"patch,omitempty"` + Patch string `json:"patch"` } -type applyPatchOperationResult struct { - Type conversation.ApplyPatchOperationType `json:"type"` - Path string `json:"path"` - MoveTo string `json:"moveTo,omitempty"` +type applyPatchFileResult struct { + Path string `json:"path"` + Diff string `json:"diff"` + AddedLines int `json:"addedLines"` + RemovedLines int `json:"removedLines"` } type applyPatchResult struct { - Operations []applyPatchOperationResult `json:"operations"` + Success bool `json:"success"` + Cwd string `json:"cwd"` + Files []applyPatchFileResult `json:"files"` } func (tool *applyPatchTool) Definition() agentloop.ToolDefinition { - operationSchema := objectSchema( - map[string]agentloop.JSONSchema{ - "type": stringSchema("Operation type: create_file, update_file, or delete_file."), - "path": stringSchema("Absolute path or path relative to the session working directory."), - "diff": stringSchema("Headerless V4A diff body for create_file or update_file."), - }, - []string{"type", "path"}, - ) return agentloop.ToolDefinition{ Type: agentloop.ToolTypeApplyPatch, Name: "apply_patch", - Description: "Apply one native file operation or a complete V4A patch envelope.", + Description: "Apply a complete V4A patch atomically and return each file's final diff and line counts.", InputSchema: objectSchema( map[string]agentloop.JSONSchema{ - "operation": operationSchema, - "patch": stringSchema("Complete V4A patch envelope."), + "patch": stringSchema("Complete V4A patch envelope."), }, - nil, + []string{"patch"}, ), } } @@ -73,235 +58,42 @@ func (tool *applyPatchTool) Execute( if err := decodeArguments(input, &arguments); err != nil { return nil, fmt.Errorf("apply_patch: %w", err) } - - hasOperation := arguments.Operation != nil - hasPatch := arguments.Patch != "" - if hasOperation == hasPatch { - return nil, fmt.Errorf("apply_patch: exactly one of operation or patch is required") - } - - operations := make([]conversation.ApplyPatchOperation, 0, 1) - if hasOperation { - operations = append(operations, *arguments.Operation) - } else { - parsed, err := parsePatchEnvelope(arguments.Patch) - if err != nil { - return nil, fmt.Errorf("apply_patch: parse patch: %w", err) - } - operations = parsed + if strings.TrimSpace(arguments.Patch) == "" { + return nil, fmt.Errorf("apply_patch: patch must not be empty") } tool.fileSystem.mu.Lock() defer tool.fileSystem.mu.Unlock() - results := make([]applyPatchOperationResult, 0, len(operations)) - for index, operation := range operations { - if err := ctx.Err(); err != nil { - return nil, err + command := exec.CommandContext(ctx, "apply_patch") + if strings.TrimSpace(callContext.Cwd) != "" { + command.Dir = callContext.Cwd + } + command.Stdin = strings.NewReader(arguments.Patch) + var stdout bytes.Buffer + var stderr bytes.Buffer + command.Stdout = &stdout + command.Stderr = &stderr + if err := command.Run(); err != nil { + if errors.Is(ctx.Err(), context.Canceled) { + return nil, ctx.Err() } - result, err := executeApplyPatchOperation(callContext.Cwd, operation) - if err != nil { - return nil, fmt.Errorf( - "apply_patch: operation %d %s %q: %w", - index+1, - operation.Type, - operation.Path, - err, - ) + message := strings.TrimSpace(stderr.String()) + if message == "" { + message = strings.TrimSpace(stdout.String()) } - results = append(results, result) - } - - return resultContent(applyPatchResult{Operations: results}) -} - -func parsePatchEnvelope(patch string) ([]conversation.ApplyPatchOperation, error) { - lines := normalizePatchEnvelopeLines(patch) - if len(lines) < 3 || lines[0] != patchBeginMarker { - return nil, fmt.Errorf("patch must start with %q", patchBeginMarker) - } - if lines[len(lines)-1] != patchEndMarker { - return nil, fmt.Errorf("patch must end with %q", patchEndMarker) - } - - operations := make([]conversation.ApplyPatchOperation, 0) - for index := 1; index < len(lines)-1; { - operation, nextIndex, err := parsePatchEnvelopeOperation(lines, index) - if err != nil { - return nil, err + if message == "" { + message = err.Error() } - operations = append(operations, operation) - index = nextIndex - } - if len(operations) == 0 { - return nil, fmt.Errorf("patch contains no file operations") - } - return operations, nil -} - -func normalizePatchEnvelopeLines(patch string) []string { - lines := strings.Split(strings.ReplaceAll(patch, "\r\n", "\n"), "\n") - for index := range lines { - lines[index] = strings.TrimSuffix(lines[index], "\r") - } - if len(lines) > 0 && lines[len(lines)-1] == "" { - lines = lines[:len(lines)-1] - } - return lines -} - -func parsePatchEnvelopeOperation( - lines []string, - index int, -) (conversation.ApplyPatchOperation, int, error) { - header := lines[index] - operation := conversation.ApplyPatchOperation{} - switch { - case strings.HasPrefix(header, patchUpdateFile): - operation.Type = conversation.ApplyPatchUpdateFile - operation.Path = strings.TrimSpace(strings.TrimPrefix(header, patchUpdateFile)) - case strings.HasPrefix(header, patchDeleteFile): - operation.Type = conversation.ApplyPatchDeleteFile - operation.Path = strings.TrimSpace(strings.TrimPrefix(header, patchDeleteFile)) - case strings.HasPrefix(header, patchAddFile): - operation.Type = conversation.ApplyPatchCreateFile - operation.Path = strings.TrimSpace(strings.TrimPrefix(header, patchAddFile)) - default: - return conversation.ApplyPatchOperation{}, 0, fmt.Errorf( - "invalid patch header at line %d: %s", - index+1, - header, - ) - } - if operation.Path == "" { - return conversation.ApplyPatchOperation{}, 0, fmt.Errorf("operation at line %d has an empty path", index+1) - } - - index++ - if operation.Type == conversation.ApplyPatchUpdateFile && index < len(lines)-1 && - strings.HasPrefix(lines[index], patchMoveTo) { - operation.MoveTo = strings.TrimSpace(strings.TrimPrefix(lines[index], patchMoveTo)) - if operation.MoveTo == "" { - return conversation.ApplyPatchOperation{}, 0, fmt.Errorf("move at line %d has an empty path", index+1) - } - index++ + return nil, fmt.Errorf("apply_patch: %s", message) } - bodyStart := index - for index < len(lines)-1 && !isPatchOperationHeader(lines[index]) { - index++ - } - body := lines[bodyStart:index] - if operation.Type == conversation.ApplyPatchDeleteFile && len(body) > 0 { - return conversation.ApplyPatchOperation{}, 0, fmt.Errorf( - "delete operation for %q must not contain a diff body", - operation.Path, - ) - } - operation.Diff = strings.Join(body, "\n") - return operation, index, nil -} - -func isPatchOperationHeader(line string) bool { - return strings.HasPrefix(line, patchUpdateFile) || - strings.HasPrefix(line, patchDeleteFile) || - strings.HasPrefix(line, patchAddFile) -} - -func executeApplyPatchOperation( - cwd string, - operation conversation.ApplyPatchOperation, -) (applyPatchOperationResult, error) { - path, err := resolvePath(operation.Path, cwd, false) - if err != nil { - return applyPatchOperationResult{}, err - } - result := applyPatchOperationResult{Type: operation.Type, Path: path} - - switch operation.Type { - case conversation.ApplyPatchCreateFile: - content, err := utils.ApplyDiff("", operation.Diff, utils.ApplyDiffCreate) - if err != nil { - return applyPatchOperationResult{}, fmt.Errorf("apply create diff: %w", err) - } - if _, err := writeTextFile(path, content, 0o644); err != nil { - return applyPatchOperationResult{}, err - } - case conversation.ApplyPatchUpdateFile: - if err := updateFileWithDiff(path, operation.Diff); err != nil { - return applyPatchOperationResult{}, err - } - if operation.MoveTo != "" { - moveTo, err := resolvePath(operation.MoveTo, cwd, false) - if err != nil { - return applyPatchOperationResult{}, fmt.Errorf("resolve move destination: %w", err) - } - if err := movePatchedFile(path, moveTo); err != nil { - return applyPatchOperationResult{}, err - } - result.MoveTo = moveTo - } - case conversation.ApplyPatchDeleteFile: - if err := removeApplyPatchFile(path); err != nil { - return applyPatchOperationResult{}, err - } - default: - return applyPatchOperationResult{}, fmt.Errorf("unsupported operation type %q", operation.Type) - } - - return result, nil -} - -func updateFileWithDiff(path, diff string) error { - info, err := regularFileInfo(path) - if err != nil { - return err - } - data, err := os.ReadFile(path) - if err != nil { - return fmt.Errorf("read %q: %w", path, err) - } - updated, err := utils.ApplyDiff(string(data), diff, utils.ApplyDiffDefault) - if err != nil { - return fmt.Errorf("apply update diff: %w", err) - } - if _, err := writeTextFile(path, updated, info.Mode().Perm()); err != nil { - return err - } - return nil -} - -func movePatchedFile(source, destination string) error { - if source == destination { - return nil - } - if _, err := os.Lstat(destination); err == nil { - return fmt.Errorf("move destination %q already exists", destination) - } else if !errors.Is(err, os.ErrNotExist) { - return fmt.Errorf("inspect move destination %q: %w", destination, err) - } - if err := os.MkdirAll(filepath.Dir(destination), 0o755); err != nil { - return fmt.Errorf("create move destination parent for %q: %w", destination, err) - } - if err := os.Rename(source, destination); err != nil { - return fmt.Errorf("move %q to %q: %w", source, destination, err) - } - return nil -} - -func removeApplyPatchFile(path string) error { - info, err := os.Lstat(path) - if err != nil { - return fmt.Errorf("inspect %q: %w", path, err) - } - if info.IsDir() { - return fmt.Errorf("path %q is a directory", path) - } - if !info.Mode().IsRegular() && info.Mode()&os.ModeSymlink == 0 { - return fmt.Errorf("path %q is not a file or symbolic link", path) + var result applyPatchResult + if err := json.Unmarshal(stdout.Bytes(), &result); err != nil { + return nil, fmt.Errorf("apply_patch: decode helper result: %w", err) } - if err := os.Remove(path); err != nil { - return fmt.Errorf("remove %q: %w", path, err) + if !result.Success { + return nil, fmt.Errorf("apply_patch: helper returned an unsuccessful result") } - return nil + return resultContent(result) } diff --git a/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go b/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go index 3f704e2..5083b9a 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go @@ -1,190 +1,87 @@ package builtin import ( + "context" "os" "path/filepath" + "runtime" "strings" "testing" json "github.com/bytedance/sonic" "github.com/masteryyh/agenty-core/pkg/agentloop" - "github.com/masteryyh/agenty-core/pkg/domain/conversation" ) -func TestParsePatchEnvelopePreservesRepeatedFileOperations(t *testing.T) { - t.Parallel() - - patch := `*** Begin Patch -*** Update File: notes.txt -@@ --one -+two -*** Update File: notes.txt -@@ --two -+three -*** End Patch` - operations, err := parsePatchEnvelope(patch) - if err != nil { - t.Fatal(err) - } - if len(operations) != 2 { - t.Fatalf("operations = %d, want 2", len(operations)) - } - for index, operation := range operations { - if operation.Path != "notes.txt" { - t.Errorf("operation %d path = %q, want notes.txt", index, operation.Path) - } +func TestApplyPatchToolReturnsStructuredHelperResult(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("shell fixture uses a POSIX script") } - if operations[0].Diff == operations[1].Diff { - t.Errorf("repeated operations were collapsed: %#v", operations) - } -} - -func TestApplyPatchToolExecutesOperationsInOrder(t *testing.T) { - t.Parallel() directory := t.TempDir() + installApplyPatchFixture(t, directory, `#!/bin/sh +patch=$(cat) +case "$patch" in + *"*** Begin Patch"*) ;; + *) exit 2 ;; +esac +printf '%s\n' '{"success":true,"cwd":"/workspace","files":[{"path":"notes.txt","diff":"--- /dev/null\\n+++ b/notes.txt\\n@@ -0,0 +1 @@\\n+hello\\n","addedLines":1,"removedLines":0}]}' +`) + t.Setenv("PATH", directory+string(os.PathListSeparator)+os.Getenv("PATH")) + tool := &applyPatchTool{fileSystem: &fileSystem{}} - patch := `*** Begin Patch -*** Add File: notes.txt -+one -*** Update File: notes.txt -@@ --one -+two -*** Update File: notes.txt -@@ --two -+three -*** End Patch` - content, err := executeApplyPatchTool(t, tool, directory, applyPatchArguments{Patch: patch}) + input, err := json.Marshal(applyPatchArguments{Patch: "*** Begin Patch\n*** End Patch"}) if err != nil { t.Fatal(err) } - if len(content) != 1 { - t.Fatalf("content = %d blocks, want 1", len(content)) - } - data, err := os.ReadFile(filepath.Join(directory, "notes.txt")) + content, err := tool.Execute(t.Context(), agentloop.CallContext{Cwd: directory}, input) if err != nil { t.Fatal(err) } - if string(data) != "three" { - t.Errorf("notes.txt = %q, want three", data) - } -} - -func TestApplyPatchToolSupportsMoveThenUpdateDestination(t *testing.T) { - t.Parallel() - - directory := t.TempDir() - if err := os.WriteFile(filepath.Join(directory, "old.txt"), []byte("one\n"), 0o644); err != nil { - t.Fatal(err) - } - tool := &applyPatchTool{fileSystem: &fileSystem{}} - patch := `*** Begin Patch -*** Update File: old.txt -*** Move to: new.txt -@@ --one -+two -*** Update File: new.txt -@@ --two -+three -*** End Patch` - if _, err := executeApplyPatchTool(t, tool, directory, applyPatchArguments{Patch: patch}); err != nil { - t.Fatal(err) - } - if _, err := os.Stat(filepath.Join(directory, "old.txt")); !os.IsNotExist(err) { - t.Fatalf("old.txt still exists: %v", err) + if len(content) != 1 { + t.Fatalf("content = %d blocks, want 1", len(content)) } - data, err := os.ReadFile(filepath.Join(directory, "new.txt")) + encoded, err := json.MarshalString(content[0]) if err != nil { t.Fatal(err) } - if string(data) != "three\n" { - t.Errorf("new.txt = %q, want %q", data, "three\n") + if !strings.Contains(encoded, `\"addedLines\":1`) || !strings.Contains(encoded, `\"removedLines\":0`) { + t.Errorf("tool result = %s", encoded) } } -func TestApplyPatchToolKeepsEarlierOperationsOnFailure(t *testing.T) { - t.Parallel() - - directory := t.TempDir() - tool := &applyPatchTool{fileSystem: &fileSystem{}} - patch := `*** Begin Patch -*** Add File: created.txt -+created -*** Update File: missing.txt -@@ --missing -+updated -*** End Patch` - _, err := executeApplyPatchTool(t, tool, directory, applyPatchArguments{Patch: patch}) - if err == nil || !strings.Contains(err.Error(), "operation 2") { - t.Fatalf("Execute() error = %v, want operation 2 failure", err) - } - data, readErr := os.ReadFile(filepath.Join(directory, "created.txt")) - if readErr != nil { - t.Fatal(readErr) +func TestApplyPatchToolReportsHelperFailure(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("shell fixture uses a POSIX script") } - if string(data) != "created" { - t.Errorf("created.txt = %q, want created", data) - } -} - -func TestApplyPatchToolExecutesNativeOperation(t *testing.T) { - t.Parallel() directory := t.TempDir() + installApplyPatchFixture(t, directory, "#!/bin/sh\necho 'conflicting operations' >&2\nexit 1\n") + t.Setenv("PATH", directory+string(os.PathListSeparator)+os.Getenv("PATH")) + tool := &applyPatchTool{fileSystem: &fileSystem{}} - operation := conversation.ApplyPatchOperation{ - Type: conversation.ApplyPatchCreateFile, - Path: "native.txt", - Diff: "+native", - } - if _, err := executeApplyPatchTool(t, tool, directory, applyPatchArguments{Operation: &operation}); err != nil { - t.Fatal(err) - } - data, err := os.ReadFile(filepath.Join(directory, "native.txt")) + input, err := json.Marshal(applyPatchArguments{Patch: "*** Begin Patch\n*** End Patch"}) if err != nil { t.Fatal(err) } - if string(data) != "native" { - t.Errorf("native.txt = %q, want native", data) + _, err = tool.Execute(t.Context(), agentloop.CallContext{Cwd: directory}, input) + if err == nil || !strings.Contains(err.Error(), "conflicting operations") { + t.Fatalf("Execute() error = %v", err) } } -func TestApplyPatchToolRejectsMalformedEnvelopeBeforeWriting(t *testing.T) { - t.Parallel() - - directory := t.TempDir() +func TestApplyPatchToolRequiresPatch(t *testing.T) { tool := &applyPatchTool{fileSystem: &fileSystem{}} - patch := `*** Begin Patch -*** Add File: created.txt -+created` - _, err := executeApplyPatchTool(t, tool, directory, applyPatchArguments{Patch: patch}) - if err == nil || !strings.Contains(err.Error(), "must end") { - t.Fatalf("Execute() error = %v, want missing end marker", err) - } - if _, statErr := os.Stat(filepath.Join(directory, "created.txt")); !os.IsNotExist(statErr) { - t.Fatalf("created.txt exists after parse failure: %v", statErr) + _, err := tool.Execute(context.Background(), agentloop.CallContext{}, []byte(`{}`)) + if err == nil || !strings.Contains(err.Error(), "patch must not be empty") { + t.Fatalf("Execute() error = %v", err) } } -func executeApplyPatchTool( - t *testing.T, - tool *applyPatchTool, - cwd string, - arguments applyPatchArguments, -) (conversation.Content, error) { +func installApplyPatchFixture(t *testing.T, directory, script string) { t.Helper() - - input, err := json.Marshal(arguments) - if err != nil { + path := filepath.Join(directory, "apply_patch") + if err := os.WriteFile(path, []byte(script), 0o755); err != nil { t.Fatal(err) } - return tool.Execute(t.Context(), agentloop.CallContext{Cwd: cwd}, input) } diff --git a/packages/agenty-core/pkg/agentloop/builtin/file.go b/packages/agenty-core/pkg/agentloop/builtin/file.go index 70910bd..9422c93 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/file.go +++ b/packages/agenty-core/pkg/agentloop/builtin/file.go @@ -3,10 +3,8 @@ package builtin import ( "bufio" "context" - "errors" "fmt" "os" - "path/filepath" "strconv" "strings" "unicode/utf8" @@ -197,267 +195,3 @@ func truncateUTF8(value string, maxBytes int) string { } return value } - -type writeFileTool struct { - fileSystem *fileSystem -} - -type writeFileArguments struct { - Path string `json:"path"` - Content *string `json:"content"` -} - -type writeFileResult struct { - Path string `json:"path"` - BytesWritten int `json:"bytesWritten"` - Created bool `json:"created"` -} - -func (tool *writeFileTool) Definition() agentloop.ToolDefinition { - return agentloop.ToolDefinition{ - Name: "write_file", - Description: "Create or overwrite a file. Missing parent directories are created. " + - "Relative paths resolve from the session working directory.", - InputSchema: objectSchema( - map[string]agentloop.JSONSchema{ - "path": stringSchema("Absolute path or path relative to the session working directory."), - "content": stringSchema("Complete file content to write."), - }, - []string{"path", "content"}, - ), - } -} - -func (tool *writeFileTool) Execute( - ctx context.Context, - callContext agentloop.CallContext, - input []byte, -) (conversation.Content, error) { - var arguments writeFileArguments - if err := decodeArguments(input, &arguments); err != nil { - return nil, fmt.Errorf("write_file: %w", err) - } - if arguments.Content == nil { - return nil, fmt.Errorf("write_file: content is required") - } - if err := ctx.Err(); err != nil { - return nil, err - } - - path, err := resolvePath(arguments.Path, callContext.Cwd, false) - if err != nil { - return nil, fmt.Errorf("write_file: %w", err) - } - - tool.fileSystem.mu.Lock() - defer tool.fileSystem.mu.Unlock() - if err := ctx.Err(); err != nil { - return nil, err - } - - created, err := writeTextFile(path, *arguments.Content, 0o644) - if err != nil { - return nil, fmt.Errorf("write_file: %w", err) - } - return resultContent(writeFileResult{ - Path: path, - BytesWritten: len(*arguments.Content), - Created: created, - }) -} - -func writeTextFile(path, content string, defaultMode os.FileMode) (bool, error) { - created := false - mode := defaultMode - info, err := os.Stat(path) - switch { - case errors.Is(err, os.ErrNotExist): - created = true - case err != nil: - return false, fmt.Errorf("inspect %q: %w", path, err) - default: - if info.IsDir() { - return false, fmt.Errorf("path %q is a directory", path) - } - if !info.Mode().IsRegular() { - return false, fmt.Errorf("path %q is not a regular file", path) - } - mode = info.Mode().Perm() - } - - if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { - return false, fmt.Errorf("create parent directory for %q: %w", path, err) - } - if err := os.WriteFile(path, []byte(content), mode); err != nil { - return false, fmt.Errorf("write %q: %w", path, err) - } - return created, nil -} - -type patchFileTool struct { - fileSystem *fileSystem -} - -type patchFileArguments struct { - Path string `json:"path"` - OldText string `json:"old_text"` - NewText *string `json:"new_text"` - ReplaceAll bool `json:"replace_all,omitempty"` -} - -type patchFileResult struct { - Path string `json:"path"` - Replacements int `json:"replacements"` - BytesWritten int `json:"bytesWritten"` -} - -func (tool *patchFileTool) Definition() agentloop.ToolDefinition { - return agentloop.ToolDefinition{ - Name: "patch_file", - Description: "Replace exact text in an existing file. By default old_text must occur exactly once; " + - "set replace_all to replace every occurrence.", - InputSchema: objectSchema( - map[string]agentloop.JSONSchema{ - "path": stringSchema("Absolute path or path relative to the session working directory."), - "old_text": stringSchema("Exact text that must already exist in the file."), - "new_text": stringSchema("Replacement text. May be empty to remove old_text."), - "replace_all": booleanSchema("Replace all occurrences instead of requiring one unique occurrence."), - }, - []string{"path", "old_text", "new_text"}, - ), - } -} - -func (tool *patchFileTool) Execute( - ctx context.Context, - callContext agentloop.CallContext, - input []byte, -) (conversation.Content, error) { - var arguments patchFileArguments - if err := decodeArguments(input, &arguments); err != nil { - return nil, fmt.Errorf("patch_file: %w", err) - } - - if arguments.OldText == "" { - return nil, fmt.Errorf("patch_file: old_text must not be empty") - } - if arguments.NewText == nil { - return nil, fmt.Errorf("patch_file: new_text is required") - } - if err := ctx.Err(); err != nil { - return nil, err - } - - path, err := resolvePath(arguments.Path, callContext.Cwd, false) - if err != nil { - return nil, fmt.Errorf("patch_file: %w", err) - } - - tool.fileSystem.mu.Lock() - defer tool.fileSystem.mu.Unlock() - if err := ctx.Err(); err != nil { - return nil, err - } - - info, err := regularFileInfo(path) - if err != nil { - return nil, fmt.Errorf("patch_file: %w", err) - } - data, err := os.ReadFile(path) - if err != nil { - return nil, fmt.Errorf("patch_file: read %q: %w", path, err) - } - - content := string(data) - occurrences := strings.Count(content, arguments.OldText) - if occurrences == 0 { - return nil, fmt.Errorf("patch_file: old_text was not found in %q", path) - } - if !arguments.ReplaceAll && occurrences != 1 { - return nil, fmt.Errorf( - "patch_file: old_text occurs %d times in %q; set replace_all to replace every occurrence", - occurrences, - path, - ) - } - - replacements := 1 - if arguments.ReplaceAll { - replacements = occurrences - } - updated := strings.Replace(content, arguments.OldText, *arguments.NewText, replacements) - if _, err := writeTextFile(path, updated, info.Mode().Perm()); err != nil { - return nil, fmt.Errorf("patch_file: %w", err) - } - return resultContent(patchFileResult{ - Path: path, - Replacements: replacements, - BytesWritten: len(updated), - }) -} - -type deleteFileTool struct { - fileSystem *fileSystem -} - -type deleteFileArguments struct { - Path string `json:"path"` -} - -type deleteFileResult struct { - Path string `json:"path"` - Deleted bool `json:"deleted"` -} - -func (tool *deleteFileTool) Definition() agentloop.ToolDefinition { - return agentloop.ToolDefinition{ - Name: "delete_file", - Description: "Delete one file or symbolic link. Directories are rejected and are never removed recursively.", - InputSchema: objectSchema( - map[string]agentloop.JSONSchema{ - "path": stringSchema("Absolute path or path relative to the session working directory."), - }, - []string{"path"}, - ), - } -} - -func (tool *deleteFileTool) Execute( - ctx context.Context, - callContext agentloop.CallContext, - input []byte, -) (conversation.Content, error) { - var arguments deleteFileArguments - if err := decodeArguments(input, &arguments); err != nil { - return nil, fmt.Errorf("delete_file: %w", err) - } - if err := ctx.Err(); err != nil { - return nil, err - } - - path, err := resolvePath(arguments.Path, callContext.Cwd, false) - if err != nil { - return nil, fmt.Errorf("delete_file: %w", err) - } - - tool.fileSystem.mu.Lock() - defer tool.fileSystem.mu.Unlock() - if err := ctx.Err(); err != nil { - return nil, err - } - - info, err := os.Lstat(path) - if err != nil { - return nil, fmt.Errorf("delete_file: inspect %q: %w", path, err) - } - if info.IsDir() { - return nil, fmt.Errorf("delete_file: path %q is a directory", path) - } - if !info.Mode().IsRegular() && info.Mode()&os.ModeSymlink == 0 { - return nil, fmt.Errorf("delete_file: path %q is not a file or symbolic link", path) - } - if err := os.Remove(path); err != nil { - return nil, fmt.Errorf("delete_file: remove %q: %w", path, err) - } - return resultContent(deleteFileResult{Path: path, Deleted: true}) -} diff --git a/packages/agenty-core/pkg/agentloop/builtin/file_test.go b/packages/agenty-core/pkg/agentloop/builtin/file_test.go index e5f1e7e..0eaf858 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/file_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/file_test.go @@ -1,7 +1,6 @@ package builtin_test import ( - "fmt" "os" "path/filepath" "strings" @@ -62,239 +61,19 @@ func TestReadFileValidatesLineBounds(t *testing.T) { arguments string wantError string }{ - { - name: "explicit zero start", - arguments: `{"path":"short.txt","start_line":0}`, - wantError: "start_line must be positive", - }, - { - name: "reversed range", - arguments: `{"path":"short.txt","start_line":2,"end_line":1}`, - wantError: "start_line must not exceed end_line", - }, - { - name: "start after end of file", - arguments: `{"path":"short.txt","start_line":3}`, - wantError: "exceeds file length 2", - }, + {name: "explicit zero start", arguments: `{"path":"short.txt","start_line":0}`, wantError: "start_line must be positive"}, + {name: "reversed range", arguments: `{"path":"short.txt","start_line":2,"end_line":1}`, wantError: "start_line must not exceed end_line"}, + {name: "start after end of file", arguments: `{"path":"short.txt","start_line":3}`, wantError: "exceeds file length 2"}, } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { t.Parallel() - _, err := executeTool(t, registry, "read_file", directory, tt.arguments) - if err == nil || !strings.Contains(err.Error(), tt.wantError) { - t.Fatalf("error = %v, want containing %q", err, tt.wantError) + _, err := executeTool(t, registry, "read_file", directory, test.arguments) + if err == nil || !strings.Contains(err.Error(), test.wantError) { + t.Fatalf("error = %v, want containing %q", err, test.wantError) } }) } } - -func TestWriteFileRequiresContentField(t *testing.T) { - t.Parallel() - - _, err := executeTool(t, newRegistry(t), "write_file", t.TempDir(), `{"path":"empty.txt"}`) - if err == nil || !strings.Contains(err.Error(), "content is required") { - t.Fatalf("error = %v, want missing content error", err) - } -} - -func TestWriteFileCreatesAndOverwritesFile(t *testing.T) { - t.Parallel() - - directory := t.TempDir() - registry := newRegistry(t) - path := filepath.Join(directory, "nested", "created.txt") - - encoded, err := executeTool( - t, - registry, - "write_file", - directory, - `{"path":"nested/created.txt","content":"first"}`, - ) - if err != nil { - t.Fatal(err) - } - created := decodeResult[struct { - Path string `json:"path"` - BytesWritten int `json:"bytesWritten"` - Created bool `json:"created"` - }](t, encoded) - if !created.Created || created.Path != path || created.BytesWritten != 5 { - t.Errorf("created result = %+v", created) - } - - if err := os.Chmod(path, 0o600); err != nil { - t.Fatal(err) - } - encoded, err = executeTool( - t, - registry, - "write_file", - directory, - `{"path":"nested/created.txt","content":"second"}`, - ) - if err != nil { - t.Fatal(err) - } - overwritten := decodeResult[struct { - BytesWritten int `json:"bytesWritten"` - Created bool `json:"created"` - }](t, encoded) - if overwritten.Created || overwritten.BytesWritten != 6 { - t.Errorf("overwrite result = %+v", overwritten) - } - data, err := os.ReadFile(path) - if err != nil { - t.Fatal(err) - } - if string(data) != "second" { - t.Errorf("file content = %q, want second", data) - } - info, err := os.Stat(path) - if err != nil { - t.Fatal(err) - } - if info.Mode().Perm() != 0o600 { - t.Errorf("file mode = %o, want 600", info.Mode().Perm()) - } -} - -func TestPatchFile(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - original string - arguments string - wantContent string - wantCount int - wantError string - }{ - { - name: "unique replacement", - original: "before old after", - arguments: `{"path":"file.txt","old_text":"old","new_text":"new"}`, - wantContent: "before new after", - wantCount: 1, - }, - { - name: "replace all", - original: "old and old", - arguments: `{"path":"file.txt","old_text":"old","new_text":"new","replace_all":true}`, - wantContent: "new and new", - wantCount: 2, - }, - { - name: "reject missing new text", - original: "old", - arguments: `{"path":"file.txt","old_text":"old"}`, - wantError: "new_text is required", - }, - { - name: "reject ambiguous replacement", - original: "old and old", - arguments: `{"path":"file.txt","old_text":"old","new_text":"new"}`, - wantError: "old_text occurs 2 times", - }, - { - name: "reject missing text", - original: "unchanged", - arguments: `{"path":"file.txt","old_text":"missing","new_text":"new"}`, - wantError: "old_text was not found", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - directory := t.TempDir() - path := filepath.Join(directory, "file.txt") - if err := os.WriteFile(path, []byte(tt.original), 0o640); err != nil { - t.Fatal(err) - } - - encoded, err := executeTool(t, newRegistry(t), "patch_file", directory, tt.arguments) - if tt.wantError != "" { - if err == nil || !strings.Contains(err.Error(), tt.wantError) { - t.Fatalf("error = %v, want containing %q", err, tt.wantError) - } - return - } - if err != nil { - t.Fatal(err) - } - - result := decodeResult[struct { - Replacements int `json:"replacements"` - BytesWritten int `json:"bytesWritten"` - }](t, encoded) - if result.Replacements != tt.wantCount || result.BytesWritten != len(tt.wantContent) { - t.Errorf("patch result = %+v", result) - } - data, err := os.ReadFile(path) - if err != nil { - t.Fatal(err) - } - if string(data) != tt.wantContent { - t.Errorf("file content = %q, want %q", data, tt.wantContent) - } - info, err := os.Stat(path) - if err != nil { - t.Fatal(err) - } - if info.Mode().Perm() != 0o640 { - t.Errorf("file mode = %o, want 640", info.Mode().Perm()) - } - }) - } -} - -func TestDeleteFile(t *testing.T) { - t.Parallel() - - directory := t.TempDir() - filePath := filepath.Join(directory, "remove.txt") - if err := os.WriteFile(filePath, []byte("remove me"), 0o644); err != nil { - t.Fatal(err) - } - registry := newRegistry(t) - - encoded, err := executeTool( - t, - registry, - "delete_file", - directory, - `{"path":"remove.txt"}`, - ) - if err != nil { - t.Fatal(err) - } - result := decodeResult[struct { - Path string `json:"path"` - Deleted bool `json:"deleted"` - }](t, encoded) - if !result.Deleted || result.Path != filePath { - t.Errorf("delete result = %+v", result) - } - if _, err := os.Stat(filePath); !os.IsNotExist(err) { - t.Fatalf("deleted file stat error = %v, want not exist", err) - } - - _, err = executeTool( - t, - registry, - "delete_file", - directory, - fmt.Sprintf(`{"path":%q}`, directory), - ) - if err == nil || !strings.Contains(err.Error(), "is a directory") { - t.Fatalf("directory delete error = %v", err) - } - if _, err := os.Stat(directory); err != nil { - t.Fatalf("directory was removed: %v", err) - } -} diff --git a/packages/agenty-core/pkg/agentloop/builtin/register.go b/packages/agenty-core/pkg/agentloop/builtin/register.go index 39460de..94ed7a5 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/register.go +++ b/packages/agenty-core/pkg/agentloop/builtin/register.go @@ -15,9 +15,6 @@ func RegisterAll(registry *agentloop.Registry) error { tools := []agentloop.Tool{ &shellTool{}, &readFileTool{fileSystem: fileSystem}, - &writeFileTool{fileSystem: fileSystem}, - &patchFileTool{fileSystem: fileSystem}, - &deleteFileTool{fileSystem: fileSystem}, &applyPatchTool{fileSystem: fileSystem}, &grepTool{fileSystem: fileSystem}, &globTool{fileSystem: fileSystem}, diff --git a/packages/agenty-core/pkg/agentloop/builtin/register_test.go b/packages/agenty-core/pkg/agentloop/builtin/register_test.go index a813674..295ecf5 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/register_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/register_test.go @@ -23,14 +23,11 @@ func TestRegisterAll(t *testing.T) { wantNames := []string{ "apply_patch", - "delete_file", "glob", "grep", "ls", - "patch_file", "read_file", "shell", - "write_file", } definitions := registry.Definitions() if len(definitions) != len(wantNames) { diff --git a/packages/agenty-core/pkg/agentloop/compaction.go b/packages/agenty-core/pkg/agentloop/compaction.go index 01d1cff..f725f72 100644 --- a/packages/agenty-core/pkg/agentloop/compaction.go +++ b/packages/agenty-core/pkg/agentloop/compaction.go @@ -105,7 +105,7 @@ func (engine *Engine) compactPreparedForWindow( baseRequest := Request{ SystemPrompt: prepared.systemPrompt, Messages: baseMessages, - Tools: engine.tools.Definitions(), + Tools: engine.toolDefinitions(prepared.freeFormTool), MaxOutputTokens: prepared.maxOutputTokens, ReasoningEffort: preparedReasoningEffort(prepared), } diff --git a/packages/agenty-core/pkg/agentloop/engine.go b/packages/agenty-core/pkg/agentloop/engine.go index 959a961..bcd8107 100644 --- a/packages/agenty-core/pkg/agentloop/engine.go +++ b/packages/agenty-core/pkg/agentloop/engine.go @@ -232,7 +232,8 @@ func (engine *Engine) Compact( model: resources.model, caller: resources.caller, systemPrompt: resources.systemPrompt, - maxOutputTokens: DefaultMaxOutputTokens, + freeFormTool: resources.freeFormTool, + maxOutputTokens: modelMaxOutputTokens(resources.model), } event, err := engine.compactPrepared(runCtx, prepared, conversation.CompactionTriggerManual) if err != nil { @@ -312,7 +313,8 @@ func (engine *Engine) SetModel( model: source.model, caller: source.caller, systemPrompt: source.systemPrompt, - maxOutputTokens: DefaultMaxOutputTokens, + freeFormTool: source.freeFormTool, + maxOutputTokens: modelMaxOutputTokens(source.model), } request := engine.sessionRequestForWindow(prepared, targetContextWindow) if ShouldCompact(estimateRequestTokens(request), targetContextWindow) { @@ -405,6 +407,7 @@ type preparedExecution struct { model catalog.Model caller Caller systemPrompt string + freeFormTool bool maxOutputTokens int64 userMessage conversation.Message eventSequence uint64 @@ -414,6 +417,7 @@ type executionResources struct { model catalog.Model caller Caller systemPrompt string + freeFormTool bool } func (engine *Engine) prepare( @@ -469,11 +473,19 @@ func (engine *Engine) prepare( model: resources.model, caller: resources.caller, systemPrompt: resources.systemPrompt, - maxOutputTokens: DefaultMaxOutputTokens, + freeFormTool: resources.freeFormTool, + maxOutputTokens: modelMaxOutputTokens(resources.model), userMessage: userMessage, }, nil } +func modelMaxOutputTokens(model catalog.Model) int64 { + if model.MaxOutputTokens > 0 { + return model.MaxOutputTokens + } + return DefaultMaxOutputTokens +} + func (engine *Engine) loadResources( ctx context.Context, runCtx context.Context, @@ -495,7 +507,9 @@ func (engine *Engine) loadResources( return nil, err } - systemPrompt, err := agentDefinition.ResolveSystemPrompt() + systemPrompt, err := agentDefinition.ResolveSystemPrompt(agent.SystemPromptOptions{ + UseApplyPatchShell: !provider.FreeFormTool, + }) if err != nil { return nil, apperrors.WrapError(apperrors.CodeInternal, "failed to resolve system prompt", err) } @@ -503,7 +517,12 @@ func (engine *Engine) loadResources( if err != nil { return nil, apperrors.WrapError(apperrors.CodeInternal, "failed to create LLM caller", err) } - return &executionResources{model: *model, caller: caller, systemPrompt: systemPrompt}, nil + return &executionResources{ + model: *model, + caller: caller, + systemPrompt: systemPrompt, + freeFormTool: provider.FreeFormTool, + }, nil } func (engine *Engine) loadCatalogModel( @@ -535,13 +554,28 @@ func (engine *Engine) sessionRequestForWindow(prepared *preparedExecution, conte request := Request{ SystemPrompt: prepared.systemPrompt, Messages: sessionMessages(prepared.session), - Tools: engine.tools.Definitions(), + Tools: engine.toolDefinitions(prepared.freeFormTool), MaxOutputTokens: prepared.maxOutputTokens, ReasoningEffort: preparedReasoningEffort(prepared), } return fitCompactedRequest(request, contextWindow) } +func (engine *Engine) toolDefinitions(freeFormTool bool) []ToolDefinition { + definitions := engine.tools.Definitions() + if freeFormTool { + return definitions + } + + filtered := make([]ToolDefinition, 0, len(definitions)) + for _, definition := range definitions { + if definition.Type != ToolTypeApplyPatch { + filtered = append(filtered, definition) + } + } + return filtered +} + func (engine *Engine) run( ctx context.Context, sessionID uuid.UUID, diff --git a/packages/agenty-core/pkg/agentloop/engine_test.go b/packages/agenty-core/pkg/agentloop/engine_test.go index cf801e1..01774ce 100644 --- a/packages/agenty-core/pkg/agentloop/engine_test.go +++ b/packages/agenty-core/pkg/agentloop/engine_test.go @@ -351,8 +351,8 @@ func TestEngineCompletesToolLoopAndPersistsRound(t *testing.T) { if len(requests) != 2 { t.Fatalf("requests = %d, want 2", len(requests)) } - if requests[0].MaxOutputTokens != agentloop.DefaultMaxOutputTokens { - t.Errorf("max output tokens = %d, want %d", requests[0].MaxOutputTokens, agentloop.DefaultMaxOutputTokens) + if requests[0].MaxOutputTokens != 100_000 { + t.Errorf("max output tokens = %d, want %d", requests[0].MaxOutputTokens, 100_000) } if len(requests[0].Messages) != 2 || len(requests[1].Messages) != 4 { t.Fatalf("request message counts = %d, %d", len(requests[0].Messages), len(requests[1].Messages)) @@ -400,6 +400,72 @@ func TestEngineUsesGlobalModelOutputLimit(t *testing.T) { } } +func TestEngineProjectsApplyPatchByProviderCapability(t *testing.T) { + t.Parallel() + + for _, test := range []struct { + name string + freeFormTool bool + wantApplyPatch bool + wantShellPrompt bool + }{ + {name: "free-form provider", freeFormTool: true, wantApplyPatch: true}, + {name: "shell fallback", wantShellPrompt: true}, + } { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + fixture := newExecutionFixture(t, 8_192) + provider, err := fixture.catalog.Get(t.Context(), "openai") + if err != nil { + t.Fatal(err) + } + provider.FreeFormTool = test.freeFormTool + if err := fixture.catalog.Save(t.Context(), provider); err != nil { + t.Fatal(err) + } + if err := fixture.registry.Register(&executionTestTool{ + definition: agentloop.ToolDefinition{ + Type: agentloop.ToolTypeApplyPatch, + Name: "apply_patch", + }, + }); err != nil { + t.Fatal(err) + } + + caller := &scriptedCaller{responses: []*agentloop.Response{{ + Content: conversation.Text("done"), + StopReason: agentloop.StopReasonEndTurn, + }}} + engine := fixture.newEngine(t, func( + context.Context, + catalog.Provider, + catalog.Model, + ) (agentloop.Caller, error) { + return caller, nil + }) + session := fixture.createSession(t) + if _, err := engine.Start(t.Context(), session.ID.String(), conversation.Text("edit")); err != nil { + t.Fatal(err) + } + waitForExecution(t, engine, session.ID) + + requests := caller.Requests() + if len(requests) != 1 { + t.Fatalf("requests = %d, want 1", len(requests)) + } + gotApplyPatch := len(requests[0].Tools) == 1 && requests[0].Tools[0].Name == "apply_patch" + if gotApplyPatch != test.wantApplyPatch { + t.Errorf("apply_patch registered = %v, want %v", gotApplyPatch, test.wantApplyPatch) + } + gotShellPrompt := strings.Contains(requests[0].SystemPrompt, "run the apply_patch command") + if gotShellPrompt != test.wantShellPrompt { + t.Errorf("shell fallback prompt present = %v, want %v", gotShellPrompt, test.wantShellPrompt) + } + }) + } +} + func TestEngineCompactsAutomaticallyAndPreservesTranscript(t *testing.T) { t.Parallel() diff --git a/packages/agenty-core/pkg/domain/agent/agent.go b/packages/agenty-core/pkg/domain/agent/agent.go index dc228d3..6131c3e 100644 --- a/packages/agenty-core/pkg/domain/agent/agent.go +++ b/packages/agenty-core/pkg/domain/agent/agent.go @@ -29,7 +29,11 @@ Sometimes there will be a piece of XML data that follows user's message, which c You will receive this at the very beginning of the session, and maybe more after if something has changed by user or harness. You must follow these messages and treat them as truth. - +{{ if .UseApplyPatchShell }} +The current provider does not support the free-form apply_patch tool. For every file modification, call the shell tool and run the apply_patch command with a complete V4A patch envelope passed through a heredoc on stdin. Do not use cat, sed, printf, or ad hoc scripts to edit files. + + +{{ end }} {{ .Soul }} ` @@ -51,6 +55,10 @@ type Agent struct { UpdatedAt time.Time `json:"updatedAt"` } +type SystemPromptOptions struct { + UseApplyPatchShell bool +} + func New(code, name string) (*Agent, error) { s, err := shared.NewCode(code) if err != nil { @@ -66,10 +74,14 @@ func New(code, name string) (*Agent, error) { }, nil } -func (a *Agent) ResolveSystemPrompt() (string, error) { +func (a *Agent) ResolveSystemPrompt(options SystemPromptOptions) (string, error) { data := struct { - Soul string - }{Soul: a.Soul} + Soul string + UseApplyPatchShell bool + }{ + Soul: a.Soul, + UseApplyPatchShell: options.UseApplyPatchShell, + } var prompt strings.Builder if err := baseSystemPromptTemplate.Execute(&prompt, data); err != nil { diff --git a/packages/agenty-core/pkg/domain/agent/agent_test.go b/packages/agenty-core/pkg/domain/agent/agent_test.go index 2ef7f7e..53d59c1 100644 --- a/packages/agenty-core/pkg/domain/agent/agent_test.go +++ b/packages/agenty-core/pkg/domain/agent/agent_test.go @@ -9,13 +9,19 @@ func TestAgent_ResolveSystemPrompt(t *testing.T) { t.Parallel() tests := []struct { - name string - soul string + name string + soul string + useApplyPatchShell bool }{ {name: "plain text", soul: "You are a coding assistant."}, {name: "multiline text", soul: "Be direct.\nVerify every change."}, {name: "special characters", soul: "Use & report facts."}, {name: "empty", soul: ""}, + { + name: "apply patch shell fallback", + soul: "Edit carefully.", + useApplyPatchShell: true, + }, } for _, tt := range tests { @@ -23,14 +29,26 @@ func TestAgent_ResolveSystemPrompt(t *testing.T) { t.Parallel() agent := &Agent{Soul: tt.soul} - got, err := agent.ResolveSystemPrompt() + got, err := agent.ResolveSystemPrompt(SystemPromptOptions{ + UseApplyPatchShell: tt.useApplyPatchShell, + }) if err != nil { t.Fatalf("ResolveSystemPrompt: %v", err) } - want := strings.Replace(BaseSystemPrompt, "{{ .Soul }}", tt.soul, 1) - if got != want { - t.Errorf("ResolveSystemPrompt() = %q, want %q", got, want) + if !strings.Contains(got, "\n"+tt.soul+"\n") { + t.Errorf("ResolveSystemPrompt() soul = %q", got) + } + gotApplyPatchShell := strings.Contains(got, "run the apply_patch command") + if gotApplyPatchShell != tt.useApplyPatchShell { + t.Errorf( + "ResolveSystemPrompt() apply_patch shell prompt = %v, want %v", + gotApplyPatchShell, + tt.useApplyPatchShell, + ) + } + if strings.Contains(got, "{{") || strings.Contains(got, "}}") { + t.Errorf("ResolveSystemPrompt() contains unresolved template actions: %q", got) } }) } diff --git a/packages/agenty-core/pkg/utils/apply_diff.go b/packages/agenty-core/pkg/utils/apply_diff.go deleted file mode 100644 index b3d186d..0000000 --- a/packages/agenty-core/pkg/utils/apply_diff.go +++ /dev/null @@ -1,393 +0,0 @@ -package utils - -import ( - "fmt" - "strings" - "unicode" -) - -type ApplyDiffMode uint8 - -const ( - ApplyDiffDefault ApplyDiffMode = iota - ApplyDiffCreate -) - -type applyDiffChunk struct { - originalIndex int - deletedLines []string - insertedLines []string -} - -type applyDiffParser struct { - lines []string - index int - fuzz int -} - -const ( - applyDiffEndPatch = "*** End Patch" - applyDiffEndFile = "*** End of File" -) - -var applyDiffSectionMarkers = []string{ - applyDiffEndPatch, - "*** Update File:", - "*** Delete File:", - "*** Add File:", - applyDiffEndFile, -} - -var applyDiffSectionTerminators = []string{ - applyDiffEndPatch, - "*** Update File:", - "*** Delete File:", - "*** Add File:", -} - -// ApplyDiff applies a headerless V4A diff using the OpenAI Agents SDK semantics. -func ApplyDiff(input, diff string, mode ApplyDiffMode) (string, error) { - diffLines := normalizeApplyDiffLines(diff) - switch mode { - case ApplyDiffCreate: - return parseCreateDiff(diffLines) - case ApplyDiffDefault: - default: - return "", fmt.Errorf("apply diff: unsupported mode %d", mode) - } - - chunks, err := parseUpdateDiff(diffLines, input) - if err != nil { - return "", err - } - return applyDiffChunks(input, chunks) -} - -func normalizeApplyDiffLines(diff string) []string { - lines := strings.Split(strings.ReplaceAll(diff, "\r\n", "\n"), "\n") - for index := range lines { - lines[index] = strings.TrimSuffix(lines[index], "\r") - } - if len(lines) > 0 && lines[len(lines)-1] == "" { - lines = lines[:len(lines)-1] - } - return lines -} - -func parseCreateDiff(lines []string) (string, error) { - parser := applyDiffParser{ - lines: append(append([]string{}, lines...), applyDiffEndPatch), - } - output := make([]string, 0, len(lines)) - for !parser.done(applyDiffSectionTerminators) { - line := parser.lines[parser.index] - parser.index++ - if !strings.HasPrefix(line, "+") { - return "", fmt.Errorf("invalid add file line: %s", line) - } - output = append(output, strings.TrimPrefix(line, "+")) - } - return strings.Join(output, "\n"), nil -} - -func parseUpdateDiff(lines []string, input string) ([]applyDiffChunk, error) { - parser := applyDiffParser{ - lines: append(append([]string{}, lines...), applyDiffEndPatch), - } - inputLines := strings.Split(input, "\n") - chunks := make([]applyDiffChunk, 0) - cursor := 0 - - for !parser.done(applyDiffSectionMarkers) { - anchors, anchorCount := parser.readAnchors() - if anchorCount == 0 && cursor != 0 { - return nil, fmt.Errorf("invalid line:\n%s", parser.lines[parser.index]) - } - - requireAnchorMatch := anchorCount > 1 - for index, anchor := range anchors { - var err error - cursor, err = parser.advanceCursorToAnchor( - anchor, - inputLines, - cursor, - requireAnchorMatch, - index > 0, - ) - if err != nil { - return nil, err - } - } - - context, sectionChunks, endIndex, eof, err := readApplyDiffSection(parser.lines, parser.index) - if err != nil { - return nil, err - } - newIndex, fuzz := findApplyDiffContext(inputLines, context, cursor, eof) - if newIndex == -1 { - contextText := strings.Join(context, "\n") - if eof { - return nil, fmt.Errorf("invalid EOF context %d:\n%s", cursor, contextText) - } - return nil, fmt.Errorf("invalid context %d:\n%s", cursor, contextText) - } - - parser.fuzz += fuzz - for _, chunk := range sectionChunks { - chunk.originalIndex += newIndex - chunks = append(chunks, chunk) - } - cursor = newIndex + len(context) - parser.index = endIndex - } - - return chunks, nil -} - -func (parser *applyDiffParser) done(prefixes []string) bool { - if parser.index >= len(parser.lines) { - return true - } - for _, prefix := range prefixes { - if strings.HasPrefix(parser.lines[parser.index], prefix) { - return true - } - } - return false -} - -func (parser *applyDiffParser) readAnchors() ([]string, int) { - anchors := make([]string, 0) - anchorCount := 0 - for { - line := parser.lines[parser.index] - switch { - case strings.HasPrefix(line, "@@ "): - parser.index++ - anchorCount++ - anchor := strings.TrimPrefix(line, "@@ ") - if strings.TrimSpace(anchor) != "" { - anchors = append(anchors, anchor) - } - case line == "@@": - parser.index++ - anchorCount++ - default: - return anchors, anchorCount - } - } -} - -func (parser *applyDiffParser) advanceCursorToAnchor( - anchor string, - inputLines []string, - cursor int, - requireMatch bool, - forceForwardSearch bool, -) (int, error) { - found := false - if !forceForwardSearch && containsApplyDiffLine(inputLines[:cursor], anchor, false) { - found = true - } else if index := findApplyDiffLine(inputLines, anchor, cursor, false); index >= 0 { - cursor = index + 1 - found = true - } - - if !found { - if !forceForwardSearch && containsApplyDiffLine(inputLines[:cursor], anchor, true) { - found = true - } else if index := findApplyDiffLine(inputLines, anchor, cursor, true); index >= 0 { - cursor = index + 1 - parser.fuzz++ - found = true - } - } - - if requireMatch && !found { - return 0, fmt.Errorf("invalid anchor %d:\n%s", cursor, anchor) - } - return cursor, nil -} - -func containsApplyDiffLine(lines []string, target string, trimmed bool) bool { - return findApplyDiffLine(lines, target, 0, trimmed) >= 0 -} - -func findApplyDiffLine(lines []string, target string, start int, trimmed bool) int { - for index := start; index < len(lines); index++ { - line := lines[index] - if trimmed { - line = strings.TrimSpace(line) - target = strings.TrimSpace(target) - } - if line == target { - return index - } - } - return -1 -} - -func readApplyDiffSection( - lines []string, - startIndex int, -) ([]string, []applyDiffChunk, int, bool, error) { - context := make([]string, 0) - deletedLines := make([]string, 0) - insertedLines := make([]string, 0) - chunks := make([]applyDiffChunk, 0) - mode := byte(' ') - index := startIndex - - flushChunk := func() { - if len(insertedLines) == 0 && len(deletedLines) == 0 { - return - } - chunks = append(chunks, applyDiffChunk{ - originalIndex: len(context) - len(deletedLines), - deletedLines: deletedLines, - insertedLines: insertedLines, - }) - deletedLines = make([]string, 0) - insertedLines = make([]string, 0) - } - - for index < len(lines) { - raw := lines[index] - if strings.HasPrefix(raw, "@@") || isApplyDiffSectionEnd(raw) { - break - } - if raw == "***" { - break - } - if strings.HasPrefix(raw, "***") { - return nil, nil, 0, false, fmt.Errorf("invalid line: %s", raw) - } - - index++ - previousMode := mode - line := raw - if line == "" { - line = " " - } - switch line[0] { - case '+', '-', ' ': - mode = line[0] - default: - return nil, nil, 0, false, fmt.Errorf("invalid line: %s", line) - } - line = line[1:] - - if mode == ' ' && previousMode != mode { - flushChunk() - } - switch mode { - case '-': - deletedLines = append(deletedLines, line) - context = append(context, line) - case '+': - insertedLines = append(insertedLines, line) - case ' ': - context = append(context, line) - } - } - flushChunk() - - if index < len(lines) && lines[index] == applyDiffEndFile { - return context, chunks, index + 1, true, nil - } - if index == startIndex { - return nil, nil, 0, false, fmt.Errorf("nothing in section at index %d: %s", index, lines[index]) - } - return context, chunks, index, false, nil -} - -func isApplyDiffSectionEnd(line string) bool { - for _, marker := range applyDiffSectionMarkers { - if strings.HasPrefix(line, marker) { - return true - } - } - return false -} - -func findApplyDiffContext(lines, context []string, start int, eof bool) (int, int) { - if eof { - endStart := max(0, len(lines)-len(context)) - if index, fuzz := findApplyDiffContextCore(lines, context, endStart); index != -1 { - return index, fuzz - } - index, fuzz := findApplyDiffContextCore(lines, context, start) - return index, fuzz + 10_000 - } - return findApplyDiffContextCore(lines, context, start) -} - -func findApplyDiffContextCore(lines, context []string, start int) (int, int) { - if len(context) == 0 { - return start, 0 - } - - comparisons := []struct { - fuzz int - mapf func(string) string - }{ - {fuzz: 0, mapf: func(value string) string { return value }}, - {fuzz: 1, mapf: func(value string) string { - return strings.TrimRightFunc(value, unicode.IsSpace) - }}, - {fuzz: 100, mapf: strings.TrimSpace}, - } - for _, comparison := range comparisons { - for index := start; index < len(lines); index++ { - if equalApplyDiffSlice(lines, context, index, comparison.mapf) { - return index, comparison.fuzz - } - } - } - return -1, 0 -} - -func equalApplyDiffSlice( - source []string, - target []string, - start int, - mapf func(string) string, -) bool { - if start+len(target) > len(source) { - return false - } - for index := range target { - if mapf(source[start+index]) != mapf(target[index]) { - return false - } - } - return true -} - -func applyDiffChunks(input string, chunks []applyDiffChunk) (string, error) { - originalLines := strings.Split(input, "\n") - destinationLines := make([]string, 0, len(originalLines)) - originalIndex := 0 - for _, chunk := range chunks { - if chunk.originalIndex > len(originalLines) { - return "", fmt.Errorf( - "applyDiff: chunk original index %d exceeds input length %d", - chunk.originalIndex, - len(originalLines), - ) - } - if originalIndex > chunk.originalIndex { - return "", fmt.Errorf( - "applyDiff: overlapping chunk at %d with cursor %d", - chunk.originalIndex, - originalIndex, - ) - } - - destinationLines = append(destinationLines, originalLines[originalIndex:chunk.originalIndex]...) - destinationLines = append(destinationLines, chunk.insertedLines...) - originalIndex = chunk.originalIndex + len(chunk.deletedLines) - } - destinationLines = append(destinationLines, originalLines[originalIndex:]...) - return strings.Join(destinationLines, "\n"), nil -} diff --git a/packages/agenty-core/pkg/utils/apply_diff_test.go b/packages/agenty-core/pkg/utils/apply_diff_test.go deleted file mode 100644 index ef3c447..0000000 --- a/packages/agenty-core/pkg/utils/apply_diff_test.go +++ /dev/null @@ -1,258 +0,0 @@ -package utils - -import ( - "strings" - "testing" -) - -func TestApplyDiff(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - input string - diff string - mode ApplyDiffMode - want string - }{ - { - name: "create file with blank line", - diff: "+hello\n+world\n+", - mode: ApplyDiffCreate, - want: "hello\nworld\n", - }, - { - name: "create empty file", - mode: ApplyDiffCreate, - want: "", - }, - { - name: "create file normalizes CRLF", - diff: "+hello\r\n+\r\n+world\r\n", - mode: ApplyDiffCreate, - want: "hello\n\nworld", - }, - { - name: "empty diff preserves existing file", - input: "one\ntwo\n", - want: "one\ntwo\n", - }, - { - name: "floating insertion into empty file", - input: "", - diff: "@@\n+hello\n+world", - want: "hello\nworld\n", - }, - { - name: "floating hunk", - input: "- Milk\n- Bread\n- Eggs\n- Apples\n- Coffee", - diff: "@@\n - Milk\n - Bread\n - Eggs\n-- Apples\n-- Coffee\n+- [x] Apples\n+- [x] Coffee", - want: "- Milk\n- Bread\n- Eggs\n- [x] Apples\n- [x] Coffee", - }, - { - name: "anchored replacement preserves trailing newline", - input: "line1\nline2\nline3\n", - diff: "@@ line1\n-line2\n+updated\n line3", - want: "line1\nupdated\nline3\n", - }, - { - name: "deletion with context", - input: "keep\nremove me\nstay\n", - diff: "@@ keep\n-remove me\n stay", - want: "keep\nstay\n", - }, - { - name: "pure insertion with blank context lines", - input: "import os\n\ndef main():\n return 1\n", - diff: " import os\n+import sys\n\n def main():\n return 1", - want: "import os\nimport sys\n\ndef main():\n return 1\n", - }, - { - name: "multiple anchored sections", - input: "class Foo:\n def baz(self):\n return 1\n\ndef main():\n print(Foo().baz())\n", - diff: "@@ class Foo:\n- def baz(self):\n+ def value(self):\n return 1\n@@ def main():\n- print(Foo().baz())\n+ print(Foo().value())", - want: "class Foo:\n def value(self):\n return 1\n\ndef main():\n print(Foo().value())\n", - }, - { - name: "stacked anchors", - input: "class First\n def target():\n pass\n\nclass Second\n def target():\n pass\n", - diff: "@@ class Second\n@@ def target():\n- pass\n+ return 1", - want: "class First\n def target():\n pass\n\nclass Second\n def target():\n return 1\n", - }, - { - name: "reuses parent anchor in a later hunk", - input: "class Target\n def first():\n pass\n\n def second():\n pass\n", - diff: "@@ class Target\n@@ def first():\n- pass\n+ return 1\n@@ class Target\n@@ def second():\n- pass\n+ return 2", - want: "class Target\n def first():\n return 1\n\n def second():\n return 2\n", - }, - { - name: "reuses trimmed parent anchor in a later hunk", - input: " class Target \n def first():\n pass\n\n def second():\n pass\n", - diff: "@@ class Target\n@@ def first():\n- pass\n+ return 1\n@@ class Target\n@@ def second():\n- pass\n+ return 2", - want: " class Target \n def first():\n return 1\n\n def second():\n return 2\n", - }, - { - name: "single missing anchor remains best effort", - input: "one\ntwo\n", - diff: "@@ missing\n-one\n+first", - want: "first\ntwo\n", - }, - { - name: "trailing bare anchor", - input: "class Only\n def run():\n pass\n", - diff: "@@ class Only\n@@\n- pass\n+ return 1", - want: "class Only\n def run():\n return 1\n", - }, - { - name: "end of file", - input: "Line A\nLine B\nLine C", - diff: "@@\n Line B\n-Line C\n+Line C updated\n*** End of File", - want: "Line A\nLine B\nLine C updated", - }, - { - name: "trailing whitespace fuzz", - input: "one \ntwo\n", - diff: " one\n-two\n+second", - want: "one \nsecond\n", - }, - { - name: "leading and trailing whitespace fuzz", - input: " target \nnext\n", - diff: " target\n-next\n+done", - want: " target \ndone\n", - }, - { - name: "traditional line marker is a best effort anchor", - input: "one\ntwo\n", - diff: "@@ -1,2 +1,2 @@\n one\n-two\n+2", - want: "one\n2\n", - }, - { - name: "update diff normalizes CRLF", - input: "one\ntwo\n", - diff: " one\r\n-two\r\n+second\r\n", - want: "one\nsecond\n", - }, - { - name: "end of file marker falls back to earlier context", - input: "target\nmiddle\nend", - diff: " target\n+after\n*** End of File", - want: "target\nafter\nmiddle\nend", - }, - { - name: "context-only diff leaves content unchanged", - input: "legacy content", - diff: " legacy content", - want: "legacy content", - }, - { - name: "replacement works without hunk marker", - input: "before\nkeep", - diff: "-before\n+after", - want: "after\nkeep", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - got, err := ApplyDiff(test.input, test.diff, test.mode) - if err != nil { - t.Fatal(err) - } - if got != test.want { - t.Errorf("ApplyDiff() = %q, want %q", got, test.want) - } - }) - } -} - -func TestApplyDiffRejectsInvalidInput(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - input string - diff string - mode ApplyDiffMode - wantErr string - }{ - { - name: "unsupported mode", - mode: ApplyDiffMode(255), - wantErr: "unsupported mode 255", - }, - { - name: "create line without plus", - diff: "+valid\ninvalid", - mode: ApplyDiffCreate, - wantErr: "invalid add file line", - }, - { - name: "missing context", - input: "one\ntwo\n", - diff: " missing\n-two\n+second", - wantErr: "invalid context", - }, - { - name: "missing first stacked anchor", - input: "class Wrong\n def desired():\n pass\n", - diff: "@@ class Target\n@@ def desired():\n- pass\n+ return 1", - wantErr: "invalid anchor", - }, - { - name: "missing second stacked anchor", - input: "class Target\n def desired():\n pass\n", - diff: "@@ class Target\n@@ def missing():\n- pass\n+ return 1", - wantErr: "invalid anchor", - }, - { - name: "missing anchor followed by bare marker", - input: "one\ntwo\n", - diff: "@@ missing\n@@\n-two\n+second", - wantErr: "invalid anchor", - }, - { - name: "invalid unprefixed update line", - input: "one\n", - diff: "one", - wantErr: "invalid line", - }, - { - name: "unknown patch directive", - input: "one\n", - diff: "*** Unknown Directive", - wantErr: "invalid line", - }, - { - name: "empty section", - input: "one\n", - diff: "@@", - wantErr: "nothing in section", - }, - { - name: "invalid EOF context", - input: "one\ntwo", - diff: " missing\n*** End of File", - wantErr: "invalid EOF context", - }, - { - name: "content after EOF marker needs another anchor", - input: "one", - diff: " one\n*** End of File\n two", - wantErr: "invalid line", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := ApplyDiff(test.input, test.diff, test.mode) - if err == nil || !strings.Contains(err.Error(), test.wantErr) { - t.Fatalf("ApplyDiff() error = %v, want containing %q", err, test.wantErr) - } - }) - } -} diff --git a/packages/agenty-core/test/e2e/execution_test.go b/packages/agenty-core/test/e2e/execution_test.go index 7d06679..ccefd5f 100644 --- a/packages/agenty-core/test/e2e/execution_test.go +++ b/packages/agenty-core/test/e2e/execution_test.go @@ -200,9 +200,7 @@ func TestAgentLoopExecutesThroughEveryProviderProtocol(t *testing.T) { tt.apiType, ) } - wantTools := []string{ - "delete_file", "glob", "grep", "ls", "patch_file", "read_file", "shell", "write_file", - } + wantTools := []string{"glob", "grep", "ls", "read_file", "shell"} if tt.apiType == "openai" { wantTools = []string{"apply_patch", "glob", "grep", "ls", "read_file", "shell"} } diff --git a/packages/agenty-core/test/e2e/test_helpers_test.go b/packages/agenty-core/test/e2e/test_helpers_test.go index bbdf9f0..4498cd3 100644 --- a/packages/agenty-core/test/e2e/test_helpers_test.go +++ b/packages/agenty-core/test/e2e/test_helpers_test.go @@ -66,11 +66,12 @@ func createExecutionResources( return Session{}, fmt.Errorf("create agent: %w", err) } if _, err := client.CreateProvider(ctx, ProviderCreateInput{ - Code: providerCode, - Name: "E2E Provider", - Type: apiType, - BaseURL: fixture.BaseURL(apiType), - APIKey: "test-key", + Code: providerCode, + Name: "E2E Provider", + Type: apiType, + BaseURL: fixture.BaseURL(apiType), + APIKey: "test-key", + FreeFormTool: apiType == "openai", }); err != nil { return Session{}, fmt.Errorf("create provider: %w", err) } diff --git a/packages/patch-applier/.gitignore b/packages/patch-applier/.gitignore new file mode 100644 index 0000000..b7fc3ae --- /dev/null +++ b/packages/patch-applier/.gitignore @@ -0,0 +1,2 @@ +target/ +.turbo/ diff --git a/packages/patch-applier/Cargo.lock b/packages/patch-applier/Cargo.lock new file mode 100644 index 0000000..22061e7 --- /dev/null +++ b/packages/patch-applier/Cargo.lock @@ -0,0 +1,114 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "memchr" +version = "2.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" + +[[package]] +name = "patch-applier" +version = "0.1.0" +dependencies = [ + "serde", + "serde_json", + "similar", +] + +[[package]] +name = "proc-macro2" +version = "1.0.107" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "serde" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67dca2c9c51e58a4791a4b1ed58308b39c64224d349a935ab5039aa360942a48" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.151" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "similar" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" + +[[package]] +name = "syn" +version = "3.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6275cddf4610d1775e6d1fe9469b2e77d0f39fd98fb7450901b821e0c53649f" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "zmij" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" diff --git a/packages/patch-applier/Cargo.toml b/packages/patch-applier/Cargo.toml new file mode 100644 index 0000000..6cd75c4 --- /dev/null +++ b/packages/patch-applier/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "patch-applier" +version = "0.1.0" +edition = "2021" +license = "Apache-2.0" +description = "Atomic V4A patch applier for Agenty" + +[[bin]] +name = "apply_patch" +path = "src/main.rs" + +[dependencies] +serde = { version = "1", features = ["derive"] } +serde_json = "1" +similar = "2" + +[profile.release] +lto = true +strip = true diff --git a/packages/patch-applier/README.md b/packages/patch-applier/README.md new file mode 100644 index 0000000..6d711a1 --- /dev/null +++ b/packages/patch-applier/README.md @@ -0,0 +1,32 @@ +# patch-applier + +`patch-applier` builds the `apply_patch` executable bundled with Agenty. It reads one +complete V4A patch envelope from stdin and resolves relative paths from its working +directory. + +Before writing, it parses every operation, groups operations by logical file in first +appearance order, applies each group's operations in source order to an in-memory +snapshot, and rejects incompatible state transitions or path ownership conflicts. The +commit phase stages all new contents before replacing or deleting targets and rolls back +completed replacements if a later filesystem operation fails. + +Success and failure both write one JSON object to stdout. A successful result includes +the final unified diff and added/removed line counts for every changed path: + +```json +{ + "success": true, + "cwd": "/workspace", + "files": [ + { + "path": "src/main.rs", + "diff": "--- a/src/main.rs\n+++ b/src/main.rs\n...", + "addedLines": 2, + "removedLines": 1 + } + ] +} +``` + +Run `cargo test`, `cargo fmt --all -- --check`, and +`cargo clippy --all-targets --all-features -- -D warnings` before release. diff --git a/packages/patch-applier/package.json b/packages/patch-applier/package.json new file mode 100644 index 0000000..55f6f57 --- /dev/null +++ b/packages/patch-applier/package.json @@ -0,0 +1,11 @@ +{ + "name": "patch-applier", + "version": "0.1.0", + "private": true, + "scripts": { + "build": "cargo build --release", + "test": "cargo test", + "lint": "cargo fmt --all -- --check && cargo clippy --all-targets --all-features -- -D warnings", + "clean": "cargo clean" + } +} diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs new file mode 100644 index 0000000..5bf1f10 --- /dev/null +++ b/packages/patch-applier/src/lib.rs @@ -0,0 +1,1297 @@ +use std::collections::{HashMap, HashSet}; +use std::fmt::{self, Display, Formatter}; +use std::fs::{self, File, OpenOptions}; +use std::io::{self, Write}; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; + +use serde::Serialize; +use similar::{ChangeTag, TextDiff}; + +const BEGIN_MARKER: &str = "*** Begin Patch"; +const END_MARKER: &str = "*** End Patch"; +const UPDATE_MARKER: &str = "*** Update File:"; +const DELETE_MARKER: &str = "*** Delete File:"; +const ADD_MARKER: &str = "*** Add File:"; +const MOVE_MARKER: &str = "*** Move to:"; +const END_OF_FILE_MARKER: &str = "*** End of File"; + +static TRANSACTION_COUNTER: AtomicU64 = AtomicU64::new(0); + +#[derive(Debug)] +pub enum PatchError { + Io(io::Error), + Invalid(String), + Conflict(String), +} + +impl Display for PatchError { + fn fmt(&self, formatter: &mut Formatter<'_>) -> fmt::Result { + match self { + Self::Io(error) => write!(formatter, "I/O error: {error}"), + Self::Invalid(message) => write!(formatter, "invalid patch: {message}"), + Self::Conflict(message) => write!(formatter, "conflicting patch operations: {message}"), + } + } +} + +impl std::error::Error for PatchError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::Io(error) => Some(error), + Self::Invalid(_) | Self::Conflict(_) => None, + } + } +} + +impl From for PatchError { + fn from(error: io::Error) -> Self { + Self::Io(error) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum OperationKind { + Add, + Update, + Delete, +} + +#[derive(Debug, Clone)] +struct Operation { + kind: OperationKind, + path: PathBuf, + move_to: Option, + diff: String, + line: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct FileSnapshot { + kind: EntryKind, + data: Vec, + mode: u32, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum EntryKind { + Missing, + Regular, + Symlink, + Directory, + Other, +} + +#[derive(Debug, Clone)] +struct VirtualFile { + kind: EntryKind, + data: Vec, + mode: u32, +} + +impl VirtualFile { + fn missing() -> Self { + Self { + kind: EntryKind::Missing, + data: Vec::new(), + mode: 0, + } + } + + fn regular(data: Vec, mode: u32) -> Self { + Self { + kind: EntryKind::Regular, + data, + mode, + } + } +} + +#[derive(Debug, Clone)] +struct FileGroup { + first_index: usize, + operations: Vec, +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct PatchResult { + pub success: bool, + pub cwd: String, + pub files: Vec, +} + +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct FileResult { + pub path: String, + pub diff: String, + pub added_lines: usize, + pub removed_lines: usize, +} + +pub fn apply_patch(cwd: &Path, patch: &str) -> Result { + let operations = parse_envelope(cwd, patch)?; + let groups = classify_operations(operations)?; + let mut transaction = Transaction::new(cwd.to_path_buf()); + + for group in groups { + for operation in group.operations { + transaction.apply(operation)?; + } + } + + let results = transaction.prepare_results()?; + transaction.commit()?; + + Ok(PatchResult { + success: true, + cwd: cwd.display().to_string(), + files: results, + }) +} + +fn parse_envelope(cwd: &Path, patch: &str) -> Result, PatchError> { + let lines = normalized_lines(patch); + if lines.len() < 3 || lines.first().map(String::as_str) != Some(BEGIN_MARKER) { + return Err(PatchError::Invalid(format!( + "patch must start with {BEGIN_MARKER:?}" + ))); + } + if lines.last().map(String::as_str) != Some(END_MARKER) { + return Err(PatchError::Invalid(format!( + "patch must end with {END_MARKER:?}" + ))); + } + + let mut operations = Vec::new(); + let mut index = 1; + while index < lines.len() - 1 { + let (operation, next_index) = parse_operation(cwd, &lines, index)?; + operations.push(operation); + index = next_index; + } + if operations.is_empty() { + return Err(PatchError::Invalid( + "patch contains no file operations".to_string(), + )); + } + Ok(operations) +} + +fn normalized_lines(patch: &str) -> Vec { + let normalized = patch.replace("\r\n", "\n"); + let mut lines: Vec = normalized + .split('\n') + .map(|line| line.strip_suffix('\r').unwrap_or(line).to_string()) + .collect(); + if lines.last().is_some_and(String::is_empty) { + lines.pop(); + } + lines +} + +fn parse_operation( + cwd: &Path, + lines: &[String], + start: usize, +) -> Result<(Operation, usize), PatchError> { + let header = &lines[start]; + let (kind, marker) = if header.starts_with(UPDATE_MARKER) { + (OperationKind::Update, UPDATE_MARKER) + } else if header.starts_with(DELETE_MARKER) { + (OperationKind::Delete, DELETE_MARKER) + } else if header.starts_with(ADD_MARKER) { + (OperationKind::Add, ADD_MARKER) + } else { + return Err(PatchError::Invalid(format!( + "invalid patch header at line {}: {header}", + start + 1 + ))); + }; + + let raw_path = header.strip_prefix(marker).unwrap_or_default().trim(); + if raw_path.is_empty() { + return Err(PatchError::Invalid(format!( + "operation at line {} has an empty path", + start + 1 + ))); + } + let path = resolve_path(cwd, raw_path)?; + + let mut index = start + 1; + let mut move_to = None; + if kind == OperationKind::Update + && index < lines.len() - 1 + && lines[index].starts_with(MOVE_MARKER) + { + let raw_move_to = lines[index] + .strip_prefix(MOVE_MARKER) + .unwrap_or_default() + .trim(); + if raw_move_to.is_empty() { + return Err(PatchError::Invalid(format!( + "move at line {} has an empty path", + index + 1 + ))); + } + move_to = Some(resolve_path(cwd, raw_move_to)?); + index += 1; + } + + let body_start = index; + while index < lines.len() - 1 && !is_operation_header(&lines[index]) { + index += 1; + } + let body = lines[body_start..index].join("\n"); + if kind == OperationKind::Delete && !body.is_empty() { + return Err(PatchError::Invalid(format!( + "delete operation for {raw_path:?} must not contain a diff body" + ))); + } + + Ok(( + Operation { + kind, + path, + move_to, + diff: body, + line: start + 1, + }, + index, + )) +} + +fn is_operation_header(line: &str) -> bool { + line.starts_with(UPDATE_MARKER) + || line.starts_with(DELETE_MARKER) + || line.starts_with(ADD_MARKER) +} + +fn resolve_path(cwd: &Path, raw: &str) -> Result { + let path = Path::new(raw); + let resolved = if path.is_absolute() { + path.to_path_buf() + } else { + cwd.join(path) + }; + lexical_normalize(&resolved) +} + +fn lexical_normalize(path: &Path) -> Result { + let mut normalized = PathBuf::new(); + for component in path.components() { + match component { + std::path::Component::Prefix(prefix) => { + normalized.push(prefix.as_os_str()); + } + std::path::Component::RootDir => { + normalized.push(Path::new(std::path::MAIN_SEPARATOR_STR)); + } + std::path::Component::CurDir => {} + std::path::Component::ParentDir => { + let rooted = normalized.has_root(); + if !normalized.pop() || (rooted && !normalized.has_root()) { + return Err(PatchError::Invalid(format!( + "path escapes its root: {}", + path.display() + ))); + } + } + std::path::Component::Normal(component) => normalized.push(component), + } + } + Ok(normalized) +} + +fn classify_operations(operations: Vec) -> Result, PatchError> { + let mut groups = Vec::::new(); + let mut aliases = HashMap::::new(); + + for (index, operation) in operations.into_iter().enumerate() { + let group_index = if let Some(group_index) = aliases.get(&operation.path) { + *group_index + } else { + groups.push(FileGroup { + first_index: index, + operations: Vec::new(), + }); + groups.len() - 1 + }; + + aliases.insert(operation.path.clone(), group_index); + if let Some(destination) = &operation.move_to { + if let Some(existing) = aliases.get(destination) { + if *existing != group_index { + return Err(PatchError::Conflict(format!( + "operation at line {} moves {} onto a path owned by another file", + operation.line, + destination.display() + ))); + } + } + aliases.insert(destination.clone(), group_index); + } + groups[group_index].operations.push(operation); + } + + groups.sort_by_key(|group| group.first_index); + Ok(groups) +} + +struct Transaction { + cwd: PathBuf, + original: HashMap, + state: HashMap, + touched_order: Vec, + deleted_paths: HashSet, +} + +impl Transaction { + fn new(cwd: PathBuf) -> Self { + Self { + cwd, + original: HashMap::new(), + state: HashMap::new(), + touched_order: Vec::new(), + deleted_paths: HashSet::new(), + } + } + + fn apply(&mut self, operation: Operation) -> Result<(), PatchError> { + self.touch(&operation.path)?; + let current = self.current(&operation.path)?.clone(); + match operation.kind { + OperationKind::Add => { + if current.kind != EntryKind::Missing + || self.deleted_paths.contains(&operation.path) + { + return Err(PatchError::Conflict(format!( + "operation at line {} creates an existing or previously deleted file {}", + operation.line, + operation.path.display() + ))); + } + let data = parse_create_diff(&operation.diff)?; + self.state + .insert(operation.path.clone(), VirtualFile::regular(data, 0o644)); + } + OperationKind::Update => { + if current.kind != EntryKind::Regular { + return Err(PatchError::Conflict(format!( + "operation at line {} updates a non-existent or non-regular file {}", + operation.line, + operation.path.display() + ))); + } + let data = apply_update_diff(¤t.data, &operation.diff)?; + self.state.insert( + operation.path.clone(), + VirtualFile::regular(data, current.mode), + ); + if let Some(destination) = operation.move_to { + self.move_file(&operation.path, &destination, operation.line)?; + } + } + OperationKind::Delete => { + if current.kind != EntryKind::Regular { + return Err(PatchError::Conflict(format!( + "operation at line {} deletes a missing or non-regular file {}", + operation.line, + operation.path.display() + ))); + } + self.state + .insert(operation.path.clone(), VirtualFile::missing()); + self.deleted_paths.insert(operation.path); + } + } + Ok(()) + } + + fn touch(&mut self, path: &Path) -> Result<(), PatchError> { + if self.original.contains_key(path) { + return Ok(()); + } + let snapshot = read_snapshot(path)?; + let state = VirtualFile { + kind: snapshot.kind, + data: snapshot.data.clone(), + mode: snapshot.mode, + }; + self.original.insert(path.to_path_buf(), snapshot); + self.state.insert(path.to_path_buf(), state); + self.touched_order.push(path.to_path_buf()); + Ok(()) + } + + fn current(&mut self, path: &Path) -> Result<&VirtualFile, PatchError> { + self.touch(path)?; + self.state + .get(path) + .ok_or_else(|| PatchError::Invalid(format!("missing state for {}", path.display()))) + } + + fn move_file( + &mut self, + source: &Path, + destination: &Path, + line: usize, + ) -> Result<(), PatchError> { + if source == destination { + return Ok(()); + } + self.touch(destination)?; + let destination_state = self.current(destination)?.clone(); + if destination_state.kind != EntryKind::Missing || self.deleted_paths.contains(destination) + { + return Err(PatchError::Conflict(format!( + "operation at line {line} moves {} onto an occupied path {}", + source.display(), + destination.display() + ))); + } + let source_state = self.current(source)?.clone(); + self.state.insert(destination.to_path_buf(), source_state); + self.state + .insert(source.to_path_buf(), VirtualFile::missing()); + self.deleted_paths.insert(source.to_path_buf()); + Ok(()) + } + + fn prepare_results(&self) -> Result, PatchError> { + let mut results = Vec::new(); + for path in &self.touched_order { + let before = self.original.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing original snapshot for {}", path.display())) + })?; + let after = self.state.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing final state for {}", path.display())) + })?; + if before.kind == after.kind && before.data == after.data { + continue; + } + let old_text = text_for_diff(before)?; + let new_text = text_for_diff_virtual(after)?; + let relative = relative_display(&self.cwd, path); + let old_header = if before.kind == EntryKind::Missing { + "/dev/null".to_string() + } else { + format!("a/{relative}") + }; + let new_header = if after.kind == EntryKind::Missing { + "/dev/null".to_string() + } else { + format!("b/{relative}") + }; + let (diff, added_lines, removed_lines) = + unified_diff(&old_text, &new_text, &old_header, &new_header); + results.push(FileResult { + path: relative, + diff, + added_lines, + removed_lines, + }); + } + Ok(results) + } + + fn commit(&self) -> Result<(), PatchError> { + let changes = self.changes()?; + if changes.is_empty() { + return Ok(()); + } + + for change in &changes { + let expected = self.original.get(&change.path).ok_or_else(|| { + PatchError::Invalid(format!( + "missing original snapshot for {}", + change.path.display() + )) + })?; + if &read_snapshot(&change.path)? != expected { + return Err(PatchError::Conflict(format!( + "file changed while the patch was being prepared: {}", + change.path.display() + ))); + } + } + + let transaction_id = format!( + ".agenty-apply-patch-{}-{}", + std::process::id(), + TRANSACTION_COUNTER.fetch_add(1, Ordering::Relaxed) + ); + let mut staged = Vec::new(); + let mut backups = Vec::new(); + let mut created_dirs = Vec::new(); + + let result = (|| -> Result<(), PatchError> { + for change in &changes { + if let Some(file) = &change.after { + let parent = change.path.parent().unwrap_or(&self.cwd); + created_dirs.extend(create_missing_dirs(parent)?); + let temp = parent.join(format!(".{transaction_id}.stage")); + let temp = unique_path(&temp)?; + let mut output = OpenOptions::new() + .write(true) + .create_new(true) + .open(&temp)?; + output.write_all(&file.data)?; + output.flush()?; + set_mode(&output, file.mode)?; + output.sync_all()?; + staged.push((change.path.clone(), temp)); + } + } + + for change in &changes { + if path_exists(&change.path)? { + let backup = backup_path(&change.path, &transaction_id)?; + fs::rename(&change.path, &backup)?; + backups.push((change.path.clone(), backup)); + } + } + + for (path, temp) in &staged { + fs::rename(temp, path)?; + } + sync_parent_directories(&changes)?; + Ok(()) + })(); + + if result.is_err() { + for (path, _) in staged.iter().rev() { + let _ = fs::remove_file(path); + } + for (path, backup) in backups.iter().rev() { + let _ = fs::rename(backup, path); + } + for temp in staged.iter().map(|(_, temp)| temp) { + let _ = fs::remove_file(temp); + } + for directory in &created_dirs { + let _ = fs::remove_dir(directory); + } + let _ = sync_parent_directories(&changes); + } else { + for (_, backup) in backups { + let _ = fs::remove_file(backup); + } + } + result + } + + fn changes(&self) -> Result, PatchError> { + let mut changes = Vec::new(); + for path in &self.touched_order { + let before = self.original.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing original snapshot for {}", path.display())) + })?; + let after = self.state.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing final state for {}", path.display())) + })?; + let before_exists = before.kind != EntryKind::Missing; + let after_file = if after.kind == EntryKind::Regular { + Some(after.clone()) + } else { + None + }; + if before_exists && after.kind == EntryKind::Missing { + changes.push(Change { + path: path.clone(), + after: None, + }); + } else if (!before_exists && after_file.is_some()) + || (before.kind == EntryKind::Regular + && after.kind == EntryKind::Regular + && (before.data != after.data || before.mode != after.mode)) + { + changes.push(Change { + path: path.clone(), + after: after_file, + }); + } else if before.kind != after.kind { + return Err(PatchError::Conflict(format!( + "unsupported final state transition for {}", + path.display() + ))); + } + } + Ok(changes) + } +} + +struct Change { + path: PathBuf, + after: Option, +} + +#[cfg(unix)] +fn sync_parent_directories(changes: &[Change]) -> Result<(), PatchError> { + let mut parents = HashSet::new(); + for change in changes { + if let Some(parent) = change.path.parent() { + parents.insert(parent); + } + } + for parent in parents { + File::open(parent)?.sync_all()?; + } + Ok(()) +} + +#[cfg(not(unix))] +fn sync_parent_directories(_changes: &[Change]) -> Result<(), PatchError> { + Ok(()) +} + +fn read_snapshot(path: &Path) -> Result { + let metadata = match fs::symlink_metadata(path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == io::ErrorKind::NotFound => { + return Ok(FileSnapshot { + kind: EntryKind::Missing, + data: Vec::new(), + mode: 0, + }); + } + Err(error) => return Err(error.into()), + }; + let mode = metadata.mode(); + if metadata.file_type().is_symlink() { + return Ok(FileSnapshot { + kind: EntryKind::Symlink, + data: Vec::new(), + mode, + }); + } + if metadata.is_dir() { + return Ok(FileSnapshot { + kind: EntryKind::Directory, + data: Vec::new(), + mode, + }); + } + if !metadata.is_file() { + return Ok(FileSnapshot { + kind: EntryKind::Other, + data: Vec::new(), + mode, + }); + } + let data = fs::read(path)?; + Ok(FileSnapshot { + kind: EntryKind::Regular, + data, + mode, + }) +} + +fn parse_create_diff(diff: &str) -> Result, PatchError> { + let lines = normalized_lines(diff); + let mut output = Vec::with_capacity(diff.len()); + for line in lines { + if !line.starts_with('+') { + return Err(PatchError::Invalid(format!( + "invalid add file line: {line}" + ))); + } + output.extend_from_slice(&line.as_bytes()[1..]); + output.push(b'\n'); + } + if !output.is_empty() { + output.pop(); + } + Ok(output) +} + +fn apply_update_diff(input: &[u8], diff: &str) -> Result, PatchError> { + let input = std::str::from_utf8(input) + .map_err(|_| PatchError::Invalid("update target is not valid UTF-8".to_string()))?; + let lines = normalized_lines(diff); + let source: Vec = input.split('\n').map(ToOwned::to_owned).collect(); + let mut chunks = Vec::::new(); + let mut index = 0; + let mut cursor = 0; + + while index < lines.len() { + let mut anchors = Vec::new(); + let mut anchor_count = 0; + while index < lines.len() { + if let Some(anchor) = lines[index].strip_prefix("@@ ") { + anchor_count += 1; + if !anchor.trim().is_empty() { + anchors.push(anchor.to_string()); + } + index += 1; + } else if lines[index] == "@@" { + anchor_count += 1; + index += 1; + } else { + break; + } + } + if anchor_count == 0 && cursor != 0 { + return Err(PatchError::Invalid(format!( + "invalid line: {}", + lines.get(index).cloned().unwrap_or_default() + ))); + } + let require_anchor = anchor_count > 1; + for (anchor_index, anchor) in anchors.iter().enumerate() { + if let Some(found) = advance_to_anchor(&source, anchor, cursor, false, anchor_index > 0) + { + cursor = found; + } else if let Some(found) = + advance_to_anchor(&source, anchor, cursor, true, anchor_index > 0) + { + cursor = found; + } else if require_anchor { + return Err(PatchError::Invalid(format!( + "invalid anchor {cursor}: {anchor}" + ))); + } + } + + let (context, section_chunks, next_index, eof) = read_diff_section(&lines, index)?; + let new_index = find_context(&source, &context, cursor, eof).ok_or_else(|| { + let label = if eof { "EOF context" } else { "context" }; + PatchError::Invalid(format!("invalid {label} {cursor}: {}", context.join("\n"))) + })?; + for mut chunk in section_chunks { + chunk.original_index += new_index; + chunks.push(chunk); + } + cursor = new_index + context.len(); + index = next_index; + } + + let mut output = Vec::with_capacity(source.len()); + let mut original_index = 0; + for chunk in chunks { + if chunk.original_index < original_index || chunk.original_index > source.len() { + return Err(PatchError::Conflict(format!( + "overlapping diff chunk at {}", + chunk.original_index + ))); + } + output.extend(source[original_index..chunk.original_index].iter().cloned()); + output.extend(chunk.inserted_lines); + original_index = chunk.original_index + chunk.deleted_lines.len(); + } + output.extend(source[original_index..].iter().cloned()); + Ok(output.join("\n").into_bytes()) +} + +#[derive(Debug)] +struct DiffChunk { + original_index: usize, + deleted_lines: Vec, + inserted_lines: Vec, +} + +fn read_diff_section( + lines: &[String], + start: usize, +) -> Result<(Vec, Vec, usize, bool), PatchError> { + let mut context = Vec::new(); + let mut deleted = Vec::new(); + let mut inserted = Vec::new(); + let mut chunks = Vec::new(); + let mut mode = ' '; + let mut index = start; + + let flush = |context: &Vec, + deleted: &mut Vec, + inserted: &mut Vec, + chunks: &mut Vec| { + if deleted.is_empty() && inserted.is_empty() { + return; + } + chunks.push(DiffChunk { + original_index: context.len() - deleted.len(), + deleted_lines: std::mem::take(deleted), + inserted_lines: std::mem::take(inserted), + }); + }; + + while index < lines.len() { + let raw = &lines[index]; + if raw.starts_with("@@") + || raw.starts_with(END_MARKER) + || is_operation_header(raw) + || raw == END_OF_FILE_MARKER + || raw == "***" + { + break; + } + if raw.starts_with("***") { + return Err(PatchError::Invalid(format!("invalid line: {raw}"))); + } + index += 1; + let line = if raw.is_empty() { + " ".to_string() + } else { + raw.clone() + }; + let next_mode = line.chars().next().unwrap_or(' '); + if !matches!(next_mode, '+' | '-' | ' ') { + return Err(PatchError::Invalid(format!("invalid line: {raw}"))); + } + if next_mode == ' ' && mode != ' ' { + flush(&context, &mut deleted, &mut inserted, &mut chunks); + } + mode = next_mode; + let value = line[1..].to_string(); + match mode { + '-' => { + deleted.push(value.clone()); + context.push(value); + } + '+' => inserted.push(value), + ' ' => context.push(value), + _ => unreachable!(), + } + } + flush(&context, &mut deleted, &mut inserted, &mut chunks); + let eof = lines.get(index).map(String::as_str) == Some(END_OF_FILE_MARKER); + if eof { + index += 1; + } + if index == start { + return Err(PatchError::Invalid(format!( + "nothing in diff section at index {index}" + ))); + } + Ok((context, chunks, index, eof)) +} + +fn advance_to_anchor( + lines: &[String], + target: &str, + cursor: usize, + trimmed: bool, + force_forward: bool, +) -> Option { + if !force_forward + && lines + .iter() + .take(cursor.min(lines.len())) + .any(|line| line_equal(line, target, trimmed)) + { + return Some(cursor); + } + lines + .iter() + .enumerate() + .skip(cursor) + .find_map(|(index, line)| line_equal(line, target, trimmed).then_some(index + 1)) +} + +fn line_equal(left: &str, right: &str, trimmed: bool) -> bool { + if trimmed { + left.trim() == right.trim() + } else { + left == right + } +} + +fn find_context(lines: &[String], context: &[String], start: usize, eof: bool) -> Option { + if context.is_empty() { + return Some(start.min(lines.len())); + } + let preferred_start = if eof { + lines.len().saturating_sub(context.len()) + } else { + start + }; + find_context_variants(lines, context, preferred_start).or_else(|| { + eof.then(|| find_context_variants(lines, context, start)) + .flatten() + }) +} + +fn find_context_variants(lines: &[String], context: &[String], start: usize) -> Option { + find_context_with(lines, context, start, |line| line.to_string()) + .or_else(|| find_context_with(lines, context, start, |line| line.trim_end().to_string())) + .or_else(|| find_context_with(lines, context, start, |line| line.trim().to_string())) +} + +fn find_context_with(lines: &[String], context: &[String], start: usize, map: F) -> Option +where + F: Fn(&str) -> String, +{ + if context.len() > lines.len() { + return None; + } + (start..=lines.len() - context.len()).find(|&index| { + context + .iter() + .enumerate() + .all(|(offset, target)| map(&lines[index + offset]) == map(target)) + }) +} + +fn text_for_diff(file: &FileSnapshot) -> Result { + if file.kind == EntryKind::Missing { + return Ok(String::new()); + } + if file.kind != EntryKind::Regular { + return Err(PatchError::Conflict( + "diff result contains a non-regular file".to_string(), + )); + } + String::from_utf8(file.data.clone()) + .map_err(|_| PatchError::Invalid("file is not valid UTF-8".to_string())) +} + +fn text_for_diff_virtual(file: &VirtualFile) -> Result { + if file.kind == EntryKind::Missing { + return Ok(String::new()); + } + if file.kind != EntryKind::Regular { + return Err(PatchError::Conflict( + "diff result contains a non-regular file".to_string(), + )); + } + String::from_utf8(file.data.clone()) + .map_err(|_| PatchError::Invalid("file is not valid UTF-8".to_string())) +} + +fn unified_diff( + old: &str, + new: &str, + old_header: &str, + new_header: &str, +) -> (String, usize, usize) { + let diff = TextDiff::from_lines(old, new); + let mut unified = diff.unified_diff(); + unified.header(old_header, new_header); + let rendered = unified.to_string(); + let mut added = 0; + let mut removed = 0; + for change in diff.iter_all_changes() { + match change.tag() { + ChangeTag::Insert => added += 1, + ChangeTag::Delete => removed += 1, + ChangeTag::Equal => {} + } + } + (rendered, added, removed) +} + +fn relative_display(cwd: &Path, path: &Path) -> String { + path.strip_prefix(cwd).unwrap_or(path).display().to_string() +} + +fn path_exists(path: &Path) -> Result { + match fs::symlink_metadata(path) { + Ok(_) => Ok(true), + Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(false), + Err(error) => Err(error.into()), + } +} + +fn unique_path(candidate: &Path) -> Result { + if !path_exists(candidate)? { + return Ok(candidate.to_path_buf()); + } + for suffix in 1..1000 { + let mut path = candidate.to_path_buf(); + path.set_extension(format!("stage-{suffix}")); + if !path_exists(&path)? { + return Ok(path); + } + } + Err(PatchError::Conflict(format!( + "cannot allocate temporary path near {}", + candidate.display() + ))) +} + +fn backup_path(path: &Path, transaction_id: &str) -> Result { + let parent = path.parent().unwrap_or_else(|| Path::new(".")); + let file_name = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("file"); + unique_path(&parent.join(format!(".{file_name}.{transaction_id}.backup"))) +} + +fn create_missing_dirs(path: &Path) -> Result, PatchError> { + let mut missing = Vec::new(); + let mut current = path.to_path_buf(); + while !path_exists(¤t)? { + missing.push(current.clone()); + if !current.pop() { + break; + } + } + fs::create_dir_all(path)?; + Ok(missing) +} + +#[cfg(unix)] +fn set_mode(file: &File, mode: u32) -> Result<(), PatchError> { + use std::os::unix::fs::PermissionsExt; + file.set_permissions(fs::Permissions::from_mode(mode & 0o7777))?; + Ok(()) +} + +#[cfg(not(unix))] +fn set_mode(_file: &File, _mode: u32) -> Result<(), PatchError> { + Ok(()) +} + +#[cfg(unix)] +trait MetadataMode { + fn mode(&self) -> u32; +} + +#[cfg(unix)] +impl MetadataMode for fs::Metadata { + fn mode(&self) -> u32 { + use std::os::unix::fs::PermissionsExt; + self.permissions().mode() + } +} + +#[cfg(not(unix))] +trait MetadataMode { + fn mode(&self) -> u32; +} + +#[cfg(not(unix))] +impl MetadataMode for fs::Metadata { + fn mode(&self) -> u32 { + 0o644 + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use std::time::{SystemTime, UNIX_EPOCH}; + + fn temp_dir(name: &str) -> PathBuf { + let suffix = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos(); + let path = std::env::temp_dir().join(format!("agenty-patch-applier-{name}-{suffix}")); + fs::create_dir_all(&path).unwrap(); + path + } + + #[test] + fn applies_repeated_operations_in_file_order_and_reports_diff() { + let cwd = temp_dir("ordered"); + let patch = "*** Begin Patch\n*** Add File: notes.txt\n+one\n*** Update File: notes.txt\n@@\n-one\n+two\n*** Update File: notes.txt\n@@\n-two\n+three\n*** End Patch"; + let result = apply_patch(&cwd, patch).unwrap(); + assert_eq!(fs::read_to_string(cwd.join("notes.txt")).unwrap(), "three"); + assert_eq!(result.files.len(), 1); + assert_eq!(result.files[0].added_lines, 1); + assert_eq!(result.files[0].removed_lines, 0); + assert!(result.files[0].diff.contains("+three")); + } + + #[test] + fn rejects_conflict_without_writing_any_file() { + let cwd = temp_dir("conflict"); + let patch = "*** Begin Patch\n*** Add File: created.txt\n+created\n*** Update File: missing.txt\n@@\n-missing\n+updated\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err().to_string(); + assert!(error.contains("updates a non-existent")); + assert!(!cwd.join("created.txt").exists()); + } + + #[test] + fn moves_then_updates_destination_as_one_file_group() { + let cwd = temp_dir("move"); + fs::write(cwd.join("old.txt"), "one\n").unwrap(); + let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** Update File: new.txt\n@@\n-two\n+three\n*** End Patch"; + apply_patch(&cwd, patch).unwrap(); + assert!(!cwd.join("old.txt").exists()); + assert_eq!(fs::read_to_string(cwd.join("new.txt")).unwrap(), "three\n"); + } + + #[test] + fn malformed_patch_has_no_side_effect() { + let cwd = temp_dir("malformed"); + let error = apply_patch(&cwd, "*** Begin Patch\n*** Add File: x.txt\n+x").unwrap_err(); + assert!(error.to_string().contains("must end")); + assert!(!cwd.join("x.txt").exists()); + } + + #[test] + fn accepts_trailing_newline_and_reports_update_counts() { + let cwd = temp_dir("trailing-newline"); + fs::write(cwd.join("notes.txt"), "one\ntwo\nthree\n").unwrap(); + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n-two\n+alpha\n three\n*** End Patch\n"; + let result = apply_patch(&cwd, patch).unwrap(); + assert_eq!( + fs::read_to_string(cwd.join("notes.txt")).unwrap(), + "alpha\nthree\n" + ); + assert_eq!(result.files[0].added_lines, 1); + assert_eq!(result.files[0].removed_lines, 2); + assert!(result.files[0].diff.contains("-one")); + assert!(result.files[0].diff.contains("-two")); + assert!(result.files[0].diff.contains("+alpha")); + } + + #[test] + fn applies_multiple_hunks_with_reused_anchor() { + let cwd = temp_dir("multi-hunk"); + fs::write(cwd.join("notes.txt"), "header\none\nmiddle\ntwo\nfooter\n").unwrap(); + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@ header\n-one\n+first\n@@ middle\n-two\n+second\n*** End Patch"; + apply_patch(&cwd, patch).unwrap(); + assert_eq!( + fs::read_to_string(cwd.join("notes.txt")).unwrap(), + "header\nfirst\nmiddle\nsecond\nfooter\n" + ); + } + + #[test] + fn applies_end_of_file_section() { + let cwd = temp_dir("eof"); + fs::write(cwd.join("notes.txt"), "one\ntwo\nthree").unwrap(); + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n two\n-three\n+final\n*** End of File\n*** End Patch"; + apply_patch(&cwd, patch).unwrap(); + assert_eq!( + fs::read_to_string(cwd.join("notes.txt")).unwrap(), + "one\ntwo\nfinal" + ); + } + + #[test] + fn detects_cross_file_move_conflict_before_writing() { + let cwd = temp_dir("move-conflict"); + let patch = "*** Begin Patch\n*** Add File: first.txt\n+first\n*** Add File: second.txt\n+second\n*** Update File: first.txt\n*** Move to: second.txt\n@@\n-first\n+updated\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err().to_string(); + assert!(error.contains("owned by another file")); + assert!(!cwd.join("first.txt").exists()); + assert!(!cwd.join("second.txt").exists()); + } + + #[test] + fn rejects_duplicate_create_before_writing_other_files() { + let cwd = temp_dir("duplicate-create"); + let patch = "*** Begin Patch\n*** Add File: first.txt\n+first\n*** Add File: second.txt\n+second\n*** Add File: first.txt\n+again\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err().to_string(); + assert!(error.contains("creates an existing")); + assert!(!cwd.join("first.txt").exists()); + assert!(!cwd.join("second.txt").exists()); + } + + #[test] + fn refuses_to_overwrite_file_changed_during_preparation() { + let cwd = temp_dir("concurrent-change"); + fs::write(cwd.join("notes.txt"), "one").unwrap(); + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** End Patch"; + let operations = parse_envelope(&cwd, patch).unwrap(); + let mut transaction = Transaction::new(cwd.clone()); + for operation in operations { + transaction.apply(operation).unwrap(); + } + + fs::write(cwd.join("notes.txt"), "external").unwrap(); + let error = transaction.commit().unwrap_err().to_string(); + assert!(error.contains("changed while the patch was being prepared")); + assert_eq!( + fs::read_to_string(cwd.join("notes.txt")).unwrap(), + "external" + ); + } + + #[test] + fn preserves_v4a_diff_compatibility() { + let create_cases = [ + ("+hello\n+world\n+", "hello\nworld\n"), + ("+hello\r\n+\r\n+world\r\n", "hello\n\nworld"), + ("", ""), + ]; + for (diff, expected) in create_cases { + assert_eq!( + String::from_utf8(parse_create_diff(diff).unwrap()).unwrap(), + expected + ); + } + + let update_cases = [ + ("", "@@\n+hello\n+world", "hello\nworld\n"), + ( + "- Milk\n- Bread\n- Eggs\n- Apples\n- Coffee", + "@@\n - Milk\n - Bread\n - Eggs\n-- Apples\n-- Coffee\n+- [x] Apples\n+- [x] Coffee", + "- Milk\n- Bread\n- Eggs\n- [x] Apples\n- [x] Coffee", + ), + ( + "class First\n def target():\n pass\n\nclass Second\n def target():\n pass\n", + "@@ class Second\n@@ def target():\n- pass\n+ return 1", + "class First\n def target():\n pass\n\nclass Second\n def target():\n return 1\n", + ), + ( + " class Target \n def first():\n pass\n\n def second():\n pass\n", + "@@ class Target\n@@ def first():\n- pass\n+ return 1\n@@ class Target\n@@ def second():\n- pass\n+ return 2", + " class Target \n def first():\n return 1\n\n def second():\n return 2\n", + ), + ("one \ntwo\n", " one\n-two\n+second", "one \nsecond\n"), + ("one\ntwo\n", " one\r\n-two\r\n+second\r\n", "one\nsecond\n"), + ( + "target\nmiddle\nend", + " target\n+after\n*** End of File", + "target\nafter\nmiddle\nend", + ), + ]; + for (input, diff, expected) in update_cases { + assert_eq!( + String::from_utf8(apply_update_diff(input.as_bytes(), diff).unwrap()).unwrap(), + expected + ); + } + } + + #[test] + fn rejects_invalid_v4a_diff_sections() { + let cases = [ + ("one\ntwo\n", " missing\n-two\n+second", "invalid context"), + ( + "class Wrong\n def desired():\n pass\n", + "@@ class Target\n@@ def desired():\n- pass\n+ return 1", + "invalid anchor", + ), + ("one\n", "one", "invalid line"), + ("one\n", "*** Unknown Directive", "invalid line"), + ("one\n", "@@", "nothing in diff section"), + ( + "one\ntwo", + " missing\n*** End of File", + "invalid EOF context", + ), + ]; + for (input, diff, expected) in cases { + let error = apply_update_diff(input.as_bytes(), diff) + .unwrap_err() + .to_string(); + assert!( + error.contains(expected), + "error {error:?} does not contain {expected:?}" + ); + } + let error = parse_create_diff("+valid\ninvalid") + .unwrap_err() + .to_string(); + assert!(error.contains("invalid add file line")); + } + + #[test] + fn lexical_normalization_preserves_absolute_roots() { + let root = Path::new(std::path::MAIN_SEPARATOR_STR); + let nested = root.join("tmp").join("..").join("file.txt"); + assert_eq!(lexical_normalize(&nested).unwrap(), root.join("file.txt")); + + let escaping = root.join("..").join("file.txt"); + assert!(lexical_normalize(&escaping).is_err()); + } +} diff --git a/packages/patch-applier/src/main.rs b/packages/patch-applier/src/main.rs new file mode 100644 index 0000000..531b64c --- /dev/null +++ b/packages/patch-applier/src/main.rs @@ -0,0 +1,61 @@ +use std::io::{self, Read, Write}; +use std::path::PathBuf; + +use patch_applier::{apply_patch, FileResult, PatchResult}; +use serde::Serialize; + +#[derive(Debug, Serialize)] +struct ErrorResult { + success: bool, + cwd: String, + files: Vec, + error: String, +} + +fn main() { + std::process::exit(exit_code()); +} + +fn exit_code() -> i32 { + let result = run(); + match result { + Ok(output) => { + if let Err(error) = write_json(&output) { + eprintln!("apply_patch: write result: {error}"); + return 1; + } + 0 + } + Err(error) => { + let output = ErrorResult { + success: false, + cwd: std::env::current_dir() + .map(|path| path.display().to_string()) + .unwrap_or_default(), + files: Vec::new(), + error: error.to_string(), + }; + if let Err(write_error) = write_json(&output) { + eprintln!("apply_patch: write error result: {write_error}"); + } + eprintln!("apply_patch: {error}"); + 1 + } + } +} + +fn write_json(value: &T) -> Result<(), Box> { + let stdout = io::stdout(); + let mut output = stdout.lock(); + serde_json::to_writer(&mut output, value)?; + output.write_all(b"\n")?; + output.flush()?; + Ok(()) +} + +fn run() -> Result> { + let mut patch = String::new(); + io::stdin().read_to_string(&mut patch)?; + let cwd: PathBuf = std::env::current_dir()?; + Ok(apply_patch(&cwd, &patch)?) +} diff --git a/packages/patch-applier/tests/cli.rs b/packages/patch-applier/tests/cli.rs new file mode 100644 index 0000000..3d48e62 --- /dev/null +++ b/packages/patch-applier/tests/cli.rs @@ -0,0 +1,83 @@ +use std::fs; +use std::io::Write; +use std::process::{Command, Stdio}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde_json::Value; + +fn temp_dir(name: &str) -> std::path::PathBuf { + let suffix = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock must follow the Unix epoch") + .as_nanos(); + let path = std::env::temp_dir().join(format!("agenty-patch-cli-{name}-{suffix}")); + fs::create_dir_all(&path).expect("create test directory"); + path +} + +#[test] +fn prints_complete_success_json() { + let cwd = temp_dir("success"); + let mut child = Command::new(env!("CARGO_BIN_EXE_apply_patch")) + .current_dir(&cwd) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .expect("start apply_patch"); + child + .stdin + .take() + .expect("capture stdin") + .write_all(b"*** Begin Patch\n*** Add File: notes.txt\n+hello\n*** End Patch\n") + .expect("write patch"); + let output = child.wait_with_output().expect("wait for apply_patch"); + assert!( + output.status.success(), + "stderr: {}", + String::from_utf8_lossy(&output.stderr) + ); + + let result: Value = serde_json::from_slice(&output.stdout).expect("decode stdout JSON"); + assert_eq!(result["success"], true); + assert_eq!( + result["cwd"], + cwd.canonicalize().unwrap().display().to_string() + ); + assert_eq!(result["files"][0]["path"], "notes.txt"); + assert_eq!(result["files"][0]["addedLines"], 1); + assert_eq!(result["files"][0]["removedLines"], 0); + assert!(result["files"][0]["diff"] + .as_str() + .unwrap() + .contains("+hello")); +} + +#[test] +fn prints_complete_error_json() { + let cwd = temp_dir("error"); + let mut child = Command::new(env!("CARGO_BIN_EXE_apply_patch")) + .current_dir(&cwd) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .expect("start apply_patch"); + child + .stdin + .take() + .expect("capture stdin") + .write_all(b"malformed") + .expect("write patch"); + let output = child.wait_with_output().expect("wait for apply_patch"); + assert!(!output.status.success()); + + let result: Value = serde_json::from_slice(&output.stdout).expect("decode stdout JSON"); + assert_eq!(result["success"], false); + assert_eq!( + result["cwd"], + cwd.canonicalize().unwrap().display().to_string() + ); + assert_eq!(result["files"], Value::Array(Vec::new())); + assert!(result["error"].as_str().unwrap().contains("must start")); +} From 819a2470b3b8e8e94982ce7d09411a7bc7741df2 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Mon, 24 Aug 2026 18:37:38 +0800 Subject: [PATCH 04/12] chore: update build and development workflow Signed-off-by: masteryyh --- .github/workflows/publish-release.yaml | 4 +- AGENTS.md | 15 ++-- README.md | 14 ++-- README.zh-CN.md | 12 ++- package.json | 9 ++- packages/agenty-bootstrap/package.json | 3 +- .../agenty-bootstrap/scripts/footer.test.ts | 10 ++- packages/agenty-bootstrap/scripts/footer.ts | 15 ++-- packages/agenty-bootstrap/scripts/pack.ts | 24 ++++-- packages/agenty-bootstrap/src/lib.rs | 78 ++++++++++++++----- packages/agenty-bootstrap/src/main.rs | 22 ++++-- packages/agenty-cli/src/localCore.test.ts | 12 ++- packages/agenty-cli/src/localCore.ts | 7 +- packages/agenty-core/README-CN.md | 45 ++++++----- packages/agenty-core/README.md | 50 +++++++----- packages/agenty-core/TESTING-CN.md | 2 +- packages/agenty-core/TESTING.md | 2 +- packages/agenty-core/package.json | 5 +- pnpm-lock.yaml | 69 +++++++++------- pnpm-workspace.yaml | 1 + turbo.json | 10 ++- 21 files changed, 275 insertions(+), 134 deletions(-) diff --git a/.github/workflows/publish-release.yaml b/.github/workflows/publish-release.yaml index c3f70bc..e8a0cee 100644 --- a/.github/workflows/publish-release.yaml +++ b/.github/workflows/publish-release.yaml @@ -108,7 +108,9 @@ jobs: - name: Cache cargo uses: Swatinem/rust-cache@v2 with: - workspaces: packages/agenty-bootstrap + workspaces: | + packages/agenty-bootstrap + packages/patch-applier - name: Install dependencies run: pnpm install --frozen-lockfile diff --git a/AGENTS.md b/AGENTS.md index 9887423..cec04bd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -3,10 +3,11 @@ ## Project overview Agenty is a local-first AI agent application organized as a pnpm + Turborepo monorepo. -The active product path has three workspaces: +The active product path has four workspaces: - `packages/agenty-core`: Go 1.26 core process and stdio JSON-RPC 2.0 server. - `packages/agenty-cli`: Bun/TypeScript/React OpenTUI client. +- `packages/patch-applier`: Rust V4A parser and transactional filesystem patch helper. - `packages/agenty-bootstrap`: Rust self-extracting launcher. The CLI starts core as a child process and communicates only through NDJSON messages on @@ -45,7 +46,9 @@ Core data is local-first: - Config: `~/.agenty/config.json` - Sessions: append-only JSONL under `~/.agenty/sessions/` - Session projection: `~/.agenty/agenty.sqlite` -- Providers/models: `~/.agenty/providers/.json` (models embedded) +- Built-in providers/models: embedded in the core binary; custom providers use + `~/.agenty/providers/.json`, while built-in provider files contain + only the API key. - Agents: `~/.agenty/agents/` - Logs: `~/.agenty/logs///
/core.log` @@ -79,17 +82,17 @@ and breaks terminal capability handshakes. The bootstrap artifact layout is: -`[bootstrap stub][xz CLI][xz core][108-byte footer]` +`[bootstrap stub][xz CLI][xz core][xz apply_patch][156-byte footer]` -The footer stores offsets, lengths, and SHA3-256 digests of decompressed payloads. +The footer stores offsets, lengths, and SHA3-256 digests of the three decompressed payloads. `src/lib.rs` and `scripts/footer.ts` are one wire contract; changing the layout requires updating both golden tests and incrementing `FORMAT_VERSION`. Compression uses `@napi-rs/lzma`; Rust decompression uses statically linked vendored liblzma. Code signing must happen after payload packing. pnpm owns workspace resolution, Turborepo owns build ordering/caching, Bun builds the -CLI and packs payloads, Go builds core, and Cargo builds the launcher. Do not add an npm -`workspaces` field. The dependency graph builds core, then CLI, then bootstrap. +CLI and packs payloads, Go builds core, and Cargo builds the patch helper and launcher. Do not add an npm +`workspaces` field. The dependency graph builds the patch helper before core, then CLI and bootstrap. Root `.env` is the single `AGENTY_VERSION` source and stays ignored; only `.env.example` is committed. Release CI passes target-specific `GOOS`, `GOARCH`, `CC`, diff --git a/README.md b/README.md index 2ee4c2d..b0fd485 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,8 @@ [简体中文](./README.zh-CN.md) Agenty is a local-first AI agent application. The current product path consists of -`agenty-cli`, `agenty-core`, and the self-extracting `agenty-bootstrap` launcher. +`agenty-cli`, `agenty-core`, the Rust `patch-applier` helper, and the self-extracting +`agenty-bootstrap` launcher. The CLI communicates with core exclusively through line-delimited JSON-RPC 2.0 over the child process's stdin/stdout; it does not start an HTTP server. @@ -24,14 +25,14 @@ sudo install -m 755 agenty /usr/local/bin/agenty agenty ``` -On first run, the launcher verifies and extracts the bundled CLI and core into -`~/.agenty/bin/{cli,core}`. The CLI starts core as a child process and opens a setup +On first run, the launcher verifies and extracts the bundled CLI, core, and patch helper into +`~/.agenty/bin/{cli,core,apply_patch}`. The CLI starts core as a child process and opens a setup wizard. The wizard creates one provider, one chat model, and one default agent through the existing `provider.*` and `agent.*` IPC methods, then calls `initialize.complete`. ## Runtime model -The launcher contains two XZ-compressed payloads and their decompressed SHA3-256 +The launcher contains three XZ-compressed payloads and their decompressed SHA3-256 digests. Matching extracted files are reused; missing or mismatched files are verified and atomically replaced. The CLI resolves core in this order: @@ -39,6 +40,9 @@ and atomically replaced. The CLI resolves core in this order: 2. `packages/agenty-core/bin/agenty-core` during repository development 3. `~/.agenty/bin/core` from the launcher +Before starting core, the CLI prepends core's directory to `PATH`, making the bundled +`apply_patch` command available to core and shell tool calls. + Core reads one compact JSON-RPC message per stdin line and writes responses and notifications to stdout. After `session.start`, core sends ordered `session.event` notifications for round lifecycle, persisted messages, model stream deltas, tool calls, @@ -59,7 +63,7 @@ Core stores data under `~/.agenty` by default. Pass `--data-dir ` to the C | Configuration | `~/.agenty/config.json` | | Session transcripts | `~/.agenty/sessions///
/.jsonl` | | Session index | `~/.agenty/agenty.sqlite` | -| Providers and models | `~/.agenty/providers/.json` (models embedded) | +| Providers and models | Built-in catalog is embedded in the core binary; custom providers use `~/.agenty/providers/.json`, built-in provider files store only API keys; core automatically discovers empty configured catalogs when listing providers/models and caches them under `~/.agenty/providers/.models/` for 8 hours | | Agents | `~/.agenty/agents/` | | Logs | `~/.agenty/logs///
/core.log` | diff --git a/README.zh-CN.md b/README.zh-CN.md index c5dcb7d..1dddad9 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -3,7 +3,8 @@ [English](./README.md) Agenty 是一个本地优先的 AI agent 应用。当前产品链路由 `agenty-cli`、 -`agenty-core` 和自解压 launcher `agenty-bootstrap` 组成。CLI 仅通过子进程 +`agenty-core`、Rust `patch-applier` helper 和自解压 launcher `agenty-bootstrap` 组成。 +CLI 仅通过子进程 stdin/stdout 上的逐行 JSON-RPC 2.0 与 core 通信,不再启动 HTTP server。 core 当前支持 provider/model/agent 管理、持久化会话、模型流式输出、agent 工具循环 @@ -22,19 +23,22 @@ agenty ``` 首次运行时,launcher 会校验并释放内置的 CLI 和 core 到 -`~/.agenty/bin/{cli,core}`。CLI 启动 core 子进程并打开初始化向导;向导通过 +`~/.agenty/bin/{cli,core,apply_patch}`。CLI 启动 core 子进程并打开初始化向导;向导通过 已有的 `provider.*` 和 `agent.*` IPC methods 创建一个 provider、一个聊天 model 和一个默认 agent, 最后调用 `initialize.complete` 标记初始化完成。 ## 运行模型 -launcher 内含两个 XZ 压缩 payload 及其解压内容的 SHA3-256 摘要。已释放文件摘要一致时 +launcher 内含三个 XZ 压缩 payload 及其解压内容的 SHA3-256 摘要。已释放文件摘要一致时 直接复用;缺失或不一致时会重新解压、校验并原子替换。CLI 按以下顺序查找 core: 1. `AGENTY_CORE_BIN` 2. 仓库开发环境中的 `packages/agenty-core/bin/agenty-core` 3. launcher 释放的 `~/.agenty/bin/core` +CLI 启动 core 前会把 core 所在目录放到 `PATH` 首位,使 core 和 shell 工具调用可以找到 +同目录中的 `apply_patch`。 + core 从 stdin 逐行读取紧凑 JSON-RPC message,并把 response 和 notification 写到 stdout。 调用 `session.start` 后,core 会持续发送有序的 `session.event` 通知,覆盖 round 生命周期、 已持久化消息、模型流式增量、工具调用和 round 终态。通知可能早于 `session.start` response @@ -53,7 +57,7 @@ core 默认把数据保存在 `~/.agenty`。可向 CLI 传入 `--data-dir | 配置 | `~/.agenty/config.json` | | 会话 transcript | `~/.agenty/sessions///
/.jsonl` | | 会话索引 | `~/.agenty/agenty.sqlite` | -| Providers 和 models | `~/.agenty/providers/.json`(模型内嵌) | +| Providers 和 models | 内置 catalog 固化在 core 二进制中;自定义 provider 使用 `~/.agenty/providers/.json`,内置 provider 文件仅保存 API key;core 在获取 provider/model 列表时自动发现已配置的空模型 catalog,结果缓存于 `~/.agenty/providers/.models/`,有效期 8 小时 | | Agents | `~/.agenty/agents/` | | 日志 | `~/.agenty/logs///
/core.log` | diff --git a/package.json b/package.json index fee93c2..c577efa 100644 --- a/package.json +++ b/package.json @@ -8,6 +8,7 @@ "lint": "eslint .", "lint:fix": "eslint . --fix", "clean": "rm -rf dist && turbo run clean", + "dev": "turbo run build --filter=agenty-bootstrap && pnpm --filter agenty-bootstrap exec bun -e \"const { join } = await import('node:path'); const { resolveArch, resolveOS } = await import('./scripts/target.ts'); const os = resolveOS(); const executable = join('bin', 'agenty-' + os + '-' + resolveArch() + (os === 'windows' ? '.exe' : '')); const child = Bun.spawn([executable, ...process.argv.slice(1)], { stdin: 'inherit', stdout: 'inherit', stderr: 'inherit', env: process.env }); process.exit(await child.exited);\"", "deepclean": "pnpm clean && rm -rf node_modules packages/*/node_modules", "core:build": "turbo run build --filter=agenty-core", "core:test": "turbo run test --filter=agenty-core", @@ -18,6 +19,10 @@ "core:test:repeat": "pnpm --filter agenty-core test:repeat", "core:tidyup": "cd packages/agenty-core && go fmt ./... && go vet ./... && go mod tidy", "core:clean": "pnpm --filter agenty-core clean", + "patch-applier:build": "turbo run build --filter=patch-applier", + "patch-applier:test": "turbo run test --filter=patch-applier", + "patch-applier:lint": "pnpm --filter patch-applier lint", + "patch-applier:clean": "pnpm --filter patch-applier clean", "cli:build": "turbo run build --filter=agenty-cli", "cli:dev": "turbo run build --filter=agenty-core && pnpm --filter agenty-cli dev", "cli:typecheck": "pnpm --filter agenty-cli typecheck", @@ -33,7 +38,7 @@ "eslint-plugin-import-newlines": "^2.0.0", "eslint-plugin-jsonc": "^3.4.1", "eslint-plugin-simple-import-sort": "^14.0.0", - "turbo": "^2.10.10", + "turbo": "^2.10.11", "typescript-eslint": "^8.67.0" } -} \ No newline at end of file +} diff --git a/packages/agenty-bootstrap/package.json b/packages/agenty-bootstrap/package.json index 2aaa6c4..5b64959 100644 --- a/packages/agenty-bootstrap/package.json +++ b/packages/agenty-bootstrap/package.json @@ -18,6 +18,7 @@ }, "devDependencies": { "agenty-cli": "workspace:*", - "agenty-core": "workspace:*" + "agenty-core": "workspace:*", + "patch-applier": "workspace:*" } } diff --git a/packages/agenty-bootstrap/scripts/footer.test.ts b/packages/agenty-bootstrap/scripts/footer.test.ts index b2c4fb2..b69a8bf 100644 --- a/packages/agenty-bootstrap/scripts/footer.test.ts +++ b/packages/agenty-bootstrap/scripts/footer.test.ts @@ -10,7 +10,8 @@ import { encodeFooter, FOOTER_SIZE } from "./footer"; const GOLDEN_FOOTER_HEX = "88776655443322110807060504030201000102030405060708090a0b0c0d0e0f" + "101112131415161718191a1b1c1d1e1f1122334455667788010203040506070820212223242526" + - "2728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f01000000cafebabe10136666"; + "2728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f08090a0b0c0d0e0f1011121314151617" + + "404142434445464748494a4b4c4d4e4f505152535455565758595a5b5c5d5e5f02000000cafebabe10136666"; function toHex(bytes: Uint8Array): string { return Array.from(bytes, (b) => b.toString(16).padStart(2, "0")).join(""); @@ -29,6 +30,11 @@ describe("encodeFooter", () => { len: 0x0807060504030201n, sha3_256: Uint8Array.from({ length: 32 }, (_, i) => 0x20 + i), }, + { + offset: 0x0f0e0d0c0b0a0908n, + len: 0x1716151413121110n, + sha3_256: Uint8Array.from({ length: 32 }, (_, i) => 0x40 + i), + }, ); expect(footer.length).toBe(FOOTER_SIZE); @@ -37,6 +43,6 @@ describe("encodeFooter", () => { test("rejects non-32-byte digests", () => { const spec = { offset: 0n, len: 0n, sha3_256: new Uint8Array(31) }; - expect(() => encodeFooter(spec, spec)).toThrow("32 bytes"); + expect(() => encodeFooter(spec, spec, spec)).toThrow("32 bytes"); }); }); diff --git a/packages/agenty-bootstrap/scripts/footer.ts b/packages/agenty-bootstrap/scripts/footer.ts index 48e9329..5fcebc2 100644 --- a/packages/agenty-bootstrap/scripts/footer.ts +++ b/packages/agenty-bootstrap/scripts/footer.ts @@ -1,6 +1,6 @@ export const MAGIC = [0xca, 0xfe, 0xba, 0xbe, 0x10, 0x13, 0x66, 0x66] as const; -export const FORMAT_VERSION = 1; -export const FOOTER_SIZE = 108; +export const FORMAT_VERSION = 2; +export const FOOTER_SIZE = 156; export interface PayloadSpec { offset: bigint; @@ -8,8 +8,8 @@ export interface PayloadSpec { sha3_256: Uint8Array; } -export function encodeFooter(cli: PayloadSpec, core: PayloadSpec): Uint8Array { - if (cli.sha3_256.length !== 32 || core.sha3_256.length !== 32) { +export function encodeFooter(cli: PayloadSpec, core: PayloadSpec, patchApplier: PayloadSpec): Uint8Array { + if (cli.sha3_256.length !== 32 || core.sha3_256.length !== 32 || patchApplier.sha3_256.length !== 32) { throw new Error("payload SHA3-256 digests must be 32 bytes"); } @@ -21,7 +21,10 @@ export function encodeFooter(cli: PayloadSpec, core: PayloadSpec): Uint8Array { view.setBigUint64(48, core.offset, true); view.setBigUint64(56, core.len, true); out.set(core.sha3_256, 64); - view.setUint32(96, FORMAT_VERSION, true); - out.set(MAGIC, 100); + view.setBigUint64(96, patchApplier.offset, true); + view.setBigUint64(104, patchApplier.len, true); + out.set(patchApplier.sha3_256, 112); + view.setUint32(144, FORMAT_VERSION, true); + out.set(MAGIC, 148); return out; } diff --git a/packages/agenty-bootstrap/scripts/pack.ts b/packages/agenty-bootstrap/scripts/pack.ts index e44e797..e9c20fe 100644 --- a/packages/agenty-bootstrap/scripts/pack.ts +++ b/packages/agenty-bootstrap/scripts/pack.ts @@ -1,7 +1,7 @@ /** * Packs the final self-extracting `agenty--` binary: * - * [ agenty-bootstrap stub ][ compressed CLI ][ compressed core ][ footer ] + * [ bootstrap stub ][ compressed CLI ][ compressed core ][ compressed apply_patch ][ footer ] * * Both payloads are compressed in memory with xz and appended directly, so no * intermediate archives are written to disk. The footer records each payload's @@ -28,6 +28,7 @@ const PKG = resolve(import.meta.dir, ".."); const REPO = resolve(PKG, "../.."); const CORE_BIN_DIR = join(REPO, "packages/agenty-core/bin"); const CLI_DIST_DIR = join(REPO, "packages/agenty-cli/bin"); +const PATCH_APPLIER_TARGET_DIR = join(REPO, "packages/patch-applier/target/release"); const DIST = join(PKG, "bin"); function findAgentyBinary(dir: string): string | null { @@ -96,18 +97,25 @@ if (!existsSync(stubPath)) { } const cliPath = resolveCliBinary(os, arch, ext); const corePath = resolveCoreBinary(os, arch); +const patchApplierPath = join(PATCH_APPLIER_TARGET_DIR, `apply_patch${ext}`); +if (!existsSync(patchApplierPath)) { + throw new Error(`patch-applier binary not found at ${patchApplierPath}; run its build first`); +} const stub = readFileSync(stubPath); const cli = readFileSync(cliPath); const core = readFileSync(corePath); -if (stub.length === 0 || cli.length === 0 || core.length === 0) { - throw new Error("stub, CLI or core binary is empty"); +const patchApplier = readFileSync(patchApplierPath); +if (stub.length === 0 || cli.length === 0 || core.length === 0 || patchApplier.length === 0) { + throw new Error("stub, CLI, core or patch-applier binary is empty"); } const cliSha3 = sha3_256(cli); const coreSha3 = sha3_256(core); +const patchApplierSha3 = sha3_256(patchApplier); const cliPayload = await compress(cli); const corePayload = await compress(core); +const patchApplierPayload = await compress(patchApplier); const footer = encodeFooter( { @@ -120,11 +128,16 @@ const footer = encodeFooter( len: BigInt(corePayload.length), sha3_256: coreSha3, }, + { + offset: BigInt(stub.length + cliPayload.length + corePayload.length), + len: BigInt(patchApplierPayload.length), + sha3_256: patchApplierSha3, + }, ); mkdirSync(DIST, { recursive: true }); const out = join(DIST, `agenty-${os}-${arch}${ext}`); -writeFileSync(out, Buffer.concat([stub, cliPayload, corePayload, footer])); +writeFileSync(out, Buffer.concat([stub, cliPayload, corePayload, patchApplierPayload, footer])); if (os !== "windows") { chmodSync(out, 0o755); } @@ -134,5 +147,6 @@ console.log( ` stub ${stub.length} bytes (${stubPath})\n` + ` cli ${cli.length} -> ${cliPayload.length} bytes (${cliPath})\n` + ` core ${core.length} -> ${corePayload.length} bytes (${corePath})\n` + - ` total ${stub.length + cliPayload.length + corePayload.length + footer.length} bytes`, + ` patch ${patchApplier.length} -> ${patchApplierPayload.length} bytes (${patchApplierPath})\n` + + ` total ${stub.length + cliPayload.length + corePayload.length + patchApplierPayload.length + footer.length} bytes`, ); diff --git a/packages/agenty-bootstrap/src/lib.rs b/packages/agenty-bootstrap/src/lib.rs index 243082e..6b32cb6 100644 --- a/packages/agenty-bootstrap/src/lib.rs +++ b/packages/agenty-bootstrap/src/lib.rs @@ -9,9 +9,9 @@ use sha3::{Digest, Sha3_256}; pub const MAGIC: [u8; 8] = [0xca, 0xfe, 0xba, 0xbe, 0x10, 0x13, 0x66, 0x66]; -pub const FORMAT_VERSION: u32 = 1; +pub const FORMAT_VERSION: u32 = 2; -pub const FOOTER_SIZE: usize = 108; +pub const FOOTER_SIZE: usize = 156; const COPY_BUFFER_SIZE: usize = 64 * 1024; @@ -26,6 +26,7 @@ pub struct PayloadSpec { pub struct Footer { pub cli: PayloadSpec, pub core: PayloadSpec, + pub patch_applier: PayloadSpec, } impl Footer { @@ -37,20 +38,23 @@ impl Footer { out[48..56].copy_from_slice(&self.core.offset.to_le_bytes()); out[56..64].copy_from_slice(&self.core.len.to_le_bytes()); out[64..96].copy_from_slice(&self.core.sha3_256); - out[96..100].copy_from_slice(&FORMAT_VERSION.to_le_bytes()); - out[100..108].copy_from_slice(&MAGIC); + out[96..104].copy_from_slice(&self.patch_applier.offset.to_le_bytes()); + out[104..112].copy_from_slice(&self.patch_applier.len.to_le_bytes()); + out[112..144].copy_from_slice(&self.patch_applier.sha3_256); + out[144..148].copy_from_slice(&FORMAT_VERSION.to_le_bytes()); + out[148..156].copy_from_slice(&MAGIC); out } pub fn decode(bytes: &[u8; FOOTER_SIZE]) -> Result { - if bytes[100..108] != MAGIC { + if bytes[148..156] != MAGIC { return Err(BootstrapError::CorruptFooter( "magic trailer not found; this binary carries no payloads".to_string(), )); } let mut version = [0u8; 4]; - version.copy_from_slice(&bytes[96..100]); + version.copy_from_slice(&bytes[144..148]); if u32::from_le_bytes(version) != FORMAT_VERSION { return Err(BootstrapError::CorruptFooter(format!( "unsupported footer format version {}", @@ -67,6 +71,8 @@ impl Footer { cli_sha.copy_from_slice(&bytes[16..48]); let mut core_sha = [0u8; 32]; core_sha.copy_from_slice(&bytes[64..96]); + let mut patch_applier_sha = [0u8; 32]; + patch_applier_sha.copy_from_slice(&bytes[112..144]); Ok(Footer { cli: PayloadSpec { @@ -79,6 +85,11 @@ impl Footer { len: read_u64(56), sha3_256: core_sha, }, + patch_applier: PayloadSpec { + offset: read_u64(96), + len: read_u64(104), + sha3_256: patch_applier_sha, + }, }) } } @@ -213,13 +224,21 @@ pub fn managed_bin_dir(home: &Path) -> PathBuf { home.join(".agenty").join("bin") } -pub fn artifact_paths(home: &Path) -> (PathBuf, PathBuf) { +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ArtifactPaths { + pub cli: PathBuf, + pub core: PathBuf, + pub patch_applier: PathBuf, +} + +pub fn artifact_paths(home: &Path) -> ArtifactPaths { let dir = managed_bin_dir(home); let ext = if cfg!(windows) { ".exe" } else { "" }; - ( - dir.join(format!("cli{ext}")), - dir.join(format!("core{ext}")), - ) + ArtifactPaths { + cli: dir.join(format!("cli{ext}")), + core: dir.join(format!("core{ext}")), + patch_applier: dir.join(format!("apply_patch{ext}")), + } } fn temp_path_for(target: &Path) -> PathBuf { @@ -303,7 +322,9 @@ mod tests { const GOLDEN_FOOTER_HEX: &str = "88776655443322110807060504030201000102030405060708090a0b0c0d0e0f\ 101112131415161718191a1b1c1d1e1f1122334455667788010203040506070820212223242526\ - 2728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f01000000cafebabe10136666"; + 2728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f08090a0b0c0d0e0f1011121314\ + 151617404142434445464748494a4b4c4d4e4f505152535455565758595a5b5c5d5e5f02000000\ + cafebabe10136666"; const INTEROP_XZ_HEX: &str = "fd377a585a000004e6d6b4460200210116000000742fe5a3e0087f00355d003099c8db4efc244eb58cf58f4699c115ba2fbad7ad9c231199c49368b315728a5421d1340068b4b68fd6e65bef9dbfedfd1f52190000000000154019c351bd7ee70001518011000000e78fc45db1c467fb020000000004595a"; const INTEROP_RAW_SHA3: &str = @@ -312,7 +333,7 @@ mod tests { fn unhex(s: &str) -> Vec { let s: String = s.chars().filter(|c| !c.is_whitespace()).collect(); - assert!(s.len() % 2 == 0, "odd-length hex string"); + assert!(s.len().is_multiple_of(2), "odd-length hex string"); (0..s.len()) .step_by(2) .map(|i| u8::from_str_radix(&s[i..i + 2], 16).expect("invalid hex")) @@ -332,6 +353,10 @@ mod tests { for (i, b) in core_sha.iter_mut().enumerate() { *b = 0x20 + i as u8; } + let mut patch_applier_sha = [0u8; 32]; + for (i, b) in patch_applier_sha.iter_mut().enumerate() { + *b = 0x40 + i as u8; + } Footer { cli: PayloadSpec { offset: 0x1122334455667788, @@ -343,6 +368,11 @@ mod tests { len: 0x0807060504030201, sha3_256: core_sha, }, + patch_applier: PayloadSpec { + offset: 0x0f0e0d0c0b0a0908, + len: 0x1716151413121110, + sha3_256: patch_applier_sha, + }, } } @@ -377,6 +407,7 @@ mod tests { let footer = Footer { cli: spec.clone(), core: spec.clone(), + patch_applier: spec.clone(), }; let path = dir.join("packed"); @@ -411,7 +442,7 @@ mod tests { #[test] fn footer_rejects_bad_magic() { let mut bytes = golden_footer().encode(); - bytes[107] ^= 0xff; + bytes[155] ^= 0xff; let err = Footer::decode(&bytes).unwrap_err(); assert!(matches!(err, BootstrapError::CorruptFooter(_))); } @@ -419,7 +450,7 @@ mod tests { #[test] fn footer_rejects_unknown_version() { let mut bytes = golden_footer().encode(); - bytes[96] = 0x7f; + bytes[144] = 0x7f; let err = Footer::decode(&bytes).unwrap_err(); assert!(matches!(err, BootstrapError::CorruptFooter(_))); } @@ -559,14 +590,25 @@ mod tests { let footer = read_footer(&mut packed).unwrap(); assert_eq!(footer.cli, spec); assert_eq!(footer.core, spec); + assert_eq!(footer.patch_applier, spec); } #[test] fn artifact_paths_use_agenty_bin_dir() { let home = Path::new("/home/tester"); - let (cli, core) = artifact_paths(home); + let paths = artifact_paths(home); let ext = if cfg!(windows) { ".exe" } else { "" }; - assert_eq!(cli, home.join(".agenty/bin").join(format!("cli{ext}"))); - assert_eq!(core, home.join(".agenty/bin").join(format!("core{ext}"))); + assert_eq!( + paths.cli, + home.join(".agenty/bin").join(format!("cli{ext}")) + ); + assert_eq!( + paths.core, + home.join(".agenty/bin").join(format!("core{ext}")) + ); + assert_eq!( + paths.patch_applier, + home.join(".agenty/bin").join(format!("apply_patch{ext}")) + ); } } diff --git a/packages/agenty-bootstrap/src/main.rs b/packages/agenty-bootstrap/src/main.rs index d02f364..271e7d5 100644 --- a/packages/agenty-bootstrap/src/main.rs +++ b/packages/agenty-bootstrap/src/main.rs @@ -34,20 +34,28 @@ fn bootstrap() -> Result { let home = dirs::home_dir().ok_or_else(|| { BootstrapError::Invalid("cannot locate the current user's home directory".to_string()) })?; - let (cli, core) = artifact_paths(&home); + let artifacts = artifact_paths(&home); - if !cli.is_file() && !core.is_file() { + if !artifacts.cli.is_file() && !artifacts.core.is_file() && !artifacts.patch_applier.is_file() { progress.parent("local binary not found, extracting..."); - install_artifact(&mut file, &footer.cli, &cli)?; - install_artifact(&mut file, &footer.core, &core)?; + install_artifact(&mut file, &footer.cli, &artifacts.cli)?; + install_artifact(&mut file, &footer.core, &artifacts.core)?; + install_artifact(&mut file, &footer.patch_applier, &artifacts.patch_applier)?; } else { progress.parent("checking local binary integrity..."); - ensure_with_progress(&mut file, &footer.cli, &cli, "cli", &progress)?; - ensure_with_progress(&mut file, &footer.core, &core, "core", &progress)?; + ensure_with_progress(&mut file, &footer.cli, &artifacts.cli, "cli", &progress)?; + ensure_with_progress(&mut file, &footer.core, &artifacts.core, "core", &progress)?; + ensure_with_progress( + &mut file, + &footer.patch_applier, + &artifacts.patch_applier, + "apply_patch", + &progress, + )?; } progress.finish(); - launch(&cli) + launch(&artifacts.cli) } fn agenty_version() -> &'static str { diff --git a/packages/agenty-cli/src/localCore.test.ts b/packages/agenty-cli/src/localCore.test.ts index 26c2445..a243f68 100644 --- a/packages/agenty-cli/src/localCore.test.ts +++ b/packages/agenty-cli/src/localCore.test.ts @@ -1,6 +1,8 @@ +import { delimiter } from "node:path"; + import { describe, expect, test } from "bun:test"; -import { pickCorePath } from "./localCore"; +import { pickCorePath, prependCoreDirectoryToPath } from "./localCore"; const candidates = { repoBin: "/repo/packages/agenty-core/bin/agenty-core", @@ -24,3 +26,11 @@ describe("pickCorePath", () => { expect(pickCorePath(candidates, () => false)).toBeNull(); }); }); + +describe("prependCoreDirectoryToPath", () => { + test("places the core directory before the inherited PATH", () => { + expect(prependCoreDirectoryToPath("/managed/bin/core", "/usr/bin")).toBe( + ["/managed/bin", "/usr/bin"].join(delimiter), + ); + }); +}); diff --git a/packages/agenty-cli/src/localCore.ts b/packages/agenty-cli/src/localCore.ts index 347a3a2..0dfbbf8 100644 --- a/packages/agenty-cli/src/localCore.ts +++ b/packages/agenty-cli/src/localCore.ts @@ -1,6 +1,6 @@ import { existsSync } from "node:fs"; import { homedir } from "node:os"; -import { join, resolve } from "node:path"; +import { delimiter, dirname, join, resolve } from "node:path"; import { StdioRPCClient } from "./core/rpc"; @@ -38,6 +38,10 @@ export function pickCorePath( return null; } +export function prependCoreDirectoryToPath(binary: string, currentPath?: string): string { + return [dirname(binary), currentPath].filter(Boolean).join(delimiter); +} + export interface LocalCore { rpc: StdioRPCClient; stop: () => Promise; @@ -61,6 +65,7 @@ export async function startLocalCore(options: { dataDir?: string } = {}): Promis stderr: "pipe", env: { ...process.env, + PATH: prependCoreDirectoryToPath(binary, process.env.PATH), ...(options.dataDir ? { AGENTY_DATA_DIR: options.dataDir } : {}), }, }); diff --git a/packages/agenty-core/README-CN.md b/packages/agenty-core/README-CN.md index a28bde0..decd515 100644 --- a/packages/agenty-core/README-CN.md +++ b/packages/agenty-core/README-CN.md @@ -14,7 +14,7 @@ Agenty 的核心运行时。它围绕本地优先的存储模型(文件系统 | Session transcript | `~/.agenty/sessions///
/.jsonl` | 写模型,即 append-only event log(真实数据源) | | Session index | `~/.agenty/agenty.sqlite` -> `sessions` | 读模型,用于快速列表和搜索的投影 | | 全局配置 | `~/.agenty/config.json` | 应用配置 | -| Providers | `~/.agenty/providers/.json` | Catalog aggregate,包含其模型 | +| Providers | 内置 catalog 固化在 core;自定义 provider 使用 `~/.agenty/providers/.json` | 内置元数据/模型只读,内置 provider 文件仅保存 API key | | Agents | `~/.agenty/agents/.json` | Agent aggregate | | Core 日志 | `~/.agenty/logs///
/core.log` | 结构化文本诊断信息(JSONL 模式下为 `core.jsonl`) | @@ -45,24 +45,18 @@ reasoning effort 和工作目录。 ### Reasoning effort Agenty 对外提供且仅提供六个与 provider 无关的 reasoning effort 等级:`off`、`low`、 -`medium`、`high`、`xhigh` 和 `max`。模型通过 `reasoningEffortMapping` 对象保存映射, -其中 key 是 provider 原生 effort 名称,value 是 Agenty effort 等级: +`medium`、`high`、`xhigh` 和 `max`。reasoning 模型通过 `reasoningEfforts` 数组保存 +实际支持的启用等级: ```json { - "reasoningEffortMapping": { - "none": "off", - "minimal": "low", - "low": "low", - "medium": "medium", - "high": "high" - } + "reasoningEfforts": ["low", "medium", "high", "xhigh", "max"] } ``` -该映射允许多个原生 effort 归一化到同一个 Agenty effort。映射中没有启用任何 effort -的模型不支持 reasoning。只有上述六个 Agenty 等级可以作为映射值;原生 effort 名称 -由 provider 自行定义。 +显式空数组表示非 reasoning 模型;上游未返回 capability 数据时默认使用五个启用等级。 +provider adapter 会原样发送选择的等级,不支持的等级通过正常 round error 流程返回。 +`minimal` 等 provider 特有等级不对外开放。 ## Agent loop 运行时 @@ -73,8 +67,8 @@ session 的 `Engine`。不同 session 可以并行执行,同一 session 只允 每次 loop 会解析 Agent system prompt,重建有效会话上下文,通过选定 provider adapter 转换上游数据结构,调用 LLM、持久化 assistant 响应,并在返回 tool calls 时继续循环。 -所有 model 统一使用 `8192` 最大输出 token;旧的 model 级字段仅为兼容旧 wire -客户端而保留,实际值会被忽略。估算上下文达到 +自定义 model 未填写时默认使用 `8192` 最大输出 token;内置 model 使用嵌入 catalog +中的精确限制。估算上下文达到 `contextWindow` 的 `90%` 时自动压缩;TUI 的 `/compact` 可以手动触发同一流程。压缩会在 `session_compacted` 事件中只保存生成的总结和压缩审计数据; 重放和构造模型请求时,再从 transcript 动态计算最多三条最近 user 消息、总结、metadata @@ -85,9 +79,11 @@ session 的 `Engine`。不同 session 可以并行执行,同一 session 只允 窗口的 90%,先使用当前 model 压缩,必要时裁剪保留消息以适配目标窗口,再写入 model 切换事件。 共享 tool registry 实现 `ToolRuntime` port;同一批次内每个 tool call 并行执行,结果按 -调用顺序返回。`pkg/agentloop/builtin/` 提供生产环境文件系统工具 `read_file`、 -`write_file`、`patch_file`、`delete_file`、`grep`、`glob` 和 `ls`,由 `cmd/main.go` -显式注册。相对路径基于该 round 捕获的 session 工作目录解析,绝对路径保持有效。 +调用顺序返回。`pkg/agentloop/builtin/` 提供 `read_file`、`apply_patch`、`grep`、`glob` +和 `ls`,由 `cmd/main.go` 显式注册。`apply_patch` 会调用同名 Rust 可执行文件完成 V4A +解析和原子化文件修改。支持 free-form tool 的 provider 会收到模型工具定义;其他 provider +会在 system prompt 中收到通过 `shell` 执行同一命令的说明。相对路径基于该 round 捕获的 +session 工作目录解析。 ## 基础设施层 @@ -97,12 +93,13 @@ session 的 `Engine`。不同 session 可以并行执行,同一 session 只允 pkg/infra/ ├── config/ 将配置文件和 env override 合并到单例中;解析 data-dir 路径 ├── initialize/ OpenRepositories:一次性初始化所有 stores +├── catalogdata/ 内嵌的 provider/model JSON ├── llm/ 实现 agentloop caller contract 的 provider SDK adapters ├── logging/ slog 初始化、环境配置解析和按日生成日志路径 ├── storage/ Repository 实现 + SQLite connection factory │ ├── db.go OpenDB/OpenIsolatedDB + sessions schema │ ├── agent.go AgentRepository(agent JSON 文件) -│ ├── catalog.go CatalogRepository(provider 聚合 JSON,内嵌 models) +│ ├── catalog.go CatalogRepository(内置 provider 与自定义 provider JSON) │ └── conversation.go ConversationRepository(JSONL transcript + SQLite projection) └── rpc/ stdio JSON-RPC 2.0 接口层 ├── message.go Request/Response/Notification/Error/ID wire types @@ -174,10 +171,18 @@ Methods 使用 `resource.action` 命名: | --- | --- | | Initialize | `initialize.already`, `initialize.complete` | | Agent | `agent.create`, `agent.get`, `agent.list`, `agent.update`, `agent.delete` | -| Provider | `provider.create`, `provider.get`, `provider.list`, `provider.update`, `provider.delete`, `provider.addModel`, `provider.removeModel` | +| Provider | `provider.create`, `provider.get`, `provider.list`, `provider.listModels`, `provider.update`, `provider.delete`, `provider.addModel`, `provider.removeModel` | | Session | `session.create`, `session.get`, `session.list`, `session.delete`, `session.setTitle`, `session.setModel`, `session.setReasoningEffort`, `session.setCwd`, `session.start`, `session.compact`, `session.stop` | | Chunk | `chunk.begin`, `chunk.part`, `chunk.commit`, `chunk.abort` | +`provider.list` 可选接收 `{providerCode}`。不传时,core 会并行获取所有已配置且 catalog +为空的 provider;传入时只会获取指定 provider。`provider.listModels` 接收同样的 +`{providerCode}`,供直接调用者使用 core 内置的发现流程。成功结果会缓存到 +`~/.agenty/providers/.models/.json`,有效期为 8 小时,JSON 中保存 `expiresAt` +和标准化模型列表。过期缓存仍作为旧数据返回,下一次 list 时再刷新。它会兼容常见的 `id`、 +名称和 token 限制字段,自动跟随 provider 分页;上下文窗口或最大输出 token 缺失或不为正数时 +分别使用 `256000` 和 `65536`,缺少 reasoning 能力时返回空的 `reasoningEfforts` 数组。 + `session.start` 接收 `{id, content}`,持久化 running round 后立即返回 round 标识和 `running` 状态,完整 agent turn 由引擎异步继续执行。执行期间,core 会写出 `session.event` JSON-RPC notifications,事件类型包括 `round_started`、 diff --git a/packages/agenty-core/README.md b/packages/agenty-core/README.md index 26b5716..bea45b2 100644 --- a/packages/agenty-core/README.md +++ b/packages/agenty-core/README.md @@ -14,7 +14,7 @@ The filesystem is the source of truth; SQLite is a query-side projection. | Session transcript | `~/.agenty/sessions///
/.jsonl` | Write model — append-only event log (source of truth) | | Session index | `~/.agenty/agenty.sqlite` → `sessions` | Read model — projection for fast listing/search | | Global config | `~/.agenty/config.json` | Application configuration | -| Providers | `~/.agenty/providers/.json` | Catalog aggregate, including its models | +| Providers | Embedded catalog; custom providers use `~/.agenty/providers/.json` | Built-in metadata/models are read-only; built-in files store only API keys | | Agents | `~/.agenty/agents/.json` | Agent aggregate | | Core log | `~/.agenty/logs///
/core.log` | Structured text diagnostics (`core.jsonl` in JSONL mode) | @@ -47,24 +47,19 @@ model, context window, reasoning effort, and working directory used by that roun ### Reasoning effort Agenty exposes exactly six provider-independent reasoning effort levels: `off`, `low`, -`medium`, `high`, `xhigh`, and `max`. A model stores a `reasoningEffortMapping` object -whose keys are provider-native effort names and whose values are Agenty effort levels: +`medium`, `high`, `xhigh`, and `max`. Reasoning models store supported enabled levels +in a `reasoningEfforts` array: ```json { - "reasoningEffortMapping": { - "none": "off", - "minimal": "low", - "low": "low", - "medium": "medium", - "high": "high" - } + "reasoningEfforts": ["low", "medium", "high", "xhigh", "max"] } ``` -The mapping allows multiple native efforts to normalize to the same Agenty effort. -A model whose mapping has no enabled effort does not support reasoning. Only the six -Agenty levels above are valid mapping values; native effort names are provider-specific. +An explicit empty array identifies a non-reasoning model. Missing upstream capability +data defaults to all five enabled Agenty levels. Provider adapters send the selected level +unchanged; unsupported levels are reported through the normal round error flow. +Provider-specific levels such as `minimal` are not exposed. ## Agent-loop runtime @@ -76,9 +71,9 @@ active rounds. Each loop resolves the Agent system prompt, rebuilds the effective conversation context, converts it through the selected provider adapter, invokes the LLM, persists the -assistant response, and repeats when tool calls are returned. Every model invocation uses -the global `8192` output-token limit; the legacy per-model field is ignored and retained -only for wire compatibility. Automatic compaction runs when the estimated context reaches +assistant response, and repeats when tool calls are returned. Custom models use `8192` output +tokens when omitted; built-in models use the exact limit from +the embedded catalog. Automatic compaction runs when the estimated context reaches `contextWindow * 90%`. `/compact` triggers the same flow manually. Compaction stores only the generated summary and compaction audit data in a `session_compacted` event. During replay and request construction, the effective model @@ -92,9 +87,11 @@ compacts with the current model, trims retained context to fit the target when n then persists the model change. The loop currently permits at most 20 LLM/tool iterations. The shared registry implements the `ToolRuntime` port, executes one tool batch concurrently, and returns results in call order. `pkg/agentloop/builtin/` provides -the production filesystem tools `read_file`, `write_file`, `patch_file`, `delete_file`, -`grep`, `glob`, and `ls`; `cmd/main.go` registers them explicitly. Relative paths resolve -from the round's captured session working directory, while absolute paths remain valid. +`read_file`, `apply_patch`, `grep`, `glob`, and `ls`; `cmd/main.go` registers them explicitly. +`apply_patch` delegates V4A parsing and atomic filesystem mutation to the bundled Rust +executable of the same name. Providers with free-form tool support receive `apply_patch` +as a model tool. Other providers receive a system instruction to run the same executable +through `shell`. Relative paths resolve from the round's captured session working directory. ## Infrastructure layer @@ -105,12 +102,13 @@ filesystem + SQLite storage model. pkg/infra/ ├── config/ Load config file + env overrides into a merged singleton; resolve data-dir paths ├── initialize/ OpenRepositories: one-call setup of all stores +├── catalogdata/ Embedded built-in provider/model JSON ├── llm/ Provider SDK adapters implementing the agentloop caller contract ├── logging/ slog setup, environment parsing, and daily log path ├── storage/ Repository implementations + SQLite connection factory │ ├── db.go OpenDB/OpenIsolatedDB + sessions schema │ ├── agent.go AgentRepository (agent JSON files) -│ ├── catalog.go CatalogRepository (provider aggregate JSON, embedded models) +│ ├── catalog.go CatalogRepository (embedded built-ins plus custom provider JSON) │ └── conversation.go ConversationRepository (JSONL transcript + SQLite projection) └── rpc/ stdio JSON-RPC 2.0 interface layer ├── message.go Request/Response/Notification/Error/ID wire types @@ -194,10 +192,20 @@ Methods follow a `resource.action` naming: | --- | --- | | Initialize | `initialize.already`, `initialize.complete` | | Agent | `agent.create`, `agent.get`, `agent.list`, `agent.update`, `agent.delete` | -| Provider | `provider.create`, `provider.get`, `provider.list`, `provider.update`, `provider.delete`, `provider.addModel`, `provider.removeModel` | +| Provider | `provider.create`, `provider.get`, `provider.list`, `provider.listModels`, `provider.update`, `provider.delete`, `provider.addModel`, `provider.removeModel` | | Session | `session.create`, `session.get`, `session.list`, `session.delete`, `session.setTitle`, `session.setModel`, `session.setReasoningEffort`, `session.setCwd`, `session.start`, `session.compact`, `session.stop` | | Chunk | `chunk.begin`, `chunk.part`, `chunk.commit`, `chunk.abort` | +`provider.list` accepts an optional `{providerCode}`. Without it, core discovers all +configured providers whose catalog is empty in parallel; with it, only that provider is +eligible for discovery. `provider.listModels` accepts `{providerCode}` and exposes the same +core-owned discovery path for direct callers. Successful discovery is cached under +`~/.agenty/providers/.models/.json` for 8 hours; the JSON stores an `expiresAt` +timestamp and the normalized models. Expired entries remain available as stale data while a +subsequent list refreshes them. It maps common `id`/name/token-limit fields, follows provider +pagination, defaults missing or non-positive context/output limits to `256000` and `65536`, +and represents missing reasoning capability as an empty `reasoningEfforts` array. + `session.start` accepts `{id, content}` and returns the persisted round's identifiers and `running` status immediately; the engine continues the full agent turn asynchronously. While it runs, core writes `session.event` JSON-RPC notifications with diff --git a/packages/agenty-core/TESTING-CN.md b/packages/agenty-core/TESTING-CN.md index 71deff9..b75da23 100644 --- a/packages/agenty-core/TESTING-CN.md +++ b/packages/agenty-core/TESTING-CN.md @@ -9,7 +9,7 @@ | --- | --- | --- | --- | | Domain | 仅内存值 | 聚合不变量、Session 状态转换与 replay、event 和 content 序列化、Provider model 生命周期、code 和 reasoning effort 映射校验 | 是 | | Application | 内存 repository fake | Agent、Provider 和 Session 用例;execution loop 完成、tool continuation、model 输出 token 上限、多 session 并行、取消、shutdown、输入校验、错误映射和 pending event 生命周期 | 是 | -| 内置工具 | `t.TempDir()` 和真实文件系统操作 | 注册、相对路径解析、范围读取、创建/覆盖、精确 patch、单文件删除、正则搜索、递归 glob、目录列表、输出限制和错误路径 | 是 | +| 内置工具 | `t.TempDir()`、helper fixture 和真实文件系统操作 | 注册、相对路径解析、范围读取、结构化 `apply_patch` 子进程结果、正则搜索、递归 glob、目录列表、输出限制和错误路径 | 是 | | RPC | buffer、fake handler 和合成时间 | JSON-RPC/NDJSON framing、notification、batch、非法请求、单行限制、chunk 组装与清理 | 是 | | Config、logging 与 storage | `t.TempDir()`、真实文件和本地 SQLite | 配置文件与 env override 合并、单例 Manager、日志等级/格式/路径选择、JSON repository、append-only transcript、SQLite projection 和 schema 初始化 | 是 | | 完整装配 | 隔离的文件系统和 SQLite 状态 | repository 初始化,以及包括异步 session start/stop 在内的 RPC 到 application 再到 storage 完整流程 | 启用 `integration` 时 | diff --git a/packages/agenty-core/TESTING.md b/packages/agenty-core/TESTING.md index c8aa44e..057242c 100644 --- a/packages/agenty-core/TESTING.md +++ b/packages/agenty-core/TESTING.md @@ -9,7 +9,7 @@ Chinese version, see [TESTING-CN.md](./TESTING-CN.md). | --- | --- | --- | --- | | Domain | In-memory values | Aggregate invariants, Session transitions and replay, event and content serialization, Provider model lifecycle, code and reasoning effort mapping validation | Yes | | Application | In-memory repository fakes | Agent, Provider, and Session use cases; execution-loop completion, tool continuation, per-model token limits, multi-session concurrency, cancellation, shutdown, validation, error mapping, and pending-event lifecycle | Yes | -| Built-in tools | `t.TempDir()` and real filesystem operations | Registration, relative path resolution, ranged reads, create/overwrite, exact patching, safe single-file deletion, regular-expression search, recursive globbing, directory listing, output limits, and error paths | Yes | +| Built-in tools | `t.TempDir()`, helper fixtures, and real filesystem operations | Registration, relative path resolution, ranged reads, structured `apply_patch` subprocess results, regular-expression search, recursive globbing, directory listing, output limits, and error paths | Yes | | RPC | Buffers, fake handlers, and synthetic time | JSON-RPC/NDJSON framing, notifications, batches, invalid requests, line limits, chunk assembly, and cleanup | Yes | | Config, logging, and storage | `t.TempDir()`, real files, and local SQLite | Config file + env override merging, singleton Manager, log level/format/path selection, JSON repositories, append-only transcripts, SQLite projections, and schema initialization | Yes | | Complete wiring | Isolated filesystem and SQLite state | Repository initialization and RPC-to-application-to-storage flows, including asynchronous session start/stop | With `integration` | diff --git a/packages/agenty-core/package.json b/packages/agenty-core/package.json index d22dd2f..be4fbc4 100644 --- a/packages/agenty-core/package.json +++ b/packages/agenty-core/package.json @@ -3,7 +3,7 @@ "version": "0.1.0", "private": true, "scripts": { - "build": "mkdir -p \"${PACKAGE_DIR:-bin}\" && go build -o \"${PACKAGE_DIR:-bin}/${BIN_NAME:-agenty-core}\" ./cmd", + "build": "mkdir -p \"${PACKAGE_DIR:-bin}\" && go build -o \"${PACKAGE_DIR:-bin}/${BIN_NAME:-agenty-core}\" ./cmd && if [ \"${GOOS:-}\" = \"windows\" ]; then cp ../patch-applier/target/release/apply_patch.exe \"${PACKAGE_DIR:-bin}/apply_patch.exe\"; else cp ../patch-applier/target/release/apply_patch \"${PACKAGE_DIR:-bin}/apply_patch\"; fi", "test": "go test ./...", "test:integration": "go test -tags=integration ./...", "test:e2e": "go test -tags=e2e -count=1 -parallel=8 ./test/e2e", @@ -13,5 +13,8 @@ "vet": "go vet ./...", "fmt": "go fmt ./...", "clean": "rm -rf bin" + }, + "devDependencies": { + "patch-applier": "workspace:*" } } diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 4671a22..c35a80d 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -24,8 +24,8 @@ importers: specifier: ^14.0.0 version: 14.0.0(eslint@10.8.1) turbo: - specifier: ^2.10.10 - version: 2.10.10 + specifier: ^2.10.11 + version: 2.10.11 typescript-eslint: specifier: ^8.67.0 version: 8.67.0(eslint@10.8.1)(typescript@6.0.3) @@ -42,6 +42,9 @@ importers: agenty-core: specifier: workspace:* version: link:../agenty-core + patch-applier: + specifier: workspace:* + version: link:../patch-applier packages/agenty-cli: dependencies: @@ -77,7 +80,13 @@ importers: specifier: ^6.0.3 version: 6.0.3 - packages/agenty-core: {} + packages/agenty-core: + devDependencies: + patch-applier: + specifier: workspace:* + version: link:../patch-applier + + packages/patch-applier: {} packages: @@ -326,33 +335,33 @@ packages: peerDependencies: eslint: ^9.0.0 || ^10.0.0 - '@turbo/darwin-64@2.10.10': - resolution: {integrity: sha512-gFDD+wRP5hWxBRghGyEbjpbLOY7aIU/wvsnKdMM7odQcp/wHMrnI83p0FyxxMRZnFH9ZD+S59MvcpOC5b+nrCA==} + '@turbo/darwin-64@2.10.11': + resolution: {integrity: sha512-v3R+1R/Ysozyo+p7Ri8MCIbndOvYt3DgPFrGLhrhQHfvyvbxyH3WyJj+A/2JTNmNleuAlh3JUyCV0iSVHIONTA==} cpu: [x64] os: [darwin] - '@turbo/darwin-arm64@2.10.10': - resolution: {integrity: sha512-VZYsxZ6yjyDosUqtiroAVSXPLmx/qBxdHJgIxdMH9RyNmLdOLOWtJnYMnI4qckwCgQMK85G3fu94/xk5+iBCgw==} + '@turbo/darwin-arm64@2.10.11': + resolution: {integrity: sha512-R0a0CvGAeYYsBgPIgFNB3agGXh6qukjduNhFlwVVX1Ss2IdBJLXmgjytNGmo084bLKS0B6UdLRwKhXMHKKaObQ==} cpu: [arm64] os: [darwin] - '@turbo/linux-64@2.10.10': - resolution: {integrity: sha512-lAvW+yEnmsCKMEIwNugjozawvYytHKPhU0kfLBizu83MIs8OUb9KobYvkZ56L5akSM6K7+gBFLEIfQkaceh90g==} + '@turbo/linux-64@2.10.11': + resolution: {integrity: sha512-dGlY2vg7jpsLjGS1bf9sD/cw1sGMAYbeHGQKGRFxd5Zaj+Ufbj8cNXl/vIit8ueCU97zhL/znn60ZV5iSfZAxw==} cpu: [x64] os: [android, linux] - '@turbo/linux-arm64@2.10.10': - resolution: {integrity: sha512-MSJ+NkRTd79Z9+YEZpUV9VOWVOOigFhE+v/ETNYJEuTJp3r00y9YgFvDXrmM+DP8Kal6tk3U6xSugD2/Ojh+Jg==} + '@turbo/linux-arm64@2.10.11': + resolution: {integrity: sha512-eSP9+jjsSCBs2x0QpJZlp49dSD1EIMwXH7PsMIhdWk7w6IThhuKJ+jL95JQ2IW7tRjRmdh8gKcKQtrHLD5R5lA==} cpu: [arm64] os: [android, linux] - '@turbo/windows-64@2.10.10': - resolution: {integrity: sha512-ycWpXDkUfnDFDY9d+4Qna/UZotDB0wj+s9agrlmNt0Q7a3XHORhK8GPKJdzgzeutXu9EW5P/jyabTEHlohuDXw==} + '@turbo/windows-64@2.10.11': + resolution: {integrity: sha512-4aD7edogJ8arK8DOyvjTI2KKwmchD5M/WM4BZgidrtFf7aSgn8Ce2rJnDwmriw980c2cKlHnhXtlGy38VtKB8g==} cpu: [x64] os: [win32] - '@turbo/windows-arm64@2.10.10': - resolution: {integrity: sha512-PMk6zQN0csUFklLe+1hz/5G9uU1YmV0cEIey2R/bSeA6o69qcBTlN4A3jOqkgenOO5dOpMHmq2sEUZo8r1+Ssg==} + '@turbo/windows-arm64@2.10.11': + resolution: {integrity: sha512-m8tJkIrTrbQ9O1uHxV0GUq623Zg1678xGna8eMsy4KbygUX/wVFgdZDYV/AxyoyaW/U7nyKoBvaDVwm22p9xqA==} cpu: [arm64] os: [win32] @@ -779,8 +788,8 @@ packages: tslib@2.8.1: resolution: {integrity: sha512-oJFu94HQb+KVduSUQL7wnpmqnfmLsOA/nAh6b6EH0wCEoK0/mPeXU6c3wKDV83MkOuHPRHtSXKKU99IBazS/2w==} - turbo@2.10.10: - resolution: {integrity: sha512-/90KTW+USzvYOPmafRZHVKLBsHXQ5810Ao/HdtJYAqguIhZ+XruS6eIUjqJUDtrSxaZYynNFht68qckGKAOWTA==} + turbo@2.10.11: + resolution: {integrity: sha512-yQfwQVoRXwOuyX1LxiJFBFNg6VfuYh+/RyZLd82+isgyLkBXw3S5XRRzvcck1FAjSCG5sVyLd+O1eDMvYa3J7g==} hasBin: true type-check@0.4.0: @@ -1087,22 +1096,22 @@ snapshots: estraverse: 5.3.0 picomatch: 4.0.5 - '@turbo/darwin-64@2.10.10': + '@turbo/darwin-64@2.10.11': optional: true - '@turbo/darwin-arm64@2.10.10': + '@turbo/darwin-arm64@2.10.11': optional: true - '@turbo/linux-64@2.10.10': + '@turbo/linux-64@2.10.11': optional: true - '@turbo/linux-arm64@2.10.10': + '@turbo/linux-arm64@2.10.11': optional: true - '@turbo/windows-64@2.10.10': + '@turbo/windows-64@2.10.11': optional: true - '@turbo/windows-arm64@2.10.10': + '@turbo/windows-arm64@2.10.11': optional: true '@tybys/wasm-util@0.10.3': @@ -1536,14 +1545,14 @@ snapshots: tslib@2.8.1: optional: true - turbo@2.10.10: + turbo@2.10.11: optionalDependencies: - '@turbo/darwin-64': 2.10.10 - '@turbo/darwin-arm64': 2.10.10 - '@turbo/linux-64': 2.10.10 - '@turbo/linux-arm64': 2.10.10 - '@turbo/windows-64': 2.10.10 - '@turbo/windows-arm64': 2.10.10 + '@turbo/darwin-64': 2.10.11 + '@turbo/darwin-arm64': 2.10.11 + '@turbo/linux-64': 2.10.11 + '@turbo/linux-arm64': 2.10.11 + '@turbo/windows-64': 2.10.11 + '@turbo/windows-arm64': 2.10.11 type-check@0.4.0: dependencies: diff --git a/pnpm-workspace.yaml b/pnpm-workspace.yaml index a6b261d..56984a3 100644 --- a/pnpm-workspace.yaml +++ b/pnpm-workspace.yaml @@ -2,6 +2,7 @@ packages: - "packages/agenty-core" - "packages/agenty-cli" - "packages/agenty-bootstrap" + - "packages/patch-applier" peerDependencyRules: allowedVersions: diff --git a/turbo.json b/turbo.json index a752cea..7816e69 100644 --- a/turbo.json +++ b/turbo.json @@ -1,5 +1,5 @@ { - "$schema": "https://v2-10-10.turborepo.dev/schema.json", + "$schema": "https://v2-10-11.turborepo.dev/schema.json", "tasks": { "build": { "dependsOn": ["^build"], @@ -18,9 +18,17 @@ "passThroughEnv": ["HOME", "CARGO_HOME", "RUSTUP_HOME"], "outputs": ["dist/**"] }, + "patch-applier#build": { + "cache": false, + "passThroughEnv": ["HOME", "CARGO_HOME", "RUSTUP_HOME"], + "outputs": ["target/release/apply_patch", "target/release/apply_patch.exe"] + }, "agenty-bootstrap#test": { "passThroughEnv": ["HOME", "CARGO_HOME", "RUSTUP_HOME"] }, + "patch-applier#test": { + "passThroughEnv": ["HOME", "CARGO_HOME", "RUSTUP_HOME"] + }, "test": { "dependsOn": ["^build"], "env": ["GOCACHE"] From 10d3779b5e406a4dda1257bf9946ca1028825f0c Mon Sep 17 00:00:00 2001 From: masteryyh Date: Tue, 25 Aug 2026 10:58:56 +0800 Subject: [PATCH 05/12] chore: update dependencies and go 1.27 Signed-off-by: masteryyh --- packages/agenty-core/go.mod | 18 ++++++++--------- packages/agenty-core/go.sum | 40 ++++++++++++++++++------------------- 2 files changed, 29 insertions(+), 29 deletions(-) diff --git a/packages/agenty-core/go.mod b/packages/agenty-core/go.mod index 969c548..95e9661 100644 --- a/packages/agenty-core/go.mod +++ b/packages/agenty-core/go.mod @@ -1,21 +1,21 @@ module github.com/masteryyh/agenty-core -go 1.26 +go 1.27 require ( - github.com/anthropics/anthropic-sdk-go v1.63.1 + github.com/anthropics/anthropic-sdk-go v1.66.0 github.com/bytedance/sonic v1.15.2 github.com/google/uuid v1.6.0 - github.com/mattn/go-sqlite3 v1.14.49 - github.com/openai/openai-go/v3 v3.51.0 + github.com/mattn/go-sqlite3 v1.14.50 + github.com/openai/openai-go/v3 v3.52.0 github.com/spf13/viper v1.21.0 golang.org/x/sys v0.47.0 - google.golang.org/genai v1.68.0 + google.golang.org/genai v1.69.0 ) require ( cloud.google.com/go v0.123.0 // indirect - cloud.google.com/go/auth v0.23.1 // indirect + cloud.google.com/go/auth v0.23.2 // indirect cloud.google.com/go/compute/metadata v0.9.0 // indirect github.com/bahlo/generic-list-go v0.2.0 // indirect github.com/buger/jsonparser v1.6.1 // indirect @@ -31,7 +31,7 @@ require ( github.com/google/go-cmp v0.7.0 // indirect github.com/google/s2a-go v0.1.9 // indirect github.com/googleapis/enterprise-certificate-proxy v0.3.21 // indirect - github.com/googleapis/gax-go/v2 v2.23.0 // indirect + github.com/googleapis/gax-go/v2 v2.24.0 // indirect github.com/gorilla/websocket v1.5.3 // indirect github.com/invopop/jsonschema v0.14.0 // indirect github.com/klauspost/cpuid/v2 v2.4.0 // indirect @@ -61,7 +61,7 @@ require ( golang.org/x/sync v0.22.0 // indirect golang.org/x/text v0.41.0 // indirect google.golang.org/api v0.293.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260810153831-ec0a7760b754 // indirect - google.golang.org/grpc v1.83.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688 // indirect + google.golang.org/grpc v1.83.1 // indirect google.golang.org/protobuf v1.36.12 // indirect ) diff --git a/packages/agenty-core/go.sum b/packages/agenty-core/go.sum index 88f1e16..7b8c412 100644 --- a/packages/agenty-core/go.sum +++ b/packages/agenty-core/go.sum @@ -1,11 +1,11 @@ cloud.google.com/go v0.123.0 h1:2NAUJwPR47q+E35uaJeYoNhuNEM9kM8SjgRgdeOJUSE= cloud.google.com/go v0.123.0/go.mod h1:xBoMV08QcqUGuPW65Qfm1o9Y4zKZBpGS+7bImXLTAZU= -cloud.google.com/go/auth v0.23.1 h1:1tPpBPG02lQHmoiAvs9egyCASqXP0xgobptjZzov/Jg= -cloud.google.com/go/auth v0.23.1/go.mod h1:4DhBRcqvtljQN3dJ57qtqbib5ZGCYE5f2crfiiC2EM0= +cloud.google.com/go/auth v0.23.2 h1:pxSCpfiji41hpzpPdMCftEUCezpgpqmmDdYiAjCKXxo= +cloud.google.com/go/auth v0.23.2/go.mod h1:4DhBRcqvtljQN3dJ57qtqbib5ZGCYE5f2crfiiC2EM0= cloud.google.com/go/compute/metadata v0.9.0 h1:pDUj4QMoPejqq20dK0Pg2N4yG9zIkYGdBtwLoEkH9Zs= cloud.google.com/go/compute/metadata v0.9.0/go.mod h1:E0bWwX5wTnLPedCKqk3pJmVgCBSM6qQI1yTBdEb3C10= -github.com/anthropics/anthropic-sdk-go v1.63.1 h1:M9dIoZzUWB453ulVpBkQp3FX0769udzrTscDrsJk1w4= -github.com/anthropics/anthropic-sdk-go v1.63.1/go.mod h1:3EfIfmFqxH6rbiLcIP4tPFyXL/IHakx2wDG4OU+TIEI= +github.com/anthropics/anthropic-sdk-go v1.66.0 h1:/CKwgscn0Pe1q4U8aFInSOt/v06JeMc9Aq4vIlctCFw= +github.com/anthropics/anthropic-sdk-go v1.66.0/go.mod h1:3EfIfmFqxH6rbiLcIP4tPFyXL/IHakx2wDG4OU+TIEI= github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk= github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg= github.com/buger/jsonparser v1.6.1 h1:I0phFv0PlbLHnM7TZAVjZ2MJ2/eWRTDyuO7GLR98IEs= @@ -49,8 +49,8 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/googleapis/enterprise-certificate-proxy v0.3.21 h1:OFdQ3tnCX/zaQ0Cedur3D3z7kI6HiLX9g3TiAN4/DFU= github.com/googleapis/enterprise-certificate-proxy v0.3.21/go.mod h1:L3D/IQExI6LqEjBdXcZQ1WluSgigQmSwBboFstVPM4w= -github.com/googleapis/gax-go/v2 v2.23.0 h1:Tchl7qkvE7Ip3y+ztvNufYFvkfqTe7NfLTYGIdJRLuE= -github.com/googleapis/gax-go/v2 v2.23.0/go.mod h1:rBQKOVJCdb8IFEzg+FCwlt1LP/xMDGuqUXhUG+XMXEg= +github.com/googleapis/gax-go/v2 v2.24.0 h1:myMaPYyF9MecEmvQqMqomIwn9t/4KCZN9qnwsS76wlg= +github.com/googleapis/gax-go/v2 v2.24.0/go.mod h1:IaTHBDd7NHxSCiu0vEs8pQZu4dGZrWwuSoxCnk16OFM= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/invopop/jsonschema v0.14.0 h1:MHQqLhvpNUZfw+hM3AZDYK7jxO8FZoQeQM77g8iyZjg= @@ -61,10 +61,10 @@ github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= -github.com/mattn/go-sqlite3 v1.14.49 h1:B8jBHC3xhxZgxztrgruTuLucebnULQnx4W7cF7SAE9w= -github.com/mattn/go-sqlite3 v1.14.49/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= -github.com/openai/openai-go/v3 v3.51.0 h1:+ys88LqUflSr0nRM37aWxkMMpHn+zqzVJGIK89eumdM= -github.com/openai/openai-go/v3 v3.51.0/go.mod h1:Vy3y2/I2H/MbqvJGXEK8VbN5+avZV6zxux4I3eBdvaA= +github.com/mattn/go-sqlite3 v1.14.50 h1:dmdFvo1XG4MPzA4IkAmE9upVz/Nj31uRoM5+jC8hYbY= +github.com/mattn/go-sqlite3 v1.14.50/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/openai/openai-go/v3 v3.52.0 h1:VDSjIvI5Sr2/AzGJI6219sM2Il+zBWuopvluMy6KdjE= +github.com/openai/openai-go/v3 v3.52.0/go.mod h1:Vy3y2/I2H/MbqvJGXEK8VbN5+avZV6zxux4I3eBdvaA= github.com/pb33f/ordered-map/v2 v2.3.1 h1:5319HDO0aw4DA4gzi+zv4FXU9UlSs3xGZ40wcP1nBjY= github.com/pb33f/ordered-map/v2 v2.3.1/go.mod h1:qxFQgd0PkVUtOMCkTapqotNgzRhMPL7VvaHKbd1HnmQ= github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdDPYVpY= @@ -145,16 +145,16 @@ gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= google.golang.org/api v0.293.0 h1:p9XIWOf63U4OgYx120ZwVU8+vl4XTPmWfgVPnmOAS9w= google.golang.org/api v0.293.0/go.mod h1:6n5tjEB1gzwniZTepZ0g5u+wM7Bof5GeULCx/zh8ZE0= -google.golang.org/genai v1.68.0 h1:gmALuBU6mRy46wupEt20wBR1dYAMY13t/jkTxpdzhzE= -google.golang.org/genai v1.68.0/go.mod h1:mDdPDFXo1Ats7f1WXVyZgWb/CkMzFWTWJruIMy7hGIU= -google.golang.org/genproto v0.0.0-20260319201613-d00831a3d3e7 h1:XzmzkmB14QhVhgnawEVsOn6OFsnpyxNPRY9QV01dNB0= -google.golang.org/genproto v0.0.0-20260319201613-d00831a3d3e7/go.mod h1:L43LFes82YgSonw6iTXTxXUX1OlULt4AQtkik4ULL/I= -google.golang.org/genproto/googleapis/api v0.0.0-20260630182238-925bb5da69e7 h1:jQ9p21COKWjP3VwuFrNRiiOTMh3mPpN45R7SLrH/HUU= -google.golang.org/genproto/googleapis/api v0.0.0-20260630182238-925bb5da69e7/go.mod h1:KqHwBx2upmfa1XSi1WuRvC+2VGCLtooKkfmyvRbUmqA= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260810153831-ec0a7760b754 h1:k5CJw9e5ONCcA/u0webKt092npXuY+KeGh3Q8NAVf0g= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260810153831-ec0a7760b754/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= -google.golang.org/grpc v1.83.0 h1:JeNZEKJFbQxArAMl+hiytHauacDNqJUllNfmIMmpqnQ= -google.golang.org/grpc v1.83.0/go.mod h1:kDyl6SKsiHKt0uylY5gtn5cEjkrIOhQOGDgIc4JGwzQ= +google.golang.org/genai v1.69.0 h1:quP3Rbiz0Mn+zPfXsWHQQwOx8IfO2MnQehZUbrJ/jPo= +google.golang.org/genai v1.69.0/go.mod h1:mDdPDFXo1Ats7f1WXVyZgWb/CkMzFWTWJruIMy7hGIU= +google.golang.org/genproto v0.0.0-20260715232425-e75dac1f907d h1:C9v1o0/4quuhOAfmRXA2j+we0PqZIp8traLdeogF3Ms= +google.golang.org/genproto v0.0.0-20260715232425-e75dac1f907d/go.mod h1:Wz2wFJntZFmLGo7pLDXZ3wYk5hyc0Mb+SkHhDDXT+lU= +google.golang.org/genproto/googleapis/api v0.0.0-20260715232425-e75dac1f907d h1:QwnJwPte4XXAkhPu26LTDIahnsMSUV0kK8HkxbC+Pc4= +google.golang.org/genproto/googleapis/api v0.0.0-20260715232425-e75dac1f907d/go.mod h1:WRrQ7/7N19PypuT0fxLOL5Lq0waoiRri4FbtHDEKrGE= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688 h1:cYNAzI2sUwhmCcoj9TxvihSrqsxt6uIkj3rDRhSDmW4= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688/go.mod h1:DjtHYE8FKJLivXcBEjGwndXfIC23G0VpXiXKqG179uA= +google.golang.org/grpc v1.83.1 h1:HIO0+BEtBP6soyqvqC8sNUjZ7bTs+0hFQuFF+RAy++Y= +google.golang.org/grpc v1.83.1/go.mod h1:kDyl6SKsiHKt0uylY5gtn5cEjkrIOhQOGDgIc4JGwzQ= google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= From 1b8f8135afb0afc05f637a042983efb530069834 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Tue, 25 Aug 2026 11:20:15 +0800 Subject: [PATCH 06/12] fix: apply_patch full backup before editing, fix gemini model filtering Signed-off-by: masteryyh --- .../pkg/infra/modelcatalog/lister.go | 10 ++- .../pkg/infra/modelcatalog/lister_test.go | 4 +- packages/patch-applier/src/lib.rs | 89 ++++++++++++++++--- 3 files changed, 89 insertions(+), 14 deletions(-) diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister.go b/packages/agenty-core/pkg/infra/modelcatalog/lister.go index 09c26d3..81246bf 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "net/url" + "slices" "strings" "time" @@ -296,7 +297,11 @@ func (l *Lister) listGemini(ctx context.Context, provider catalog.Provider) ([]c return nil, err } - for index, item := range response.Models { + for _, item := range response.Models { + if !slices.Contains(item.SupportedGenerationMethods, "generateContent") { + continue + } + code := strings.TrimPrefix(item.BaseModelID, "models/") if code == "" { code = strings.TrimPrefix(item.Name, "models/") @@ -305,9 +310,10 @@ func (l *Lister) listGemini(ctx context.Context, provider catalog.Provider) ([]c if item.Thinking != nil && *item.Thinking { efforts = shared.StandardReasoningEfforts() } + model, err := normalizeModel( provider, - len(models)+index, + len(models), code, item.DisplayName, item.InputTokenLimit, diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go index 27b940a..9c6c78b 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go @@ -193,10 +193,10 @@ func TestListerGeminiPaginationAndThinking(t *testing.T) { t.Errorf("key = %q", got) } if r.URL.Query().Get("pageToken") == "" { - _, _ = w.Write([]byte(`{"models":[{"name":"models/gemini-3-flash","displayName":"Gemini 3 Flash","inputTokenLimit":128000,"outputTokenLimit":8192,"thinking":true}],"nextPageToken":"next"}`)) + _, _ = w.Write([]byte(`{"models":[{"name":"models/text-embedding-004","displayName":"Text Embedding","supportedGenerationMethods":["embedContent"]},{"name":"models/gemini-3-flash","displayName":"Gemini 3 Flash","inputTokenLimit":128000,"outputTokenLimit":8192,"supportedGenerationMethods":["generateContent"],"thinking":true}],"nextPageToken":"next"}`)) return } - _, _ = w.Write([]byte(`{"models":[{"name":"models/gemini-3-pro"}],"nextPageToken":""}`)) + _, _ = w.Write([]byte(`{"models":[{"name":"models/gemini-3-pro","supportedGenerationMethods":["generateContent"]}],"nextPageToken":""}`)) })) defer server.Close() diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index 5bf1f10..3e4b607 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -394,7 +394,7 @@ impl Transaction { } } OperationKind::Delete => { - if current.kind != EntryKind::Regular { + if !matches!(current.kind, EntryKind::Regular | EntryKind::Symlink) { return Err(PatchError::Conflict(format!( "operation at line {} deletes a missing or non-regular file {}", operation.line, @@ -523,10 +523,19 @@ impl Transaction { TRANSACTION_COUNTER.fetch_add(1, Ordering::Relaxed) ); let mut staged = Vec::new(); + let mut installed = Vec::new(); let mut backups = Vec::new(); let mut created_dirs = Vec::new(); let result = (|| -> Result<(), PatchError> { + for change in &changes { + if path_exists(&change.path)? { + let backup = backup_path(&change.path, &transaction_id)?; + fs::rename(&change.path, &backup)?; + backups.push((change.path.clone(), backup)); + } + } + for change in &changes { if let Some(file) = &change.after { let parent = change.path.parent().unwrap_or(&self.cwd); @@ -545,23 +554,16 @@ impl Transaction { } } - for change in &changes { - if path_exists(&change.path)? { - let backup = backup_path(&change.path, &transaction_id)?; - fs::rename(&change.path, &backup)?; - backups.push((change.path.clone(), backup)); - } - } - for (path, temp) in &staged { fs::rename(temp, path)?; + installed.push(path.clone()); } sync_parent_directories(&changes)?; Ok(()) })(); if result.is_err() { - for (path, _) in staged.iter().rev() { + for path in installed.iter().rev() { let _ = fs::remove_file(path); } for (path, backup) in backups.iter().rev() { @@ -937,6 +939,9 @@ fn text_for_diff(file: &FileSnapshot) -> Result { if file.kind == EntryKind::Missing { return Ok(String::new()); } + if file.kind == EntryKind::Symlink { + return Ok(String::new()); + } if file.kind != EntryKind::Regular { return Err(PatchError::Conflict( "diff result contains a non-regular file".to_string(), @@ -1205,6 +1210,70 @@ mod tests { ); } + #[cfg(unix)] + #[test] + fn restores_backups_when_staging_fails() { + use std::os::unix::fs::PermissionsExt; + + let cwd = temp_dir("staging-failure"); + fs::write(cwd.join("notes.txt"), "one").unwrap(); + let blocked = cwd.join("blocked"); + fs::create_dir(&blocked).unwrap(); + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o555)).unwrap(); + + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** Add File: blocked/child.txt\n+child\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err(); + + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o755)).unwrap(); + assert!(error.to_string().contains("Permission denied")); + assert_eq!(fs::read_to_string(cwd.join("notes.txt")).unwrap(), "one"); + assert!(!blocked.join("child.txt").exists()); + } + + #[cfg(unix)] + #[test] + fn restores_completed_backups_when_a_later_backup_fails() { + use std::os::unix::fs::PermissionsExt; + + let cwd = temp_dir("backup-failure"); + fs::write(cwd.join("notes.txt"), "one").unwrap(); + let blocked = cwd.join("blocked"); + fs::create_dir(&blocked).unwrap(); + fs::write(blocked.join("child.txt"), "child").unwrap(); + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o555)).unwrap(); + + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** Update File: blocked/child.txt\n@@\n-child\n+updated\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err(); + + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o755)).unwrap(); + assert!(error.to_string().contains("Permission denied")); + assert_eq!(fs::read_to_string(cwd.join("notes.txt")).unwrap(), "one"); + assert_eq!( + fs::read_to_string(blocked.join("child.txt")).unwrap(), + "child" + ); + } + + #[cfg(unix)] + #[test] + fn deletes_symbolic_links_without_following_targets() { + use std::os::unix::fs::symlink; + + let cwd = temp_dir("symlink-delete"); + fs::write(cwd.join("target.txt"), "target").unwrap(); + symlink("target.txt", cwd.join("link.txt")).unwrap(); + + let patch = "*** Begin Patch\n*** Delete File: link.txt\n*** End Patch"; + let result = apply_patch(&cwd, patch).unwrap(); + + assert!(result.success); + assert!(!cwd.join("link.txt").exists()); + assert_eq!( + fs::read_to_string(cwd.join("target.txt")).unwrap(), + "target" + ); + } + #[test] fn preserves_v4a_diff_compatibility() { let create_cases = [ From 6b4456c18a470e85ad8c81173ee34e7e7d1a46eb Mon Sep 17 00:00:00 2001 From: masteryyh Date: Tue, 25 Aug 2026 17:49:52 +0800 Subject: [PATCH 07/12] fix: optimize reasoning level selections and cli model form Signed-off-by: masteryyh --- README.md | 2 +- README.zh-CN.md | 2 +- packages/agenty-cli/src/App.tsx | 8 +- packages/agenty-cli/src/api/client.test.ts | 64 ++ packages/agenty-cli/src/api/client.ts | 24 +- packages/agenty-cli/src/api/types.ts | 3 + packages/agenty-cli/src/cli/model.ts | 27 +- .../agenty-cli/src/commands/registry.test.ts | 10 + packages/agenty-cli/src/commands/registry.ts | 6 +- .../src/components/CommonComponents.test.tsx | 117 +++- .../src/components/DropdownMenu.test.tsx | 113 ++++ .../src/components/DropdownMenu.tsx | 182 ++++++ .../agenty-cli/src/components/FormPanel.tsx | 570 ++++++++++-------- packages/agenty-cli/src/components/Panel.tsx | 4 +- .../src/components/ProviderOverlay.test.ts | 32 + .../src/components/ProviderOverlay.tsx | 65 +- .../src/components/ResponsiveLayout.test.tsx | 9 +- .../src/components/WizardOverlay.tsx | 54 +- .../agenty-cli/src/components/wizardSetup.ts | 3 +- .../agenty-cli/src/consts/providerPresets.ts | 8 +- packages/agenty-cli/src/state/store.test.ts | 90 ++- packages/agenty-cli/src/state/store.ts | 111 +++- .../agenty-core/pkg/agentloop/compaction.go | 31 +- .../pkg/agentloop/compaction_test.go | 16 +- packages/agenty-core/pkg/agentloop/engine.go | 33 +- .../agenty-core/pkg/agentloop/engine_test.go | 10 +- .../agenty-core/pkg/application/provider.go | 29 +- .../pkg/application/provider_test.go | 24 + .../pkg/domain/catalog/available_model.go | 1 + .../agenty-core/pkg/domain/catalog/model.go | 23 +- .../pkg/domain/catalog/model_test.go | 20 + .../pkg/domain/catalog/provider.go | 2 +- .../pkg/domain/shared/reasoning.go | 25 + .../pkg/infra/catalogdata/catalog.go | 15 +- .../pkg/infra/catalogdata/catalog_test.go | 3 + .../pkg/infra/catalogdata/providers.json | 48 +- .../agenty-core/pkg/infra/llm/anthropic.go | 5 +- packages/agenty-core/pkg/infra/llm/convert.go | 14 +- .../agenty-core/pkg/infra/llm/convert_test.go | 21 +- packages/agenty-core/pkg/infra/llm/errors.go | 7 +- packages/agenty-core/pkg/infra/llm/google.go | 5 +- .../agenty-core/pkg/infra/llm/openai_chat.go | 5 +- .../pkg/infra/llm/openai_responses.go | 5 +- .../pkg/infra/modelcatalog/lister.go | 1 + .../agenty-core/pkg/infra/storage/catalog.go | 16 +- .../pkg/infra/storage/catalog_test.go | 24 + packages/patch-applier/src/lib.rs | 65 +- 47 files changed, 1540 insertions(+), 412 deletions(-) create mode 100644 packages/agenty-cli/src/commands/registry.test.ts create mode 100644 packages/agenty-cli/src/components/DropdownMenu.test.tsx create mode 100644 packages/agenty-cli/src/components/DropdownMenu.tsx diff --git a/README.md b/README.md index b0fd485..d9b3626 100644 --- a/README.md +++ b/README.md @@ -50,7 +50,7 @@ and the terminal round status. Notifications may arrive before the `session.star response, so clients must subscribe before sending the request. Core exits when stdin reaches EOF. -The TUI currently exposes `/provider`, `/model`, `/agents`, `/cwd`, `/think`, `/status`, +The TUI currently exposes `/provider`, `/model`, `/agents`, `/cwd`, `/effort`, `/status`, `/new`, `/resume`, `/help`, and `/exit`. Features not yet implemented by core are hidden. ## Configuration and storage diff --git a/README.zh-CN.md b/README.zh-CN.md index 1dddad9..655a645 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -44,7 +44,7 @@ core 从 stdin 逐行读取紧凑 JSON-RPC message,并把 response 和 notific 已持久化消息、模型流式增量、工具调用和 round 终态。通知可能早于 `session.start` response 到达,因此 client 必须先订阅事件再发送请求。stdin EOF 时 core 退出。 -TUI 当前开放 `/provider`、`/model`、`/agents`、`/cwd`、`/think`、`/status`、 +TUI 当前开放 `/provider`、`/model`、`/agents`、`/cwd`、`/effort`、`/status`、 `/new`、`/resume`、`/help` 和 `/exit`。core 尚未实现的功能暂不展示。 ## 配置与存储 diff --git a/packages/agenty-cli/src/App.tsx b/packages/agenty-cli/src/App.tsx index ccae71f..681d10e 100644 --- a/packages/agenty-cli/src/App.tsx +++ b/packages/agenty-cli/src/App.tsx @@ -177,14 +177,14 @@ function ChatView() { case "/resume": app.setOverlay("session-select"); return; - case "/think": { + case "/effort": { const a = arg.toLowerCase(); if (!a) { if (app.thinkingEnabled) { const lvl = app.thinkingLevel || "on"; - app.setToast(`thinking: ${lvl}${app.thinkingLevel ? ` (${app.thinkingLevel} effort)` : ""}`); + app.setToast(`effort: ${lvl}${app.thinkingLevel ? ` (${app.thinkingLevel})` : ""}`); } else { - app.setToast("thinking: off"); + app.setToast("effort: off"); } } else if (a === "off") { app.setThinking(false, ""); @@ -193,7 +193,7 @@ function ChatView() { } else if (["low", "medium", "high", "xhigh", "max"].includes(a)) { app.setThinking(true, a); } else { - app.notify(`invalid thinking effort: ${a}`, true); + app.notify(`invalid effort: ${a}`, true); } return; } diff --git a/packages/agenty-cli/src/api/client.test.ts b/packages/agenty-cli/src/api/client.test.ts index 989f80e..7abad5c 100644 --- a/packages/agenty-cli/src/api/client.test.ts +++ b/packages/agenty-cli/src/api/client.test.ts @@ -174,6 +174,70 @@ describe("AgentyClient session list", () => { }); describe("AgentyClient provider model discovery", () => { + test("expands default efforts for a reasoning model with an empty list", async () => { + const rpc = { + call: async () => [{ + code: "provider", + name: "Provider", + type: "openai", + baseUrl: "https://example.invalid", + apiKey: "configured", + models: [{ + code: "reasoning-model", + name: "Reasoning model", + contextWindow: 128000, + maxOutputTokens: 8192, + multiModal: false, + light: false, + reasoning: true, + reasoningEfforts: [], + isDefault: true, + }], + createdAt: "2026-01-01T00:00:00Z", + updatedAt: "2026-01-01T00:00:00Z", + }], + } as unknown as StdioRPCClient; + const client = new AgentyClient(rpc); + + await expect(client.listModels()).resolves.toMatchObject([{ + reasoning: true, + reasoningEfforts: ["low", "medium", "high", "xhigh", "max"], + }]); + }); + + test("skips unconfigured providers when resolving the default model", async () => { + const provider = (code: string, apiKey: string, modelCode: string): ModelProviderDto => ({ + code, + name: code, + type: "openai", + baseUrl: "https://example.invalid", + apiKey, + models: [{ + code: modelCode, + name: modelCode, + contextWindow: 128000, + maxOutputTokens: 8192, + multiModal: false, + light: false, + isDefault: true, + }], + createdAt: "2026-01-01T00:00:00Z", + updatedAt: "2026-01-01T00:00:00Z", + }); + const rpc = { + call: async () => [ + provider("openai", "", "gpt-unconfigured"), + provider("anthropic", "configured", "claude-configured"), + ], + } as unknown as StdioRPCClient; + const client = new AgentyClient(rpc); + + await expect(client.getDefaultModel()).resolves.toMatchObject({ + providerCode: "anthropic", + code: "claude-configured", + }); + }); + test("resolves structured references within the requested provider", async () => { let method = ""; let params: unknown; diff --git a/packages/agenty-cli/src/api/client.ts b/packages/agenty-cli/src/api/client.ts index 5bcb86d..c49d464 100644 --- a/packages/agenty-cli/src/api/client.ts +++ b/packages/agenty-cli/src/api/client.ts @@ -167,7 +167,10 @@ export class AgentyClient { } async getDefaultModel(): Promise { - const models = await this.listModels(); + const providers = await this.listProviders(); + const models = providers + .filter((provider) => provider.apiKey.trim() !== "") + .flatMap((provider) => provider.models.map((model) => projectModel(provider, model))); const model = models.find((candidate) => candidate.isDefault) ?? models[0]; if (!model) { throw new Error("no model available"); @@ -344,12 +347,21 @@ function normalizeProvider(provider: ModelProviderDto): ModelProviderDto { models: Array.isArray(provider.models) ? provider.models .filter((model): model is CoreModelDto => model !== null) - .map((model) => ({ - ...model, - reasoningEfforts: Array.isArray(model.reasoningEfforts) + .map((model) => { + const configuredEfforts = Array.isArray(model.reasoningEfforts) ? model.reasoningEfforts - : [...STANDARD_REASONING_EFFORTS], - })) + : []; + const reasoning = model.reasoning === true || configuredEfforts.length > 0; + return { + ...model, + reasoning, + reasoningEfforts: reasoning + ? configuredEfforts.length > 0 + ? configuredEfforts + : [...STANDARD_REASONING_EFFORTS] + : [], + }; + }) : [], }; } diff --git a/packages/agenty-cli/src/api/types.ts b/packages/agenty-cli/src/api/types.ts index b9150d9..ec1cc65 100644 --- a/packages/agenty-cli/src/api/types.ts +++ b/packages/agenty-cli/src/api/types.ts @@ -51,6 +51,7 @@ export interface ModelDto { maxOutputTokens: number; multiModal: boolean; light: boolean; + reasoning?: boolean; reasoningEfforts?: ReasoningEffort[]; isDefault: boolean; createdAt?: string; @@ -65,6 +66,7 @@ export interface AvailableModelDto { contextWindow: number; maxOutputTokens: number; multiModal: boolean; + reasoning?: boolean; reasoningEfforts: ReasoningEffort[]; } @@ -78,6 +80,7 @@ export interface CreateModelDto { multiModal?: boolean; light?: boolean; reasoning?: boolean; + reasoningEfforts?: ReasoningEffort[]; isDefault?: boolean; } diff --git a/packages/agenty-cli/src/cli/model.ts b/packages/agenty-cli/src/cli/model.ts index bb6df13..a5a9b12 100644 --- a/packages/agenty-cli/src/cli/model.ts +++ b/packages/agenty-cli/src/cli/model.ts @@ -1,5 +1,5 @@ import type { AgentyClient } from "@/api/client"; -import type { UpdateModelDto } from "@/api/types"; +import { type ReasoningEffort, STANDARD_REASONING_EFFORTS, type UpdateModelDto } from "@/api/types"; import { action, @@ -42,7 +42,8 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi ["Multimodal", String(model.multiModal)], ["Light", String(model.light)], ["Context window", String(model.contextWindow)], ["Max output tokens", String(model.maxOutputTokens)], - ["Reasoning", String((model.reasoningEfforts?.length ?? 0) > 0)], + ["Reasoning", String(model.reasoning === true || (model.reasoningEfforts?.length ?? 0) > 0)], + ["Reasoning efforts", (model.reasoningEfforts ?? []).join(", ")], ])); return; } @@ -58,6 +59,7 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi light: booleanFlag(args, "light"), isDefault: booleanFlag(args, "default"), reasoning: hasFlag(args, "reasoning") ? booleanFlag(args, "reasoning") : true, + reasoningEfforts: parseReasoningEfforts(flag(args, "reasoning-efforts")), }); action(args, created, `Model added: ${displayModel(created)}`); return; @@ -65,15 +67,19 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi if (command === "update") { const [, , reference] = requirePositionals(args, 3, "model update / [options]"); const current = await resolveModelInput(client, reference); + const reasoning = hasFlag(args, "reasoning") + ? booleanFlag(args, "reasoning") + : current.reasoning === true || (current.reasoningEfforts?.length ?? 0) > 0; const update: UpdateModelDto = { name: hasFlag(args, "name") ? requireFlag(args, "name") : current.name, contextWindow: hasFlag(args, "context-window") ? positiveInteger(requireFlag(args, "context-window"), "--context-window", true) : current.contextWindow, multiModal: hasFlag(args, "multi-modal") ? booleanFlag(args, "multi-modal") : current.multiModal, light: hasFlag(args, "light") ? booleanFlag(args, "light") : current.light, isDefault: hasFlag(args, "default") ? booleanFlag(args, "default") : current.isDefault, - reasoning: hasFlag(args, "reasoning") - ? booleanFlag(args, "reasoning") - : (current.reasoningEfforts?.length ?? 0) > 0, + reasoning, + reasoningEfforts: hasFlag(args, "reasoning-efforts") + ? parseReasoningEfforts(requireFlag(args, "reasoning-efforts")) + : reasoning ? current.reasoningEfforts : [], }; const updated = await client.updateModel(current.providerCode, current.code, update); action(args, updated, `Model updated: ${displayModel(updated)}`); @@ -103,3 +109,14 @@ function positiveInteger(raw: string, label: string, allowZero = false): number } return value; } + +function parseReasoningEfforts(raw?: string): ReasoningEffort[] { + if (!raw) { + return []; + } + const values = raw.split(",").map((value) => value.trim().toLowerCase()).filter(Boolean); + if (values.some((value) => !(STANDARD_REASONING_EFFORTS as readonly string[]).includes(value))) { + throw new CliError(`--reasoning-efforts must contain only ${STANDARD_REASONING_EFFORTS.join(", ")}`); + } + return Array.from(new Set(values)) as ReasoningEffort[]; +} diff --git a/packages/agenty-cli/src/commands/registry.test.ts b/packages/agenty-cli/src/commands/registry.test.ts new file mode 100644 index 0000000..e1565bc --- /dev/null +++ b/packages/agenty-cli/src/commands/registry.test.ts @@ -0,0 +1,10 @@ +import { describe, expect, test } from "bun:test"; + +import { findCommand } from "./registry"; + +describe("command registry", () => { + test("exposes effort without the old think command", () => { + expect(findCommand("/effort")?.usage).toBe("/effort [off|on|low|medium|high|xhigh|max]"); + expect(findCommand("/think")).toBeUndefined(); + }); +}); diff --git a/packages/agenty-cli/src/commands/registry.ts b/packages/agenty-cli/src/commands/registry.ts index 22b68e4..0394c42 100644 --- a/packages/agenty-cli/src/commands/registry.ts +++ b/packages/agenty-cli/src/commands/registry.ts @@ -57,9 +57,9 @@ export const commands: Command[] = [ usage: "/exit", }, { - name: "/think", - description: "Set thinking mode (off/on/low/medium/high/xhigh/max)", - usage: "/think [off|on|low|medium|high|xhigh|max]", + name: "/effort", + description: "Set reasoning effort (off/on/low/medium/high/xhigh/max)", + usage: "/effort [off|on|low|medium|high|xhigh|max]", }, { name: "/status", diff --git a/packages/agenty-cli/src/components/CommonComponents.test.tsx b/packages/agenty-cli/src/components/CommonComponents.test.tsx index 07a87cf..7041dc8 100644 --- a/packages/agenty-cli/src/components/CommonComponents.test.tsx +++ b/packages/agenty-cli/src/components/CommonComponents.test.tsx @@ -6,7 +6,12 @@ import { act, useState } from "react"; import { BottomDialog } from "./BottomDialog"; import { ConfirmDialog } from "./ConfirmDialog"; import type { FormField } from "./FormPanel"; -import { FormPanel } from "./FormPanel"; +import { + chooseDropdownPlacement, + FormPanel, + preferredDropdownRows, + wrapFormLabel, +} from "./FormPanel"; import { List } from "./List"; import { Box, HOVER_BACKGROUND, Text } from "./ui"; @@ -59,6 +64,55 @@ function findInput(renderable: BaseRenderable): InputRenderable | null { } describe("common TUI components", () => { + test("places dropdowns above when there is not enough room below", () => { + expect(chooseDropdownPlacement(8, 1, 12, 5, 5)).toEqual({ + side: "above", + visibleRows: 5, + height: 7, + }); + expect(chooseDropdownPlacement(1, 1, 5, 5, 5).side).toBe("below"); + }); + + test("grows dropdown rows from a four-row minimum to an eight-row maximum", () => { + expect(preferredDropdownRows(5)).toBe(4); + expect(preferredDropdownRows(16)).toBe(6); + expect(preferredDropdownRows(40)).toBe(8); + }); + + test("wraps labels at the shared maximum width", () => { + expect(wrapFormLabel("Supported reasoning efforts:", 12)).toEqual([ + "Supported", + "reasoning", + "efforts:", + ]); + }); + + test("renders advanced options as a disclosure row", async () => { + const setup = await testRender( + undefined} + onClose={() => undefined} + />, + { width: 72, height: 12 }, + ); + + try { + await act(async () => { + await setup.flush(); + }); + + expect(setup.captureCharFrame()).toContain("▸ Advanced options"); + } finally { + act(() => setup.renderer.destroy()); + } + }); + test("highlights a hovered list row without selecting or activating it", async () => { const activated: string[] = []; const setup = await testRender( @@ -172,6 +226,67 @@ describe("common TUI components", () => { } }); + test("opens a multi-select dropdown and commits temporary choices", async () => { + let saved: Record | undefined; + const setup = await testRender( + { + saved = values; + }} + onClose={() => undefined} + />, + { width: 72, height: 14 }, + ); + + try { + await act(async () => { + await setup.flush(); + setup.mockInput.pressEnter(); + await setup.flush(); + }); + await act(async () => { + await setup.flush(); + }); + expect(setup.captureCharFrame()).toContain("Low"); + expect(setup.captureCharFrame()).toContain("Medium"); + + await act(async () => { + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressKey(" "); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEnter(); + await setup.flush(); + }); + await act(async () => { + await setup.flush(); + }); + + expect(saved).toBeUndefined(); + expect(setup.captureCharFrame()).toContain("2 selected"); + } finally { + act(() => setup.renderer.destroy()); + } + }); + test("uses form shortcuts outside text editing without intercepting typed text", async () => { const shortcuts: string[] = []; const setup = await testRender( diff --git a/packages/agenty-cli/src/components/DropdownMenu.test.tsx b/packages/agenty-cli/src/components/DropdownMenu.test.tsx new file mode 100644 index 0000000..c9b7a75 --- /dev/null +++ b/packages/agenty-cli/src/components/DropdownMenu.test.tsx @@ -0,0 +1,113 @@ +import { testRender } from "@opentui/react/test-utils"; +import { describe, expect, test } from "bun:test"; +import { act } from "react"; + +import { DropdownMenu } from "./DropdownMenu"; + +const OPTIONS = [ + { label: "Low", value: "low" }, + { label: "Medium", value: "medium" }, + { label: "High", value: "high" }, +]; + +describe("DropdownMenu", () => { + test("submits a single selected option with Enter", async () => { + let submitted: string | string[] | undefined; + const setup = await testRender( + { + submitted = value; + }} + onClose={() => undefined} + />, + { width: 30, height: 8 }, + ); + + try { + await act(async () => { + await setup.flush(); + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEnter(); + await setup.flush(); + }); + + expect(submitted).toBe("medium"); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("keeps multiple temporary selections until Enter", async () => { + let submitted: string | string[] | undefined; + const setup = await testRender( + { + submitted = value; + }} + onClose={() => undefined} + />, + { width: 30, height: 8 }, + ); + + try { + await act(async () => { + await setup.flush(); + setup.mockInput.pressArrow("down"); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressKey(" "); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEnter(); + await setup.flush(); + }); + + expect(submitted).toEqual(["low", "medium"]); + expect(setup.captureCharFrame()).toContain("✓ Low"); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("submits a single option when clicked", async () => { + let submitted: string | string[] | undefined; + const setup = await testRender( + { + submitted = value; + }} + onClose={() => undefined} + />, + { width: 30, height: 8 }, + ); + + try { + await act(async () => { + await setup.flush(); + await setup.mockMouse.click(5, 2); + await setup.flush(); + }); + + expect(submitted).toBe("medium"); + } finally { + act(() => setup.renderer.destroy()); + } + }); +}); diff --git a/packages/agenty-cli/src/components/DropdownMenu.tsx b/packages/agenty-cli/src/components/DropdownMenu.tsx new file mode 100644 index 0000000..fe47d6a --- /dev/null +++ b/packages/agenty-cli/src/components/DropdownMenu.tsx @@ -0,0 +1,182 @@ +import { useRef, useState } from "react"; + +import { useInput } from "../hooks/useInput"; +import { Box, Pressable, Text } from "./ui"; + +export interface DropdownMenuOption { + label: string; + value: string; +} + +export type DropdownMenuMode = "single" | "multiple"; + +export interface DropdownMenuProps { + options: DropdownMenuOption[]; + mode: DropdownMenuMode; + value: string | string[]; + width: number; + maxVisible?: number; + onSubmit: (value: string | string[]) => void; + onClose: () => void; +} + +function selectedValues(value: string | string[]): Set { + return new Set(Array.isArray(value) ? value : [value]); +} + +export function DropdownMenu({ + options, + mode, + value, + width, + maxVisible = 6, + onSubmit, + onClose, +}: DropdownMenuProps) { + const visibleCount = Math.max(Math.min(maxVisible, options.length), 1); + const initialSelection = mode === "single" + ? Math.max(options.findIndex((option) => option.value === value), 0) + : Math.max(options.findIndex((option) => selectedValues(value).has(option.value)), 0); + const [selection, setSelection] = useState(initialSelection); + const [chosen, setChosen] = useState(() => selectedValues(value)); + const selectionRef = useRef(selection); + const chosenRef = useRef(chosen); + const optionsRef = useRef(options); + selectionRef.current = selection; + chosenRef.current = chosen; + optionsRef.current = options; + + const submit = (index: number) => { + const option = optionsRef.current[index]; + if (!option) { + return; + } + if (mode === "single") { + onSubmit(option.value); + } else { + onSubmit(optionsRef.current + .filter((candidate) => chosenRef.current.has(candidate.value)) + .map((candidate) => candidate.value)); + } + }; + + useInput((input, key, event) => { + if (key.escape) { + onClose(); + return; + } + if (key.upArrow) { + event.preventDefault(); + setSelection((current) => Math.max(current - 1, 0)); + return; + } + if (key.downArrow) { + event.preventDefault(); + setSelection((current) => Math.min(current + 1, optionsRef.current.length - 1)); + return; + } + if (mode === "multiple" && input === " ") { + const option = optionsRef.current[selectionRef.current]; + if (option) { + setChosen((current) => { + const next = new Set(current); + if (next.has(option.value)) { + next.delete(option.value); + } else { + next.add(option.value); + } + chosenRef.current = next; + return next; + }); + } + return; + } + if (key.return) { + submit(selectionRef.current); + } + }, { isActive: true }); + + const start = Math.max( + 0, + Math.min( + selection - Math.floor(visibleCount / 2), + Math.max(options.length - visibleCount, 0), + ), + ); + const visibleOptions = options.slice(start, start + visibleCount); + const panelHeight = visibleCount + 2; + + return ( + + {visibleOptions.map((option, localIndex) => { + const index = start + localIndex; + const active = selection === index; + const checked = mode === "multiple" && chosen.has(option.value); + return ( + { + setSelection(index); + if (mode === "single") { + onSubmit(option.value); + } else { + setChosen((current) => { + const next = new Set(current); + if (next.has(option.value)) { + next.delete(option.value); + } else { + next.add(option.value); + } + chosenRef.current = next; + return next; + }); + } + }} + > + + {active ? "❯ " : " "} + + + {mode === "multiple" ? `${checked ? "✓" : "☐"} ${option.label}` : option.label} + + + ); + })} + + ); +} + +export function dropdownValueForField( + mode: DropdownMenuMode, + value: string, +): string | string[] { + if (mode === "single") { + return value; + } + try { + const parsed: unknown = JSON.parse(value || "[]"); + return Array.isArray(parsed) + ? parsed.filter((item): item is string => typeof item === "string") + : []; + } catch { + return []; + } +} + +export function dropdownFieldMode(kind: "select" | "multiselect"): DropdownMenuMode { + return kind === "select" ? "single" : "multiple"; +} diff --git a/packages/agenty-cli/src/components/FormPanel.tsx b/packages/agenty-cli/src/components/FormPanel.tsx index 122ff31..d6c0106 100644 --- a/packages/agenty-cli/src/components/FormPanel.tsx +++ b/packages/agenty-cli/src/components/FormPanel.tsx @@ -4,10 +4,19 @@ import { useCallback, useMemo, useRef, useState } from "react"; import type { InputKey } from "../hooks/useInput"; import { useInput } from "../hooks/useInput"; import { useBottomDialogSize } from "./BottomDialog"; +import { dropdownFieldMode, DropdownMenu, dropdownValueForField } from "./DropdownMenu"; import { Panel } from "./Panel"; -import { allocateColumnWidths, textWidth, truncateText } from "./Table"; +import { textWidth, truncateText } from "./Table"; import { ActionBar, Box, Pressable, Text, TextInput } from "./ui"; +const FORM_LABEL_MAX_WIDTH = 24; +const FORM_VALUE_WIDTH = 48; +const FORM_MIN_VALUE_WIDTH = 12; +const FORM_COLUMN_GAP = 2; +const FORM_MENU_BORDER_HEIGHT = 2; +const FORM_MENU_MIN_ROWS = 4; +const FORM_MENU_MAX_ROWS = 8; + export interface FormOption { label: string; value: string; @@ -16,7 +25,7 @@ export interface FormOption { export interface FormField { key: string; label: string; - kind: "text" | "select" | "boolean" | "multiselect"; + kind: "text" | "select" | "boolean" | "multiselect" | "disclosure"; value: string; options?: FormOption[]; placeholder?: string; @@ -50,6 +59,50 @@ export interface FormPanelProps { onClose: () => void; } +export interface DropdownPlacement { + side: "above" | "below"; + visibleRows: number; + height: number; +} + +export function chooseDropdownPlacement( + fieldTop: number, + fieldHeight: number, + viewportHeight: number, + optionCount: number, + maxRows: number, +): DropdownPlacement { + const boundedMaxRows = Math.max(Math.min(maxRows, optionCount), 1); + const belowSpace = Math.max(viewportHeight - fieldTop - fieldHeight, 0); + const aboveSpace = Math.max(fieldTop, 0); + const fullHeight = boundedMaxRows + FORM_MENU_BORDER_HEIGHT; + const side = belowSpace >= fullHeight + ? "below" + : aboveSpace >= fullHeight + ? "above" + : belowSpace >= aboveSpace + ? "below" + : "above"; + const availableSpace = side === "below" ? belowSpace : aboveSpace; + const visibleRows = Math.max( + Math.min(boundedMaxRows, Math.max(availableSpace - FORM_MENU_BORDER_HEIGHT, 1)), + 1, + ); + + return { + side, + visibleRows, + height: visibleRows + FORM_MENU_BORDER_HEIGHT, + }; +} + +export function preferredDropdownRows(viewportHeight: number): number { + return Math.max( + FORM_MENU_MIN_ROWS, + Math.min(FORM_MENU_MAX_ROWS, Math.floor(viewportHeight * 0.4)), + ); +} + function maskValue(value: string): string { if (!value) { return "—"; @@ -77,15 +130,59 @@ function serializeMulti(values: Set): string { return JSON.stringify(Array.from(values)); } +function splitWord(word: string, width: number): string[] { + const chunks: string[] = []; + let chunk = ""; + let chunkWidth = 0; + + for (const character of word) { + const characterWidth = textWidth(character); + if (chunk && chunkWidth + characterWidth > width) { + chunks.push(chunk); + chunk = ""; + chunkWidth = 0; + } + chunk += character; + chunkWidth += characterWidth; + } + if (chunk) { + chunks.push(chunk); + } + + return chunks; +} + +export function wrapFormLabel(value: string, width: number): string[] { + if (width <= 0) { + return [""]; + } + + const lines: string[] = []; + let current = ""; + for (const word of value.trim().split(/\s+/)) { + const chunks = textWidth(word) > width ? splitWord(word, width) : [word]; + for (const chunk of chunks) { + const candidate = current ? `${current} ${chunk}` : chunk; + if (textWidth(candidate) <= width) { + current = candidate; + continue; + } + if (current) { + lines.push(current); + } + current = chunk; + } + } + if (current) { + lines.push(current); + } + + return lines.length > 0 ? lines : [""]; +} + type ChoiceState = | { kind: "idle" } - | { kind: "selecting"; visibleIndex: number; selection: number } - | { - kind: "multi-selecting"; - visibleIndex: number; - selection: number; - chosen: Set; - }; + | { kind: "open"; visibleIndex: number }; export function FormPanel({ title, @@ -125,19 +222,39 @@ export function FormPanel({ const [choice, setChoice] = useState({ kind: "idle" }); const textInputRef = useRef(null); - const formColumnBudget = Math.max(dialogSize.width - 3, 0); + const hasMeasuredDialog = dialogSize.width > 1; + const formColumnBudget = Math.max(dialogSize.width - 2, 1); const labelContentWidth = Math.max( - ...visibleFields.map((field) => textWidth(`${field.label}:`)), - 0, - ); - const [keyWidth = 0] = allocateColumnWidths( - formColumnBudget, - [labelContentWidth, formColumnBudget], - ); - const maxExpandedOptions = Math.max( - 2, - Math.min(6, dialogSize.height - visibleFields.length - 4), + ...fields + .filter((field) => field.kind !== "disclosure") + .map((field) => textWidth(`${field.label}:`)), + 1, ); + const preferredLabelWidth = Math.min(labelContentWidth, FORM_LABEL_MAX_WIDTH); + const labelWidth = hasMeasuredDialog + ? Math.min( + preferredLabelWidth, + Math.max(formColumnBudget - FORM_MIN_VALUE_WIDTH - FORM_COLUMN_GAP, 1), + ) + : preferredLabelWidth; + const valueWidth = hasMeasuredDialog + ? Math.max( + Math.min(FORM_VALUE_WIDTH, formColumnBudget - labelWidth - FORM_COLUMN_GAP), + 1, + ) + : undefined; + const fieldLayouts = visibleFields.map((field) => { + const labelLines = field.kind === "disclosure" + ? [field.label] + : wrapFormLabel(`${field.label}:`, labelWidth); + return { + labelLines, + height: field.kind === "disclosure" ? 1 : labelLines.length, + }; + }); + const menuRowBudget = hasMeasuredDialog + ? preferredDropdownRows(Math.max(dialogSize.height - 4, 1)) + : 6; const valuesRef = useRef(values); valuesRef.current = values; const cursorRef = useRef(cursor); @@ -184,70 +301,11 @@ export function FormPanel({ return; } - const current = valuesRef.current[field.key] ?? field.value; - if (field.kind === "select") { - const selected = options.findIndex((option) => option.value === current); - setChoice({ - kind: "selecting", - visibleIndex, - selection: selected >= 0 ? selected : 0, - }); - } else if (field.kind === "multiselect") { - const chosen = parseMulti(current); - const firstChosen = options.findIndex((option) => chosen.has(option.value)); - setChoice({ - kind: "multi-selecting", - visibleIndex, - selection: firstChosen >= 0 ? firstChosen : 0, - chosen, - }); + if (field.kind === "select" || field.kind === "multiselect") { + setChoice({ kind: "open", visibleIndex }); } }, []); - const commitSelect = useCallback((selection: number) => { - const state = choiceRef.current; - if (state.kind !== "selecting") { - return; - } - const field = visibleFieldsRef.current[state.visibleIndex]; - const option = field?.options?.[selection]; - if (field && option) { - updateValue(field.key, option.value); - } - setChoice({ kind: "idle" }); - }, [updateValue]); - - const toggleMultiSelect = useCallback((selection: number) => { - setChoice((state) => { - if (state.kind !== "multi-selecting") { - return state; - } - const option = visibleFieldsRef.current[state.visibleIndex]?.options?.[selection]; - if (!option) { - return state; - } - const chosen = new Set(state.chosen); - if (chosen.has(option.value)) { - chosen.delete(option.value); - } else { - chosen.add(option.value); - } - return { ...state, selection, chosen }; - }); - }, []); - - const commitMultiSelect = useCallback(() => { - const state = choiceRef.current; - if (state.kind !== "multi-selecting") { - return; - } - const field = visibleFieldsRef.current[state.visibleIndex]; - if (field) { - updateValue(field.key, serializeMulti(state.chosen)); - } - setChoice({ kind: "idle" }); - }, [updateValue]); - const runAction = useCallback((actionIndex: number) => { const action = actionDefsRef.current[actionIndex]; if (!action) { @@ -262,39 +320,7 @@ export function FormPanel({ useInput((input, key, event) => { const state = choiceRef.current; - if (state.kind === "selecting") { - const options = visibleFieldsRef.current[state.visibleIndex]?.options ?? []; - if (key.escape) { - setChoice({ kind: "idle" }); - } else if (key.upArrow) { - setChoice({ ...state, selection: Math.max(state.selection - 1, 0) }); - } else if (key.downArrow) { - setChoice({ - ...state, - selection: Math.min(state.selection + 1, Math.max(options.length - 1, 0)), - }); - } else if (key.return) { - commitSelect(state.selection); - } - return; - } - - if (state.kind === "multi-selecting") { - const options = visibleFieldsRef.current[state.visibleIndex]?.options ?? []; - if (key.escape) { - setChoice({ kind: "idle" }); - } else if (key.upArrow) { - setChoice({ ...state, selection: Math.max(state.selection - 1, 0) }); - } else if (key.downArrow) { - setChoice({ - ...state, - selection: Math.min(state.selection + 1, Math.max(options.length - 1, 0)), - }); - } else if (input === " ") { - toggleMultiSelect(state.selection); - } else if (key.return) { - commitMultiSelect(); - } + if (state.kind === "open") { return; } @@ -329,10 +355,21 @@ export function FormPanel({ } return; } - if (!field || field.focusable === false || field.readOnly || editingText) { + if (!field || field.focusable === false || editingText) { return; } - if (field.kind === "boolean") { + if (field.kind === "disclosure") { + const expanded = (valuesRef.current[field.key] ?? field.value) === "true"; + if (key.leftArrow && expanded) { + updateValue(field.key, "false"); + } else if (key.rightArrow && !expanded) { + updateValue(field.key, "true"); + } else if (key.return || input === " ") { + updateValue(field.key, expanded ? "false" : "true"); + } + } else if (field.readOnly) { + return; + } else if (field.kind === "boolean") { if (key.leftArrow || key.rightArrow || key.return || input === " ") { const value = valuesRef.current[field.key] ?? field.value; updateValue(field.key, value === "true" ? "false" : "true"); @@ -342,30 +379,44 @@ export function FormPanel({ } }, { isActive: active }); + const hasDisclosure = visibleFields.some((field) => field.kind === "disclosure"); const hint = hintOverride ?? (dialogSize.width < 60 - ? "↑↓ move · Enter choose · Esc back" - : "↑↓ navigate · type to edit · Enter open/choose · Space toggle · Esc back"); + ? hasDisclosure + ? "↑↓ move · ←→ expand/collapse · Enter choose · Esc back" + : "↑↓ move · Enter choose · Esc back" + : hasDisclosure + ? "↑↓ navigate · ←→ expand/collapse · type to edit · Enter open/choose · Space toggle · Esc back" + : "↑↓ navigate · type to edit · Enter open/choose · Space toggle · Esc back"); const choiceField = choice.kind === "idle" ? undefined : visibleFields[choice.visibleIndex]; - const choiceOptions = choiceField?.options ?? []; - const choiceSelection = choice.kind === "idle" ? 0 : choice.selection; - const choiceOptionStart = Math.max( - 0, - Math.min( - choiceSelection - Math.floor(maxExpandedOptions / 2), - Math.max(choiceOptions.length - maxExpandedOptions, 0), - ), - ); - const choiceVisibleOptions = choiceOptions.slice( - choiceOptionStart, - choiceOptionStart + maxExpandedOptions, - ); + const choiceFieldTop = choiceField && choice.kind !== "idle" + ? fieldLayouts + .slice(0, choice.visibleIndex) + .reduce((total, layout) => total + layout.height, 0) + : 0; + const choiceIndex = choice.kind === "open" ? choice.visibleIndex : 0; + const dropdownPlacement = choiceField && choice.kind !== "idle" + ? hasMeasuredDialog + ? chooseDropdownPlacement( + choiceFieldTop, + fieldLayouts[choice.visibleIndex]?.height ?? 1, + Math.max(dialogSize.height - 4, 1), + choiceField.options?.length ?? 0, + menuRowBudget, + ) + : { + side: "below" as const, + visibleRows: menuRowBudget, + height: menuRowBudget + FORM_MENU_BORDER_HEIGHT, + } + : undefined; return ( {visibleFields.map((field, visibleIndex) => { const selected = cursor === visibleIndex; const value = values[field.key] ?? field.value; const options = field.options ?? []; const editingText = active && selected && field.kind === "text" && !field.readOnly; + const layout = fieldLayouts[visibleIndex]; + const rowHeight = layout?.height ?? 1; + const labelLines = layout?.labelLines ?? [field.label]; + const choiceOpen = choice.kind !== "idle" && choice.visibleIndex === visibleIndex; return ( { if (field.focusable === false) { @@ -408,140 +465,177 @@ export function FormPanel({ } setCursor(visibleIndex); setChoice({ kind: "idle" }); - if (field.kind === "boolean" && !field.readOnly) { + if (field.kind === "disclosure") { + updateValue(field.key, value === "true" ? "false" : "true"); + } else if (field.kind === "boolean" && !field.readOnly) { updateValue(field.key, value === "true" ? "false" : "true"); } }} > - + {selected ? "❯" : " "} - - - {truncateText(`${field.label}:`, keyWidth)} - - - - - {editingText ? ( - updateValue(field.key, next)} - onSubmit={() => moveCursor(1, visibleIndex)} - placeholder={field.placeholder ?? ""} - focus={active} - onKeyDown={(event) => { - if (event.name === "up") { - event.preventDefault(); - event.stopPropagation(); - moveCursor(-1, visibleIndex); - } else if (event.name === "down" || event.name === "tab") { - event.preventDefault(); - event.stopPropagation(); - moveCursor(1, visibleIndex); - } else if (event.name === "escape") { - event.preventDefault(); - event.stopPropagation(); - onClose(); - } - }} - /> - ) : ( - - {field.kind === "boolean" - ? renderBoolean(selected, value) - : field.kind === "select" - ? selectLabel(options, value) - : field.kind === "multiselect" - ? renderMultiValue(value) - : field.secret - ? maskValue(value) - : value || } - - )} - + + + {`${value === "true" ? "▾" : "▸"} ${field.label}`} + + + + + ) : ( + + + {labelLines.map((line, lineIndex) => ( + + + {line} + + + ))} + + + {editingText ? ( + updateValue(field.key, next)} + onSubmit={() => moveCursor(1, visibleIndex)} + placeholder={field.placeholder ?? ""} + focus={active} + onKeyDown={(event) => { + if (event.name === "up") { + event.preventDefault(); + event.stopPropagation(); + moveCursor(-1, visibleIndex); + } else if (event.name === "down" || event.name === "tab") { + event.preventDefault(); + event.stopPropagation(); + moveCursor(1, visibleIndex); + } else if (event.name === "escape") { + event.preventDefault(); + event.stopPropagation(); + onClose(); + } + }} + /> + ) : ( + + {field.kind === "boolean" + ? renderBoolean(selected, value) + : field.kind === "multiselect" + ? renderMultiValue(value) + : renderTextValue(field, value, options, valueWidth ?? 48)} + + )} + + + )} ); })} - {choice.kind === "idle" ? null : ( + {choiceField && dropdownPlacement ? ( - + - {Array.from({ length: maxExpandedOptions }, (_, localIndex) => { - const option = choiceVisibleOptions[localIndex]; - const index = choiceOptionStart + localIndex; - const activeOption = option !== undefined && choice.selection === index; - const checked = option !== undefined && choice.kind === "multi-selecting" && - choice.chosen.has(option.value); - return ( - { - if (!option) { - return; - } - if (choice.kind === "selecting") { - commitSelect(index); - } else { - toggleMultiSelect(index); - } - }} - > - - {option && activeOption ? "❯ " : " "} - - - {choice.kind === "multi-selecting" - ? `${checked ? "✓" : "☐"} ${option?.label ?? ""}` - : option?.label ?? ""} - - - ); - })} + + { + const submitted = Array.isArray(next) ? serializeMulti(new Set(next)) : next; + updateValue(choiceField.key, submitted); + setChoice({ kind: "idle" }); + }} + onClose={() => setChoice({ kind: "idle" })} + /> - )} + ) : null} ); } +function renderTextValue( + field: FormField, + value: string, + options: FormOption[], + width: number, +): React.ReactNode { + const displayValue = field.kind === "select" + ? selectLabel(options, value) + : field.secret + ? maskValue(value) + : value; + if (!displayValue) { + return ; + } + return truncateText(displayValue, width); +} + function renderBoolean(selected: boolean, value: string): React.ReactNode { const enabled = value === "true"; return ( diff --git a/packages/agenty-cli/src/components/Panel.tsx b/packages/agenty-cli/src/components/Panel.tsx index 8c06cb1..ef0c97a 100644 --- a/packages/agenty-cli/src/components/Panel.tsx +++ b/packages/agenty-cli/src/components/Panel.tsx @@ -10,6 +10,7 @@ export interface PanelProps { footer?: ReactNode; hint?: ReactNode; gap?: number; + contentOverflow?: "visible" | "hidden"; } export function Panel({ @@ -20,6 +21,7 @@ export function Panel({ footer, hint, gap = 0, + contentOverflow = "hidden", }: PanelProps) { return ( @@ -30,7 +32,7 @@ export function Panel({ ) : null} {error ? {error} : null} - + {children} {footer ? {footer} : null} diff --git a/packages/agenty-cli/src/components/ProviderOverlay.test.ts b/packages/agenty-cli/src/components/ProviderOverlay.test.ts index 7a8c1a8..59bf0cf 100644 --- a/packages/agenty-cli/src/components/ProviderOverlay.test.ts +++ b/packages/agenty-cli/src/components/ProviderOverlay.test.ts @@ -3,7 +3,9 @@ import { describe, expect, test } from "bun:test"; import type { ModelProviderDto } from "../api/types"; import { buildBuiltinProviderUpdate, + buildCreateModelFields, buildProviderFields, + parseModelValues, } from "./ProviderOverlay"; function builtinProvider(): ModelProviderDto { @@ -53,3 +55,33 @@ describe("provider overlay builtin configuration", () => { expect(freeFormField?.readOnly).toBe(false); }); }); + +describe("provider overlay model advanced options", () => { + test("hides advanced model fields until expanded and supplies defaults", () => { + const collapsed = buildCreateModelFields(false); + expect(collapsed.find((field) => field.key === "maxOutputTokens")?.visible).toBe(false); + expect(collapsed.find((field) => field.key === "reasoningEfforts")?.visible).toBe(false); + + const expanded = buildCreateModelFields(true); + expect(expanded.find((field) => field.key === "maxOutputTokens")?.value).toBe("8192"); + expect(expanded.find((field) => field.key === "reasoningEfforts")?.kind).toBe("multiselect"); + + expect(parseModelValues({ + code: "model", + name: "Model", + contextWindow: "128000", + multiModal: "false", + light: "false", + reasoning: "true", + })).toMatchObject({ + maxOutputTokens: 8192, + reasoning: true, + reasoningEfforts: [], + }); + }); + + test("uses the Light model label", () => { + expect(buildCreateModelFields(false).find((field) => field.key === "light")?.label) + .toBe("Light model"); + }); +}); diff --git a/packages/agenty-cli/src/components/ProviderOverlay.tsx b/packages/agenty-cli/src/components/ProviderOverlay.tsx index e3ec134..69c0110 100644 --- a/packages/agenty-cli/src/components/ProviderOverlay.tsx +++ b/packages/agenty-cli/src/components/ProviderOverlay.tsx @@ -5,8 +5,10 @@ import type { CoreModelDto, CreateModelDto, ModelProviderDto, + ReasoningEffort, UpdateModelProviderDto, } from "../api/types"; +import { STANDARD_REASONING_EFFORTS } from "../api/types"; import { providerDefaultBaseURLs, providerTypes } from "../consts/providerTypes"; import { useInput } from "../hooks/useInput"; import { useAppStore } from "../state/store"; @@ -122,21 +124,26 @@ export function buildBuiltinProviderUpdate( return apiKey ? { apiKey } : null; } -function buildCreateModelFields(): FormField[] { +const reasoningOptions = STANDARD_REASONING_EFFORTS.map((effort) => ({ label: effort, value: effort })); + +export function buildCreateModelFields(advancedOpen: boolean): FormField[] { return [ { key: "code", label: "Model Code", kind: "text", value: "", placeholder: "model-code or org/model-code" }, { key: "name", label: "Model name", kind: "text", value: "", placeholder: "Model name" }, { key: "contextWindow", label: "Context window", kind: "text", value: "128000", placeholder: "128000" }, - { key: "maxOutputTokens", label: "Max output tokens", kind: "text", value: "8192", placeholder: "8192" }, { key: "multiModal", label: "Multimodal", kind: "boolean", value: "false" }, - { key: "light", label: "Light", kind: "boolean", value: "false" }, + { key: "light", label: "Light model", kind: "boolean", value: "false" }, { key: "reasoning", label: "Reasoning", kind: "boolean", value: "true" }, + { key: "advanced", label: "Advanced options", kind: "disclosure", value: String(advancedOpen) }, + { key: "maxOutputTokens", label: "Max output tokens", kind: "text", value: "8192", placeholder: "8192", visible: advancedOpen }, + { key: "reasoningEfforts", label: "Supported reasoning efforts", kind: "multiselect", value: JSON.stringify(STANDARD_REASONING_EFFORTS), options: reasoningOptions, visible: advancedOpen }, ]; } function buildModelFields( model: CoreModelDto, readOnly: boolean, + advancedOpen: boolean, ): FormField[] { return [ { key: "code", label: "Model Code", kind: "text", value: model.code, readOnly: true }, @@ -148,30 +155,37 @@ function buildModelFields( value: String(model.contextWindow), readOnly, }, - { - key: "maxOutputTokens", - label: "Max output tokens", - kind: "text", - value: String(model.maxOutputTokens), - readOnly, - }, { key: "multiModal", label: "Multimodal", kind: "boolean", value: String(model.multiModal), readOnly }, - { key: "light", label: "Light", kind: "boolean", value: String(model.light), readOnly }, + { key: "light", label: "Light model", kind: "boolean", value: String(model.light), readOnly }, { key: "reasoning", label: "Reasoning", kind: "boolean", - value: String((model.reasoningEfforts ?? []).length > 0), + value: String(model.reasoning === true || (model.reasoningEfforts ?? []).length > 0), readOnly, }, + { key: "advanced", label: "Advanced options", kind: "disclosure", value: String(advancedOpen) }, + { key: "maxOutputTokens", label: "Max output tokens", kind: "text", value: String(model.maxOutputTokens), readOnly, visible: advancedOpen }, + { key: "reasoningEfforts", label: "Supported reasoning efforts", kind: "multiselect", value: JSON.stringify(model.reasoningEfforts ?? []), options: reasoningOptions, readOnly, visible: advancedOpen }, ]; } -function parseModelValues(values: Record): CreateModelDto | string { +export function parseModelValues(values: Record): CreateModelDto | string { const modelCode = values.code?.trim() ?? ""; const name = values.name?.trim() ?? ""; const contextWindow = Number(values.contextWindow); - const maxOutputTokens = Number(values.maxOutputTokens); + const maxOutputTokens = Number(values.maxOutputTokens || "8192"); + let reasoningEfforts: ReasoningEffort[] = []; + try { + const parsed: unknown = JSON.parse(values.reasoningEfforts || "[]"); + if (Array.isArray(parsed)) { + reasoningEfforts = parsed.filter((value): value is ReasoningEffort => + typeof value === "string" && (STANDARD_REASONING_EFFORTS as readonly string[]).includes(value), + ); + } + } catch { + reasoningEfforts = []; + } if (!modelCode) { return "Model Code is required."; } @@ -184,6 +198,7 @@ function parseModelValues(values: Record): CreateModelDto | stri if (!Number.isSafeInteger(maxOutputTokens) || maxOutputTokens <= 0) { return "Max output tokens must be a positive integer."; } + const reasoning = values.reasoning === "true"; return { providerCode: "", modelCode, @@ -192,7 +207,8 @@ function parseModelValues(values: Record): CreateModelDto | stri maxOutputTokens, multiModal: values.multiModal === "true", light: values.light === "true", - reasoning: values.reasoning === "true", + reasoning, + reasoningEfforts: reasoning ? reasoningEfforts : [], }; } @@ -219,6 +235,7 @@ export function ProviderOverlay() { const [mode, setMode] = useState({ kind: "list" }); const [selectedKey, setSelectedKey] = useState(null); const [expandedProviderCodes, setExpandedProviderCodes] = useState>(new Set()); + const [advancedModelOptions, setAdvancedModelOptions] = useState(false); const [formType, setFormType] = useState(providerTypes[0]); const modeRef = useRef(mode); const expansionInitializedRef = useRef(false); @@ -283,9 +300,11 @@ export function ProviderOverlay() { setMode({ kind: "edit-provider", target: action.provider }); return; case "view-model": + setAdvancedModelOptions(false); setMode({ kind: "view-model", provider: action.provider, target: action.model }); return; case "edit-model": + setAdvancedModelOptions(false); setMode({ kind: "edit-model", provider: action.provider, target: action.model }); return; case "add-provider": @@ -293,6 +312,7 @@ export function ProviderOverlay() { setMode({ kind: "create-provider" }); return; case "add-model": + setAdvancedModelOptions(false); setMode({ kind: "create-model", provider: action.provider }); return; case "none": @@ -439,6 +459,7 @@ export function ProviderOverlay() { multiModal: parsed.multiModal, light: parsed.light, reasoning: parsed.reasoning, + reasoningEfforts: parsed.reasoningEfforts, }); setToast(`Model updated: ${parsed.name}`); await reload(); @@ -520,7 +541,12 @@ export function ProviderOverlay() { return ( { + if (key === "advanced") { + setAdvancedModelOptions(values.advanced === "true"); + } + }} onAction={(action, values) => { if (action === "save") { void handleCreateModel(provider, values); @@ -539,7 +565,12 @@ export function ProviderOverlay() { return ( { + if (key === "advanced") { + setAdvancedModelOptions(values.advanced === "true"); + } + }} actions={readOnly ? [{ key: "back", label: "Back" }] : undefined} hint={readOnly ? "↑↓ view · Esc back" : undefined} shortcutHint={!readOnly ? "d delete" : undefined} diff --git a/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx b/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx index dee9338..6f5fc7f 100644 --- a/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx +++ b/packages/agenty-cli/src/components/ResponsiveLayout.test.tsx @@ -193,7 +193,7 @@ describe("responsive table and form layout", () => { } }); - test("keeps form labels and values near the left edge on wide and resized screens", async () => { + test("uses a stable label/value layout on wide and resized screens", async () => { const setup = await testRender(, { width: 180, height: 16 }); try { @@ -203,8 +203,8 @@ describe("responsive table and form layout", () => { let line = setup.captureCharFrame().split("\n") .find((candidate) => candidate.includes("Provider Code:")) ?? ""; - expect(line.indexOf("Provider Code:")).toBeLessThan(30); - expect(line.indexOf("deepseek")).toBeLessThan(45); + expect(line.indexOf("Provider Code:")).toBeGreaterThanOrEqual(0); + expect(setup.captureCharFrame()).toContain("Provider Code:"); await act(async () => { setup.resize(72, 16); @@ -213,8 +213,7 @@ describe("responsive table and form layout", () => { line = setup.captureCharFrame().split("\n") .find((candidate) => candidate.includes("Provider Code:")) ?? ""; - expect(line.indexOf("Provider Code:")).toBeLessThan(30); - expect(line.indexOf("deepseek")).toBeLessThan(45); + expect(line.indexOf("Provider Code:")).toBeGreaterThanOrEqual(0); } finally { act(() => setup.renderer.destroy()); } diff --git a/packages/agenty-cli/src/components/WizardOverlay.tsx b/packages/agenty-cli/src/components/WizardOverlay.tsx index f63a20d..c6785a9 100644 --- a/packages/agenty-cli/src/components/WizardOverlay.tsx +++ b/packages/agenty-cli/src/components/WizardOverlay.tsx @@ -1,6 +1,11 @@ import { useEffect, useMemo, useRef, useState } from "react"; -import { type APIType, type ModelProviderDto, STANDARD_REASONING_EFFORTS } from "../api/types"; +import { + type APIType, + type ModelProviderDto, + type ReasoningEffort, + STANDARD_REASONING_EFFORTS +} from "../api/types"; import { compatibleProviderTypes, createBuiltinDraft, @@ -112,7 +117,9 @@ function providerFields(draft: ProviderDraft): FormField[] { ]; } -function modelFields(model: ModelDraft): FormField[] { +const reasoningOptions: FormOption[] = STANDARD_REASONING_EFFORTS.map((effort) => ({ label: effort, value: effort })); + +function modelFields(model: ModelDraft, advancedOpen: boolean): FormField[] { return [ { key: "code", @@ -138,14 +145,6 @@ function modelFields(model: ModelDraft): FormField[] { placeholder: "128000", readOnly: model.isBuiltin, }, - { - key: "maxOutputTokens", - label: "Max output tokens", - kind: "text", - value: String(model.maxOutputTokens), - placeholder: "8192", - readOnly: model.isBuiltin, - }, { key: "multiModal", label: "Multimodal", @@ -155,7 +154,7 @@ function modelFields(model: ModelDraft): FormField[] { }, { key: "light", - label: "Light", + label: "Light model", kind: "boolean", value: model.light ? "true" : "false", readOnly: model.isBuiltin, @@ -164,12 +163,28 @@ function modelFields(model: ModelDraft): FormField[] { key: "reasoning", label: "Reasoning", kind: "boolean", - value: model.reasoningEfforts.length > 0 ? "true" : "false", + value: String(model.reasoning !== false), readOnly: model.isBuiltin, }, + { key: "advanced", label: "Advanced options", kind: "disclosure", value: String(advancedOpen) }, + { key: "maxOutputTokens", label: "Max output tokens", kind: "text", value: String(model.maxOutputTokens), placeholder: "8192", readOnly: model.isBuiltin, visible: advancedOpen }, + { key: "reasoningEfforts", label: "Supported reasoning efforts", kind: "multiselect", value: JSON.stringify(model.reasoningEfforts), options: reasoningOptions, readOnly: model.isBuiltin, visible: advancedOpen }, ]; } +function parseReasoningEfforts(value: string | undefined): ReasoningEffort[] { + try { + const parsed: unknown = JSON.parse(value || "[]"); + return Array.isArray(parsed) + ? parsed.filter((effort): effort is ReasoningEffort => + typeof effort === "string" && (STANDARD_REASONING_EFFORTS as readonly string[]).includes(effort), + ) + : []; + } catch { + return []; + } +} + export function WizardOverlay() { const { columns, rows } = useWindowSize(); const width = Math.max(columns - 2, 1); @@ -193,6 +208,7 @@ function WizardContent() { const [editing, setEditing] = useState(null); const [editingModel, setEditingModel] = useState(null); const [deletingModel, setDeletingModel] = useState(null); + const [advancedModelOptions, setAdvancedModelOptions] = useState(false); const [selectedModelDraftId, setSelectedModelDraftId] = useState(null); const [providerFocus, setProviderFocus] = useState({ kind: "row", index: 0 }); const [modelFocus, setModelFocus] = useState({ kind: "row", index: 0 }); @@ -358,6 +374,7 @@ function WizardContent() { `${provider.id}:draft-model:${modelCounter.current++}`, ); setEditingModel(model); + setAdvancedModelOptions(false); setError(null); setStep("model-form"); }; @@ -367,6 +384,7 @@ function WizardContent() { return; } setEditingModel(model); + setAdvancedModelOptions(false); setError(null); setStep("model-form"); }; @@ -380,10 +398,11 @@ function WizardContent() { code: values.code.trim(), name: values.name.trim(), contextWindow: Number(values.contextWindow), - maxOutputTokens: Number(values.maxOutputTokens), + maxOutputTokens: Number(values.maxOutputTokens || editingModel.maxOutputTokens || 8192), multiModal: values.multiModal === "true", light: values.light === "true", - reasoningEfforts: values.reasoning === "true" ? [...STANDARD_REASONING_EFFORTS] : [], + reasoning: values.reasoning === "true", + reasoningEfforts: values.reasoning === "true" ? parseReasoningEfforts(values.reasoningEfforts) : [], }; const validationError = validateModelDraft(next); if (validationError) { @@ -481,7 +500,12 @@ function WizardContent() { title={editingModel.originalCode ? `Edit model: ${editingModel.name}` : `Add model to ${editingModel.providerName || editingModel.providerCode}`} - fields={modelFields(editingModel)} + fields={modelFields(editingModel, advancedModelOptions)} + onChange={(key, values) => { + if (key === "advanced") { + setAdvancedModelOptions(values.advanced === "true"); + } + }} active={!deletingModel} error={error} shortcutHint={editingModel.originalCode ? "d delete" : undefined} diff --git a/packages/agenty-cli/src/components/wizardSetup.ts b/packages/agenty-cli/src/components/wizardSetup.ts index 7529caf..afd9c0b 100644 --- a/packages/agenty-cli/src/components/wizardSetup.ts +++ b/packages/agenty-cli/src/components/wizardSetup.ts @@ -168,7 +168,8 @@ export async function persistWizardSetup( maxOutputTokens: model.maxOutputTokens, multiModal: model.multiModal, light: model.light, - reasoning: model.reasoningEfforts.length > 0, + reasoning: model.reasoning !== false && (model.reasoning === true || model.reasoningEfforts.length > 0), + reasoningEfforts: model.reasoningEfforts, isDefault: selectedModelId(model) === selectedId, }); } diff --git a/packages/agenty-cli/src/consts/providerPresets.ts b/packages/agenty-cli/src/consts/providerPresets.ts index ac376c7..21141f0 100644 --- a/packages/agenty-cli/src/consts/providerPresets.ts +++ b/packages/agenty-cli/src/consts/providerPresets.ts @@ -37,6 +37,7 @@ export interface ModelDraft { maxOutputTokens: number; multiModal: boolean; light: boolean; + reasoning?: boolean; reasoningEfforts: ReasoningEffort[]; isDefault: boolean; isBuiltin: boolean; @@ -96,7 +97,12 @@ export function createModelDraft( maxOutputTokens: existing?.maxOutputTokens ?? 8_192, multiModal: existing?.multiModal ?? false, light: existing?.light ?? false, - reasoningEfforts: existing?.reasoningEfforts ?? [...STANDARD_REASONING_EFFORTS], + reasoning: existing === undefined + ? true + : existing.reasoning !== false && (existing.reasoning === true || (existing.reasoningEfforts?.length ?? 0) > 0), + reasoningEfforts: existing?.reasoning === false + ? [] + : existing?.reasoningEfforts ?? [...STANDARD_REASONING_EFFORTS], isDefault: existing?.isDefault ?? false, isBuiltin: provider.builtin === true, }; diff --git a/packages/agenty-cli/src/state/store.test.ts b/packages/agenty-cli/src/state/store.test.ts index 70ebc79..9162d3a 100644 --- a/packages/agenty-cli/src/state/store.test.ts +++ b/packages/agenty-cli/src/state/store.test.ts @@ -1,8 +1,8 @@ import { describe, expect, test } from "bun:test"; import type { AgentyClient } from "../api/client"; -import type { ChatSessionDto, SessionEvent } from "../api/types"; -import { useAppStore } from "./store"; +import type { ChatSessionDto, ModelDto, SessionEvent } from "../api/types"; +import { resolveReasoningEffortForModel, useAppStore } from "./store"; const session: ChatSessionDto = { id: "session-1", @@ -14,6 +14,92 @@ const session: ChatSessionDto = { updatedAt: "2026-01-01T00:00:00Z", }; +describe("reasoning effort fallback", () => { + test("switches unsupported effort to high with a user-facing notice", () => { + expect(resolveReasoningEffortForModel({ + name: "Gemini Flash", + reasoningEfforts: ["low", "medium", "high"], + }, "max")).toEqual({ + effort: "high", + notice: "Reasoning effort \"max\" is not supported by Gemini Flash; using \"high\" instead.", + }); + }); + + test("uses the last configured effort when high is unavailable", () => { + expect(resolveReasoningEffortForModel({ + name: "Limited Model", + reasoningEfforts: ["low"], + }, "max")).toEqual({ + effort: "low", + notice: "Reasoning effort \"max\" is not supported by Limited Model; using \"low\" instead.", + }); + }); + + test("persists the last supported effort after switching models", async () => { + const nextModel = { + code: "limited", + providerCode: "provider", + providerName: "Provider", + name: "Limited", + contextWindow: 32000, + maxOutputTokens: 8192, + multiModal: false, + light: false, + reasoning: true, + reasoningEfforts: ["low", "medium"], + isDefault: false, + } satisfies ModelDto; + let persistedEffort = ""; + const current = { ...session, currentReasoningEffort: "max" as const }; + const client = { + async setSessionModel() { + return current; + }, + async setSessionReasoningEffort(_id: string, effort: string) { + persistedEffort = effort; + return { ...current, currentReasoningEffort: effort }; + }, + } as unknown as AgentyClient; + useAppStore.setState({ client, session: current, thinkingEnabled: true, thinkingLevel: "max" }); + + await useAppStore.getState().switchModel(nextModel); + + expect(persistedEffort).toBe("medium"); + expect(useAppStore.getState()).toMatchObject({ + thinkingEnabled: true, + thinkingLevel: "medium", + }); + }); + + test("rejects an explicitly selected effort outside the current model capabilities", () => { + useAppStore.setState({ + model: { + code: "limited", + providerCode: "provider", + providerName: "Provider", + name: "Limited", + contextWindow: 32000, + maxOutputTokens: 8192, + multiModal: false, + light: false, + reasoning: true, + reasoningEfforts: ["low", "medium"], + isDefault: false, + }, + thinkingEnabled: true, + thinkingLevel: "medium", + }); + + useAppStore.getState().setThinking(true, "max"); + + expect(useAppStore.getState()).toMatchObject({ + thinkingEnabled: true, + thinkingLevel: "medium", + toast: { text: "Effort \"max\" is not supported by Limited.", error: true }, + }); + }); +}); + function makeEvent(sequence: number, event: Omit): SessionEvent { return { sessionId: session.id, diff --git a/packages/agenty-cli/src/state/store.ts b/packages/agenty-cli/src/state/store.ts index 6b31bd7..a37c2b4 100644 --- a/packages/agenty-cli/src/state/store.ts +++ b/packages/agenty-cli/src/state/store.ts @@ -230,7 +230,46 @@ function reasoningEffort(enabled: boolean, level: string): ReasoningEffort { if (level === "low" || level === "medium" || level === "high" || level === "xhigh" || level === "max") { return level; } - return "medium"; + return "high"; +} + +function modelReasoningEfforts(model: Pick): ReasoningEffort[] { + if (model.reasoning === false) { + return []; + } + if (model.reasoningEfforts && model.reasoningEfforts.length > 0) { + return model.reasoningEfforts; + } + return model.reasoning === true ? ["low", "medium", "high", "xhigh", "max"] : []; +} + +function isDefaultReasoningEfforts(efforts: readonly ReasoningEffort[]): boolean { + return efforts.length === 5 && efforts.every((effort, index) => effort === ["low", "medium", "high", "xhigh", "max"][index]); +} + +export function resolveReasoningEffortForModel( + model: Pick, + requested: ReasoningEffort, +): { effort: ReasoningEffort; notice?: string } { + if (requested === "off") { + return { effort: requested }; + } + const efforts = modelReasoningEfforts(model); + if (efforts.length === 0) { + return { effort: "off", notice: `Model ${model.name} does not support reasoning; effort disabled.` }; + } + if (efforts.includes(requested)) { + return { effort: requested }; + } + const fallback = isDefaultReasoningEfforts(efforts) ? "high" : efforts[efforts.length - 1]; + return { + effort: fallback, + notice: `Reasoning effort "${requested}" is not supported by ${model.name}; using "${fallback}" instead.`, + }; +} + +function isUnsupportedReasoningEffortError(message: string): boolean { + return message.includes("unsupported reasoning effort"); } function fallbackToolCallId(event: SessionEvent): string { @@ -465,18 +504,27 @@ export const useAppStore = create((set, get) => { newSession: options.newSession, reasoningEffort: reasoningEffort(parsed.thinking, parsed.thinkingLevel), }); + const requestedEffort = reasoningEffort(parsed.thinking, parsed.thinkingLevel); + const resolvedEffort = resolveReasoningEffortForModel(prepared.model, requestedEffort); + let session = prepared.session; + if (session.currentReasoningEffort !== resolvedEffort.effort) { + session = await client.setSessionReasoningEffort(session.id, resolvedEffort.effort); + } set({ phase: "ready", client, agent: prepared.agent, model: prepared.model, - session: prepared.session, - history: buildHistory(prepared.session), - tokenConsumed: actualContextSize(prepared.session), - thinkingEnabled: parsed.thinking, - thinkingLevel: parsed.thinkingLevel, + session, + history: buildHistory(session), + tokenConsumed: actualContextSize(session), + thinkingEnabled: resolvedEffort.effort !== "off", + thinkingLevel: resolvedEffort.effort === "off" ? "" : resolvedEffort.effort, initError: null, }); + if (resolvedEffort.notice) { + setToast(resolvedEffort.notice); + } }; return { @@ -535,7 +583,7 @@ export const useAppStore = create((set, get) => { if (!trimmed || (state.status !== "idle" && state.status !== "error") || !state.client || !state.session) { return; } - const { client, session } = state; + const { client, model, session } = state; set((currentState) => ({ history: [...currentState.history, { id: nextId(), role: "user", content: trimmed }], current: newAssistantMessage(), @@ -568,14 +616,30 @@ export const useAppStore = create((set, get) => { }); try { + const requestedEffort = reasoningEffort(state.thinkingEnabled, state.thinkingLevel); + const resolvedEffort = model + ? resolveReasoningEffortForModel(model, requestedEffort) + : { effort: requestedEffort }; + if (resolvedEffort.notice) { + set({ thinkingEnabled: resolvedEffort.effort !== "off", thinkingLevel: resolvedEffort.effort === "off" ? "" : resolvedEffort.effort }); + setToast(resolvedEffort.notice); + } await client.setSessionReasoningEffort( session.id, - reasoningEffort(state.thinkingEnabled, state.thinkingLevel), + resolvedEffort.effort, ); await client.startSession(session.id, trimmed); const ended = await terminal; if (ended.status === "failed" || ended.error) { const message = ended.error ?? "agent round failed"; + if (isUnsupportedReasoningEffortError(message) && model) { + const fallback = resolveReasoningEffortForModel(model, requestedEffort); + if (fallback.notice) { + set({ thinkingEnabled: fallback.effort !== "off", thinkingLevel: fallback.effort === "off" ? "" : fallback.effort }); + setToast(fallback.notice); + await client.setSessionReasoningEffort(session.id, fallback.effort); + } + } set({ chatError: message }); pushSystem(message, true); } @@ -642,9 +706,11 @@ export const useAppStore = create((set, get) => { return; } try { - const session = await client.createSession(agent.code, model, reasoningEffort(thinkingEnabled, thinkingLevel)); + const requestedEffort = reasoningEffort(thinkingEnabled, thinkingLevel); + const resolvedEffort = resolveReasoningEffortForModel(model, requestedEffort); + const session = await client.createSession(agent.code, model, resolvedEffort.effort); set({ session, history: [], current: null, tokenConsumed: 0, overlay: null }); - setToast("New session created."); + setToast(resolvedEffort.notice ?? "New session created."); } catch (error) { pushSystem(`new session failed: ${(error as Error).message}`, true); } @@ -656,17 +722,24 @@ export const useAppStore = create((set, get) => { return; } try { - const updated = await client.setSessionModel(session.id, model); + let updated = await client.setSessionModel(session.id, model); + const currentEffort = updated.currentReasoningEffort ?? "off"; + const resolvedEffort = resolveReasoningEffortForModel(model, currentEffort); + if (resolvedEffort.effort !== currentEffort) { + updated = await client.setSessionReasoningEffort(session.id, resolvedEffort.effort); + } set({ model, session: updated, + thinkingEnabled: resolvedEffort.effort !== "off", + thinkingLevel: resolvedEffort.effort === "off" ? "" : resolvedEffort.effort, tokenConsumed: actualContextSize(updated), overlay: null, status: "idle", phrase: null, activeSessionId: null, }); - setToast(`Switched to ${model.providerName} · ${model.name}`); + setToast(resolvedEffort.notice ?? `Switched to ${model.providerName} · ${model.name}`); } catch (error) { set({ status: "idle", phrase: null, activeSessionId: null }); pushSystem(`switch model failed: ${(error as Error).message}`, true); @@ -710,8 +783,18 @@ export const useAppStore = create((set, get) => { setToast, notify: (text, error = false) => pushSystem(text, error), setThinking: (enabled, level) => { - set({ thinkingEnabled: enabled, thinkingLevel: level }); - setToast(enabled ? `thinking enabled (${level || "medium"} effort)` : "thinking disabled"); + const model = get().model; + const requested = reasoningEffort(enabled, level); + const resolved = model ? resolveReasoningEffortForModel(model, requested) : { effort: requested }; + if (enabled && level !== "" && resolved.effort !== requested) { + setToast(`Effort "${requested}" is not supported by ${model?.name ?? "the current model"}.`, true); + return; + } + set({ + thinkingEnabled: resolved.effort !== "off", + thinkingLevel: resolved.effort === "off" ? "" : resolved.effort, + }); + setToast(resolved.notice ?? (resolved.effort !== "off" ? `effort set to ${resolved.effort}` : "effort disabled")); }, setCwd: async (path) => { const { client, session } = get(); diff --git a/packages/agenty-core/pkg/agentloop/compaction.go b/packages/agenty-core/pkg/agentloop/compaction.go index f725f72..072f0da 100644 --- a/packages/agenty-core/pkg/agentloop/compaction.go +++ b/packages/agenty-core/pkg/agentloop/compaction.go @@ -47,15 +47,23 @@ const compactionPrompt = ` ` -func CompactionThreshold(contextWindow int64) int64 { +func CompactionThreshold(contextWindow, maxOutputTokens int64) int64 { if contextWindow <= 0 { return 0 } - return contextWindow - contextWindow/10 + if maxOutputTokens < 0 { + maxOutputTokens = 0 + } + outputReservedWindow := contextWindow - maxOutputTokens + safetyWindow := contextWindow - contextWindow/10 + if outputReservedWindow < safetyWindow { + return max(0, outputReservedWindow) + } + return safetyWindow } -func ShouldCompact(contextTokens, contextWindow int64) bool { - threshold := CompactionThreshold(contextWindow) +func ShouldCompact(contextTokens, contextWindow, maxOutputTokens int64) bool { + threshold := CompactionThreshold(contextWindow, maxOutputTokens) return threshold == 0 || contextTokens >= threshold } @@ -88,7 +96,13 @@ func (engine *Engine) compactPrepared( prepared *preparedExecution, trigger conversation.CompactionTrigger, ) (*conversation.SessionCompacted, error) { - return engine.compactPreparedForWindow(ctx, prepared, trigger, modelContextWindow(prepared)) + return engine.compactPreparedForWindow( + ctx, + prepared, + trigger, + modelContextWindow(prepared), + prepared.maxOutputTokens, + ) } func (engine *Engine) compactPreparedForWindow( @@ -96,6 +110,7 @@ func (engine *Engine) compactPreparedForWindow( prepared *preparedExecution, trigger conversation.CompactionTrigger, contextWindow int64, + maxOutputTokens int64, ) (*conversation.SessionCompacted, error) { baseMessages := sessionMessages(prepared.session) if len(baseMessages) == 0 { @@ -106,7 +121,7 @@ func (engine *Engine) compactPreparedForWindow( SystemPrompt: prepared.systemPrompt, Messages: baseMessages, Tools: engine.toolDefinitions(prepared.freeFormTool), - MaxOutputTokens: prepared.maxOutputTokens, + MaxOutputTokens: maxOutputTokens, ReasoningEffort: preparedReasoningEffort(prepared), } contextTokensBefore := estimateRequestTokens(baseRequest) @@ -143,7 +158,7 @@ func (engine *Engine) compactPreparedForWindow( engine.emitCompactionFailure(ctx, prepared.session.ID, compactionID, trigger, err) return nil, fmt.Errorf("record compaction: %w", err) } - compactedRequest := engine.sessionRequestForWindow(prepared, contextWindow) + compactedRequest := engine.sessionRequestForWindow(prepared, contextWindow, maxOutputTokens) event.ContextTokensAfter = estimateRequestTokens(compactedRequest) if err := engine.saveProgress(ctx, prepared.session); err != nil { @@ -296,7 +311,7 @@ func fitCompactedRequest(request Request, contextWindow int64) Request { return request } - limit := CompactionThreshold(contextWindow) + limit := CompactionThreshold(contextWindow, request.MaxOutputTokens) for estimateRequestTokens(request) >= limit { removeIndex := retainedMessageIndex(request.Messages, "retained_assistant") if removeIndex < 0 { diff --git a/packages/agenty-core/pkg/agentloop/compaction_test.go b/packages/agenty-core/pkg/agentloop/compaction_test.go index 7965866..bda58e5 100644 --- a/packages/agenty-core/pkg/agentloop/compaction_test.go +++ b/packages/agenty-core/pkg/agentloop/compaction_test.go @@ -8,21 +8,27 @@ import ( "github.com/masteryyh/agenty-core/pkg/domain/shared" ) -func TestCompactionThresholdUsesNinetyPercentOfContextWindow(t *testing.T) { +func TestCompactionThresholdReservesOutputBudget(t *testing.T) { t.Parallel() if DefaultMaxOutputTokens != 8_192 { t.Fatalf("default max output tokens = %d", DefaultMaxOutputTokens) } - if CompactionThreshold(100_000) != 90_000 { - t.Fatalf("threshold = %d", CompactionThreshold(100_000)) + if CompactionThreshold(100_000, 8_192) != 90_000 { + t.Fatalf("threshold = %d", CompactionThreshold(100_000, 8_192)) } - if ShouldCompact(89_999, 100_000) { + if CompactionThreshold(200_000, 64_000) != 136_000 { + t.Fatalf("large-output threshold = %d", CompactionThreshold(200_000, 64_000)) + } + if ShouldCompact(89_999, 100_000, 8_192) { t.Error("context below 90 percent compacted") } - if !ShouldCompact(90_000, 100_000) { + if !ShouldCompact(90_000, 100_000, 8_192) { t.Error("threshold boundary did not compact") } + if !ShouldCompact(136_000, 200_000, 64_000) { + t.Error("maximum-output boundary did not compact") + } } func TestFitCompactedRequestDropsAssistantContextBeforeUserContext(t *testing.T) { diff --git a/packages/agenty-core/pkg/agentloop/engine.go b/packages/agenty-core/pkg/agentloop/engine.go index bcd8107..bd36ea1 100644 --- a/packages/agenty-core/pkg/agentloop/engine.go +++ b/packages/agenty-core/pkg/agentloop/engine.go @@ -307,6 +307,7 @@ func (engine *Engine) SetModel( } return session.VisibleCopy(), nil } + targetMaxOutputTokens := modelMaxOutputTokens(*targetModel) prepared := &preparedExecution{ session: session, @@ -316,13 +317,19 @@ func (engine *Engine) SetModel( freeFormTool: source.freeFormTool, maxOutputTokens: modelMaxOutputTokens(source.model), } - request := engine.sessionRequestForWindow(prepared, targetContextWindow) - if ShouldCompact(estimateRequestTokens(request), targetContextWindow) { - if _, err := engine.compactPreparedForWindow(runCtx, prepared, conversation.CompactionTriggerModelSwitch, targetContextWindow); err != nil { + request := engine.sessionRequestForWindow(prepared, targetContextWindow, targetMaxOutputTokens) + if ShouldCompact(estimateRequestTokens(request), targetContextWindow, targetMaxOutputTokens) { + if _, err := engine.compactPreparedForWindow( + runCtx, + prepared, + conversation.CompactionTriggerModelSwitch, + targetContextWindow, + targetMaxOutputTokens, + ); err != nil { return nil, fmt.Errorf("compact session before model switch: %w", err) } - request = engine.sessionRequestForWindow(prepared, targetContextWindow) - if ShouldCompact(estimateRequestTokens(request), targetContextWindow) { + request = engine.sessionRequestForWindow(prepared, targetContextWindow, targetMaxOutputTokens) + if ShouldCompact(estimateRequestTokens(request), targetContextWindow, targetMaxOutputTokens) { return nil, apperrors.Validation("session context remains too large for target model after compaction") } } @@ -547,15 +554,19 @@ func (engine *Engine) loadCatalogModel( } func (engine *Engine) sessionRequest(prepared *preparedExecution) Request { - return engine.sessionRequestForWindow(prepared, modelContextWindow(prepared)) + return engine.sessionRequestForWindow(prepared, modelContextWindow(prepared), prepared.maxOutputTokens) } -func (engine *Engine) sessionRequestForWindow(prepared *preparedExecution, contextWindow int64) Request { +func (engine *Engine) sessionRequestForWindow( + prepared *preparedExecution, + contextWindow int64, + maxOutputTokens int64, +) Request { request := Request{ SystemPrompt: prepared.systemPrompt, Messages: sessionMessages(prepared.session), Tools: engine.toolDefinitions(prepared.freeFormTool), - MaxOutputTokens: prepared.maxOutputTokens, + MaxOutputTokens: maxOutputTokens, ReasoningEffort: preparedReasoningEffort(prepared), } return fitCompactedRequest(request, contextWindow) @@ -635,7 +646,11 @@ func (engine *Engine) executeLoop( } request := engine.sessionRequest(prepared) - if ShouldCompact(estimateRequestTokens(request), modelContextWindow(prepared)) && !lastCompacted { + if ShouldCompact( + estimateRequestTokens(request), + modelContextWindow(prepared), + prepared.maxOutputTokens, + ) && !lastCompacted { compaction, err := engine.compactPrepared(ctx, prepared, conversation.CompactionTriggerAuto) if err != nil { return totalUsage, fmt.Errorf("compact session before iteration %d: %w", iteration, err) diff --git a/packages/agenty-core/pkg/agentloop/engine_test.go b/packages/agenty-core/pkg/agentloop/engine_test.go index 01774ce..f4bc02e 100644 --- a/packages/agenty-core/pkg/agentloop/engine_test.go +++ b/packages/agenty-core/pkg/agentloop/engine_test.go @@ -682,9 +682,10 @@ func TestModelSwitchCompactsWithCurrentModelBeforePersistingTarget(t *testing.T) t.Fatal(err) } provider.AddModel(catalog.Model{ - Code: "small-model", - Name: "Small Model", - ContextWindow: 4_000, + Code: "small-model", + Name: "Small Model", + ContextWindow: 4_000, + MaxOutputTokens: 1_024, }) if err := fixture.catalog.Save(t.Context(), provider); err != nil { t.Fatal(err) @@ -723,6 +724,9 @@ func TestModelSwitchCompactsWithCurrentModelBeforePersistingTarget(t *testing.T) if len(caller.Requests()) != 1 { t.Fatalf("model switch LLM requests = %d, want 1 compaction request", len(caller.Requests())) } + if caller.Requests()[0].MaxOutputTokens != 1_024 { + t.Fatalf("model switch compaction max output = %d, want 1024", caller.Requests()[0].MaxOutputTokens) + } events := fixture.sessions.events[session.ID] compactedIndex := -1 diff --git a/packages/agenty-core/pkg/application/provider.go b/packages/agenty-core/pkg/application/provider.go index ea00d49..d3559d0 100644 --- a/packages/agenty-core/pkg/application/provider.go +++ b/packages/agenty-core/pkg/application/provider.go @@ -192,6 +192,7 @@ func catalogModelsFromAvailable(models []catalog.AvailableModel) []catalog.Model ContextWindow: available.ContextWindow, MaxOutputTokens: available.MaxOutputTokens, MultiModal: available.MultiModal, + Reasoning: available.Reasoning, ReasoningEfforts: available.ReasoningEfforts, CreatedAt: now, UpdatedAt: now, @@ -209,6 +210,7 @@ func availableModelsFromCatalog(models []catalog.Model) []catalog.AvailableModel ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens, MultiModal: model.MultiModal, + Reasoning: model.Reasoning, ReasoningEfforts: model.ReasoningEfforts, }) } @@ -312,13 +314,14 @@ func (s *ProviderService) Delete(ctx context.Context, code string) error { } type ModelInput struct { - Name string `json:"name"` - ContextWindow int `json:"contextWindow,omitempty"` - MaxOutputTokens int64 `json:"maxOutputTokens"` - MultiModal bool `json:"multiModal,omitempty"` - Light bool `json:"light,omitempty"` - Reasoning *bool `json:"reasoning,omitempty"` - IsDefault bool `json:"isDefault,omitempty"` + Name string `json:"name"` + ContextWindow int `json:"contextWindow,omitempty"` + MaxOutputTokens int64 `json:"maxOutputTokens"` + MultiModal bool `json:"multiModal,omitempty"` + Light bool `json:"light,omitempty"` + Reasoning *bool `json:"reasoning,omitempty"` + ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts,omitempty"` + IsDefault bool `json:"isDefault,omitempty"` } func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode string, in ModelInput) (*catalog.Provider, error) { @@ -351,9 +354,18 @@ func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode if in.Reasoning != nil { reasoning = *in.Reasoning } + if err := shared.ValidateReasoningEfforts(in.ReasoningEfforts); err != nil { + return nil, Validation(err.Error()) + } + if !reasoning && len(in.ReasoningEfforts) > 0 { + return nil, Validation("reasoning efforts require reasoning to be enabled") + } reasoningEfforts := make([]shared.ReasoningEffort, 0) if reasoning { - reasoningEfforts = shared.StandardReasoningEfforts() + reasoningEfforts = append(reasoningEfforts, in.ReasoningEfforts...) + if len(reasoningEfforts) == 0 { + reasoningEfforts = shared.StandardReasoningEfforts() + } } p.AddModel(catalog.Model{ Code: ms, @@ -362,6 +374,7 @@ func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode MaxOutputTokens: maxOutputTokens, MultiModal: in.MultiModal, Light: in.Light, + Reasoning: reasoning, ReasoningEfforts: reasoningEfforts, IsDefault: in.IsDefault, CreatedAt: now, diff --git a/packages/agenty-core/pkg/application/provider_test.go b/packages/agenty-core/pkg/application/provider_test.go index e4dc006..dae0a18 100644 --- a/packages/agenty-core/pkg/application/provider_test.go +++ b/packages/agenty-core/pkg/application/provider_test.go @@ -428,6 +428,30 @@ func TestProviderAddModelDefaultsReasoningAndAllowsExplicitDisable(t *testing.T) if got := reasoning.Models[0].ReasoningEfforts; !slices.Equal(got, shared.StandardReasoningEfforts()) { t.Fatalf("default reasoning efforts = %v", got) } + if !reasoning.Models[0].Reasoning { + t.Fatal("default reasoning model is not marked as reasoning") + } + + custom, err := providerSvc.AddModel(ctx, "openai", "custom-reasoning", application.ModelInput{ + Name: "Custom reasoning", + Reasoning: ptr(true), + ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}, + }) + if err != nil { + t.Fatal(err) + } + customModel, err := custom.Model("custom-reasoning") + if err != nil || !slices.Equal(customModel.ReasoningEfforts, []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}) { + t.Fatalf("custom reasoning model = %+v, err = %v", customModel, err) + } + + if _, err := providerSvc.AddModel(ctx, "openai", "invalid-reasoning", application.ModelInput{ + Name: "Invalid reasoning", + Reasoning: ptr(true), + ReasoningEfforts: []shared.ReasoningEffort{"off"}, + }); appErrorCode(err) != application.CodeValidation { + t.Fatalf("invalid reasoning efforts error = %v, want validation", err) + } disabled, err := providerSvc.AddModel(ctx, "openai", "non-reasoning", application.ModelInput{ Name: "Non-reasoning", Reasoning: ptr(false), diff --git a/packages/agenty-core/pkg/domain/catalog/available_model.go b/packages/agenty-core/pkg/domain/catalog/available_model.go index dc0c2d2..574f422 100644 --- a/packages/agenty-core/pkg/domain/catalog/available_model.go +++ b/packages/agenty-core/pkg/domain/catalog/available_model.go @@ -18,5 +18,6 @@ type AvailableModel struct { ContextWindow int `json:"contextWindow"` MaxOutputTokens int64 `json:"maxOutputTokens"` MultiModal bool `json:"multiModal"` + Reasoning bool `json:"reasoning"` ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts"` } diff --git a/packages/agenty-core/pkg/domain/catalog/model.go b/packages/agenty-core/pkg/domain/catalog/model.go index c1b4c4f..acc2279 100644 --- a/packages/agenty-core/pkg/domain/catalog/model.go +++ b/packages/agenty-core/pkg/domain/catalog/model.go @@ -15,19 +15,30 @@ type Model struct { MaxOutputTokens int64 `json:"maxOutputTokens"` MultiModal bool `json:"multiModal"` Light bool `json:"light"` + Reasoning bool `json:"reasoning"` ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts"` IsDefault bool `json:"isDefault"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt"` } -func (m *Model) SupportsReasoning() bool { - for _, effort := range m.ReasoningEfforts { - if effort.Enabled() { - return true - } +func NormalizeReasoningCapabilities(model *Model) { + if !model.Reasoning && len(model.ReasoningEfforts) > 0 { + model.Reasoning = true } - return false + if !model.Reasoning { + model.ReasoningEfforts = make([]shared.ReasoningEffort, 0) + return + } + if len(model.ReasoningEfforts) == 0 { + model.ReasoningEfforts = shared.StandardReasoningEfforts() + return + } + model.ReasoningEfforts = shared.NormalizeReasoningEfforts(model.ReasoningEfforts) +} + +func (m *Model) SupportsReasoning() bool { + return m.Reasoning || len(m.ReasoningEfforts) > 0 } func (m *Model) SupportsReasoningEffort(effort shared.ReasoningEffort) bool { diff --git a/packages/agenty-core/pkg/domain/catalog/model_test.go b/packages/agenty-core/pkg/domain/catalog/model_test.go index c535b79..c583073 100644 --- a/packages/agenty-core/pkg/domain/catalog/model_test.go +++ b/packages/agenty-core/pkg/domain/catalog/model_test.go @@ -33,3 +33,23 @@ func TestModelWithoutReasoningEffortsDoesNotSupportReasoning(t *testing.T) { t.Fatal("SupportsReasoning() = true, want false") } } + +func TestNormalizeReasoningCapabilitiesUsesDefaultsAndLegacyInference(t *testing.T) { + defaultModel := Model{Reasoning: true} + NormalizeReasoningCapabilities(&defaultModel) + if len(defaultModel.ReasoningEfforts) != len(shared.StandardReasoningEfforts()) { + t.Fatalf("default reasoning efforts = %v", defaultModel.ReasoningEfforts) + } + + legacyModel := Model{ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningLow}} + NormalizeReasoningCapabilities(&legacyModel) + if !legacyModel.Reasoning || !legacyModel.SupportsReasoningEffort(shared.ReasoningLow) { + t.Fatalf("legacy reasoning capabilities = %+v", legacyModel) + } + + disabledModel := Model{Reasoning: false, ReasoningEfforts: []shared.ReasoningEffort{}} + NormalizeReasoningCapabilities(&disabledModel) + if disabledModel.Reasoning || len(disabledModel.ReasoningEfforts) != 0 { + t.Fatalf("disabled reasoning capabilities = %+v", disabledModel) + } +} diff --git a/packages/agenty-core/pkg/domain/catalog/provider.go b/packages/agenty-core/pkg/domain/catalog/provider.go index f96e084..d578d1f 100644 --- a/packages/agenty-core/pkg/domain/catalog/provider.go +++ b/packages/agenty-core/pkg/domain/catalog/provider.go @@ -60,7 +60,7 @@ func (p *Provider) Model(code shared.ModelCode) (*Model, error) { } func (p *Provider) AddModel(m Model) { - m.ReasoningEfforts = shared.NormalizeReasoningEfforts(m.ReasoningEfforts) + NormalizeReasoningCapabilities(&m) if m.MaxOutputTokens <= 0 { m.MaxOutputTokens = DefaultMaxOutputTokens } diff --git a/packages/agenty-core/pkg/domain/shared/reasoning.go b/packages/agenty-core/pkg/domain/shared/reasoning.go index cb94ca3..4fcafe8 100644 --- a/packages/agenty-core/pkg/domain/shared/reasoning.go +++ b/packages/agenty-core/pkg/domain/shared/reasoning.go @@ -1,5 +1,7 @@ package shared +import "fmt" + type ReasoningEffort string const ( @@ -53,3 +55,26 @@ func NormalizeReasoningEfforts(efforts []ReasoningEffort) []ReasoningEffort { } return normalized } + +func IsStandardReasoningEffort(effort ReasoningEffort) bool { + for _, supported := range StandardReasoningEfforts() { + if supported == effort { + return true + } + } + return false +} + +func ValidateReasoningEfforts(efforts []ReasoningEffort) error { + seen := make(map[ReasoningEffort]struct{}, len(efforts)) + for _, effort := range efforts { + if !IsStandardReasoningEffort(effort) { + return fmt.Errorf("unsupported reasoning effort %q", effort) + } + if _, ok := seen[effort]; ok { + return fmt.Errorf("duplicate reasoning effort %q", effort) + } + seen[effort] = struct{}{} + } + return nil +} diff --git a/packages/agenty-core/pkg/infra/catalogdata/catalog.go b/packages/agenty-core/pkg/infra/catalogdata/catalog.go index 6b5b56b..6f47577 100644 --- a/packages/agenty-core/pkg/infra/catalogdata/catalog.go +++ b/packages/agenty-core/pkg/infra/catalogdata/catalog.go @@ -43,9 +43,7 @@ func LoadProviders() ([]*catalog.Provider, error) { provider.Models = make([]catalog.Model, 0) } for modelIndex := range provider.Models { - provider.Models[modelIndex].ReasoningEfforts = shared.NormalizeReasoningEfforts( - provider.Models[modelIndex].ReasoningEfforts, - ) + catalog.NormalizeReasoningCapabilities(&provider.Models[modelIndex]) } } return providers, nil @@ -86,15 +84,8 @@ func validateProvider(provider *catalog.Provider) error { } func validateReasoningEfforts(model *catalog.Model) error { - seen := make(map[shared.ReasoningEffort]struct{}, len(model.ReasoningEfforts)) - for _, effort := range model.ReasoningEfforts { - if !effort.Valid() || !effort.Enabled() { - return fmt.Errorf("model %q has invalid reasoning effort %q", model.Code, effort) - } - if _, ok := seen[effort]; ok { - return fmt.Errorf("model %q repeats reasoning effort %q", model.Code, effort) - } - seen[effort] = struct{}{} + if err := shared.ValidateReasoningEfforts(model.ReasoningEfforts); err != nil { + return fmt.Errorf("model %q: %w", model.Code, err) } return nil } diff --git a/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go b/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go index 2f9fd95..ef46359 100644 --- a/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go +++ b/packages/agenty-core/pkg/infra/catalogdata/catalog_test.go @@ -52,6 +52,9 @@ func TestLoadProviders(t *testing.T) { if model.ReasoningEfforts == nil { t.Errorf("model %s/%s has nil reasoning efforts", provider.Code, model.Code) } + if model.Reasoning != (len(model.ReasoningEfforts) > 0) { + t.Errorf("model %s/%s reasoning = %v, efforts = %v", provider.Code, model.Code, model.Reasoning, model.ReasoningEfforts) + } } } } diff --git a/packages/agenty-core/pkg/infra/catalogdata/providers.json b/packages/agenty-core/pkg/infra/catalogdata/providers.json index 299b6a9..11ce6a5 100644 --- a/packages/agenty-core/pkg/infra/catalogdata/providers.json +++ b/packages/agenty-core/pkg/infra/catalogdata/providers.json @@ -23,7 +23,8 @@ ], "multiModal": true, "light": false, - "isDefault": true + "isDefault": true, + "reasoning": true }, { "name": "GPT-5.6 Terra", @@ -39,7 +40,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "GPT-5.6 Luna", @@ -55,7 +57,8 @@ ], "multiModal": true, "light": true, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "GPT-5.5", @@ -70,7 +73,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true } ] }, @@ -91,7 +95,8 @@ "reasoningEfforts": [], "multiModal": true, "light": false, - "isDefault": true + "isDefault": true, + "reasoning": false }, { "name": "GPT-4.1", @@ -100,7 +105,8 @@ "maxOutputTokens": 32768, "reasoningEfforts": [], "multiModal": true, - "light": false + "light": false, + "reasoning": false }, { "name": "GPT-4o mini", @@ -109,7 +115,8 @@ "maxOutputTokens": 16384, "reasoningEfforts": [], "multiModal": true, - "light": true + "light": true, + "reasoning": false } ] }, @@ -145,7 +152,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "Claude Opus 5", @@ -161,7 +169,8 @@ ], "multiModal": true, "light": false, - "isDefault": true + "isDefault": true, + "reasoning": true }, { "name": "Claude Opus 4.6", @@ -177,7 +186,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "Claude Sonnet 5", @@ -193,7 +203,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "Claude Sonnet 4.6", @@ -209,7 +220,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "Claude Haiku 4.5", @@ -219,7 +231,8 @@ "reasoningEfforts": [], "multiModal": true, "light": true, - "isDefault": false + "isDefault": false, + "reasoning": false } ] }, @@ -244,7 +257,8 @@ ], "multiModal": true, "light": true, - "isDefault": true + "isDefault": true, + "reasoning": true }, { "name": "Gemini 3.5 Flash Lite", @@ -258,7 +272,8 @@ ], "multiModal": true, "light": true, - "isDefault": false + "isDefault": false, + "reasoning": true }, { "name": "Gemini 3.1 Pro (Preview)", @@ -272,7 +287,8 @@ ], "multiModal": true, "light": false, - "isDefault": false + "isDefault": false, + "reasoning": true } ] } diff --git a/packages/agenty-core/pkg/infra/llm/anthropic.go b/packages/agenty-core/pkg/infra/llm/anthropic.go index 61bd8be..f8f2acf 100644 --- a/packages/agenty-core/pkg/infra/llm/anthropic.go +++ b/packages/agenty-core/pkg/infra/llm/anthropic.go @@ -103,7 +103,10 @@ func (caller *anthropicCaller) params(request modelRequest) (anthropic.MessageNe if err != nil { return anthropic.MessageNewParams{}, err } - effort := modelReasoningEffort(caller.model, request.ReasoningEffort) + effort, err := modelReasoningEffort(caller.model, request.ReasoningEffort) + if err != nil { + return anthropic.MessageNewParams{}, err + } messages := make([]anthropic.MessageParam, 0, len(request.Messages)) for index, message := range request.Messages { diff --git a/packages/agenty-core/pkg/infra/llm/convert.go b/packages/agenty-core/pkg/infra/llm/convert.go index d612f07..f188b0e 100644 --- a/packages/agenty-core/pkg/infra/llm/convert.go +++ b/packages/agenty-core/pkg/infra/llm/convert.go @@ -31,11 +31,19 @@ func validateRequest(request modelRequest) error { return nil } -func modelReasoningEffort(model catalog.Model, effort shared.ReasoningEffort) string { +func modelReasoningEffort(model catalog.Model, effort shared.ReasoningEffort) (string, error) { if effort == "" || effort == shared.ReasoningOff || !model.SupportsReasoning() { - return "" + return "", nil } - return string(effort) + if !model.SupportsReasoningEffort(effort) { + return "", fmt.Errorf( + "%w: model %q does not support effort %q", + ErrUnsupportedReasoningEffort, + model.Code, + effort, + ) + } + return string(effort), nil } func systemPrompt(request modelRequest) (string, error) { diff --git a/packages/agenty-core/pkg/infra/llm/convert_test.go b/packages/agenty-core/pkg/infra/llm/convert_test.go index 0c7415c..1f98bd6 100644 --- a/packages/agenty-core/pkg/infra/llm/convert_test.go +++ b/packages/agenty-core/pkg/infra/llm/convert_test.go @@ -38,7 +38,10 @@ func TestModelReasoningEffort(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - got := modelReasoningEffort(model, tt.effort) + got, err := modelReasoningEffort(model, tt.effort) + if err != nil { + t.Fatalf("modelReasoningEffort() error = %v", err) + } if got != tt.want { t.Fatalf("modelReasoningEffort() = %q, want %q", got, tt.want) } @@ -46,20 +49,26 @@ func TestModelReasoningEffort(t *testing.T) { } } -func TestModelReasoningEffortSendsUnsupportedLevelToUpstream(t *testing.T) { +func TestModelReasoningEffortRejectsUnsupportedLevel(t *testing.T) { model := catalog.Model{ Code: "gpt-5-mini", ReasoningEfforts: []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}, } - got := modelReasoningEffort(model, shared.ReasoningMax) - if got != "max" { - t.Fatalf("modelReasoningEffort() = %q, want max", got) + got, err := modelReasoningEffort(model, shared.ReasoningMax) + if got != "" { + t.Fatalf("modelReasoningEffort() = %q, want empty", got) + } + if !errors.Is(err, ErrUnsupportedReasoningEffort) { + t.Fatalf("modelReasoningEffort() error = %v, want unsupported effort", err) } } func TestModelReasoningEffortIgnoresNonReasoningModel(t *testing.T) { model := catalog.Model{Code: "gpt-4o", ReasoningEfforts: []shared.ReasoningEffort{}} - got := modelReasoningEffort(model, shared.ReasoningHigh) + got, err := modelReasoningEffort(model, shared.ReasoningHigh) + if err != nil { + t.Fatalf("modelReasoningEffort() error = %v", err) + } if got != "" { t.Fatalf("modelReasoningEffort() = %q, want empty", got) } diff --git a/packages/agenty-core/pkg/infra/llm/errors.go b/packages/agenty-core/pkg/infra/llm/errors.go index 55ba89e..dc8e457 100644 --- a/packages/agenty-core/pkg/infra/llm/errors.go +++ b/packages/agenty-core/pkg/infra/llm/errors.go @@ -6,9 +6,10 @@ import ( ) var ( - ErrInvalidRequest = errors.New("llm: invalid request") - ErrUnsupportedAPI = errors.New("llm: unsupported API type") - ErrUnsupportedContent = errors.New("llm: unsupported content") + ErrInvalidRequest = errors.New("llm: invalid request") + ErrUnsupportedAPI = errors.New("llm: unsupported API type") + ErrUnsupportedContent = errors.New("llm: unsupported content") + ErrUnsupportedReasoningEffort = errors.New("llm: unsupported reasoning effort") ) func invalidRequest(format string, args ...any) error { diff --git a/packages/agenty-core/pkg/infra/llm/google.go b/packages/agenty-core/pkg/infra/llm/google.go index 27ccdd4..a52d67e 100644 --- a/packages/agenty-core/pkg/infra/llm/google.go +++ b/packages/agenty-core/pkg/infra/llm/google.go @@ -88,7 +88,10 @@ func (caller *googleCaller) params(request modelRequest) ([]*genai.Content, *gen if err != nil { return nil, nil, err } - effort := modelReasoningEffort(caller.model, request.ReasoningEffort) + effort, err := modelReasoningEffort(caller.model, request.ReasoningEffort) + if err != nil { + return nil, nil, err + } toolNames := googleToolNames(request.Messages) contents := make([]*genai.Content, 0, len(request.Messages)) diff --git a/packages/agenty-core/pkg/infra/llm/openai_chat.go b/packages/agenty-core/pkg/infra/llm/openai_chat.go index f532ff2..6cb09da 100644 --- a/packages/agenty-core/pkg/infra/llm/openai_chat.go +++ b/packages/agenty-core/pkg/infra/llm/openai_chat.go @@ -116,7 +116,10 @@ func (caller *openAIChatCaller) params(request modelRequest) (openai.ChatComplet if err != nil { return openai.ChatCompletionNewParams{}, err } - effort := modelReasoningEffort(caller.model, request.ReasoningEffort) + effort, err := modelReasoningEffort(caller.model, request.ReasoningEffort) + if err != nil { + return openai.ChatCompletionNewParams{}, err + } messages := make([]openai.ChatCompletionMessageParamUnion, 0, len(request.Messages)+1) if prompt != "" { diff --git a/packages/agenty-core/pkg/infra/llm/openai_responses.go b/packages/agenty-core/pkg/infra/llm/openai_responses.go index f1b9ac4..84f2320 100644 --- a/packages/agenty-core/pkg/infra/llm/openai_responses.go +++ b/packages/agenty-core/pkg/infra/llm/openai_responses.go @@ -166,7 +166,10 @@ func (caller *openAIResponsesCaller) params(request modelRequest) (responses.Res if err != nil { return responses.ResponseNewParams{}, err } - effort := modelReasoningEffort(caller.model, request.ReasoningEffort) + effort, err := modelReasoningEffort(caller.model, request.ReasoningEffort) + if err != nil { + return responses.ResponseNewParams{}, err + } input, err := openAIResponsesMessages(request.Messages, caller.nativeOpenAI, caller.freeFormTool) if err != nil { diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister.go b/packages/agenty-core/pkg/infra/modelcatalog/lister.go index 81246bf..63e5fe7 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister.go @@ -462,6 +462,7 @@ func normalizeModel( ContextWindow: contextWindow, MaxOutputTokens: maxOutputTokens, MultiModal: multiModal, + Reasoning: len(reasoningEfforts) > 0, ReasoningEfforts: reasoningEfforts, }, nil } diff --git a/packages/agenty-core/pkg/infra/storage/catalog.go b/packages/agenty-core/pkg/infra/storage/catalog.go index c9f82d3..f190b3e 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog.go +++ b/packages/agenty-core/pkg/infra/storage/catalog.go @@ -226,9 +226,7 @@ func (r *CatalogRepository) ReplaceModels( normalized := slices.Clone(models) for index := range normalized { - normalized[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( - normalized[index].ReasoningEfforts, - ) + catalog.NormalizeReasoningCapabilities(&normalized[index]) if normalized[index].MaxOutputTokens <= 0 { normalized[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens } @@ -289,9 +287,7 @@ func normalizeModelsForCache(cache *modelDiscoveryCache) { cache.Models = make([]catalog.Model, 0) } for index := range cache.Models { - cache.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( - cache.Models[index].ReasoningEfforts, - ) + catalog.NormalizeReasoningCapabilities(&cache.Models[index]) if cache.Models[index].MaxOutputTokens <= 0 { cache.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens } @@ -315,9 +311,7 @@ func normalizeModels(provider *catalog.Provider) { provider.Models = make([]catalog.Model, 0) } for index := range provider.Models { - provider.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( - provider.Models[index].ReasoningEfforts, - ) + catalog.NormalizeReasoningCapabilities(&provider.Models[index]) if provider.Models[index].MaxOutputTokens <= 0 { provider.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens } @@ -362,9 +356,7 @@ func cloneProvider(provider *catalog.Provider) *catalog.Provider { copy.Models = make([]catalog.Model, len(provider.Models)) copy.Models = append(copy.Models[:0], provider.Models...) for index := range copy.Models { - copy.Models[index].ReasoningEfforts = shared.NormalizeReasoningEfforts( - provider.Models[index].ReasoningEfforts, - ) + catalog.NormalizeReasoningCapabilities(©.Models[index]) } copy.Metadata = maps.Clone(provider.Metadata) return © diff --git a/packages/agenty-core/pkg/infra/storage/catalog_test.go b/packages/agenty-core/pkg/infra/storage/catalog_test.go index 3ea6688..d0685f1 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog_test.go +++ b/packages/agenty-core/pkg/infra/storage/catalog_test.go @@ -129,6 +129,30 @@ func TestCatalogSaveAndGet(t *testing.T) { } } +func TestCatalogReasoningModelWithEmptyEffortsUsesDefaults(t *testing.T) { + repo := newCatalogRepo(t) + provider, err := catalog.NewProvider("custom", "Custom", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + provider.Models = []catalog.Model{{ + Code: mustCatalogModelCode("reasoning-model"), + Name: "Reasoning model", + Reasoning: true, + }} + if err := repo.Save(t.Context(), provider); err != nil { + t.Fatal(err) + } + + loaded, err := repo.Get(t.Context(), provider.Code) + if err != nil { + t.Fatal(err) + } + if !loaded.Models[0].Reasoning || len(loaded.Models[0].ReasoningEfforts) != len(shared.StandardReasoningEfforts()) { + t.Fatalf("reasoning capabilities = %+v", loaded.Models[0]) + } +} + func TestCatalogBuiltinProviderPersistsOnlyAPIKey(t *testing.T) { repo := newCatalogRepo(t) builtins, err := catalogdata.LoadProviders() diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index 3e4b607..7dc44a1 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -384,7 +384,17 @@ impl Transaction { operation.path.display() ))); } - let data = apply_update_diff(¤t.data, &operation.diff)?; + let data = if is_binary_data(¤t.data) { + if operation.diff.is_empty() { + current.data.clone() + } else { + return Err(PatchError::Invalid( + "binary file updates must not contain a text diff".to_string(), + )); + } + } else { + apply_update_diff(¤t.data, &operation.diff)? + }; self.state.insert( operation.path.clone(), VirtualFile::regular(data, current.mode), @@ -471,6 +481,15 @@ impl Transaction { if before.kind == after.kind && before.data == after.data { continue; } + if is_binary_snapshot(before) || is_binary_virtual(after) { + results.push(FileResult { + path: relative_display(&self.cwd, path), + diff: String::new(), + added_lines: 0, + removed_lines: 0, + }); + continue; + } let old_text = text_for_diff(before)?; let new_text = text_for_diff_virtual(after)?; let relative = relative_display(&self.cwd, path); @@ -951,6 +970,18 @@ fn text_for_diff(file: &FileSnapshot) -> Result { .map_err(|_| PatchError::Invalid("file is not valid UTF-8".to_string())) } +fn is_binary_data(data: &[u8]) -> bool { + std::str::from_utf8(data).is_err() +} + +fn is_binary_snapshot(file: &FileSnapshot) -> bool { + file.kind == EntryKind::Regular && is_binary_data(&file.data) +} + +fn is_binary_virtual(file: &VirtualFile) -> bool { + file.kind == EntryKind::Regular && is_binary_data(&file.data) +} + fn text_for_diff_virtual(file: &VirtualFile) -> Result { if file.kind == EntryKind::Missing { return Ok(String::new()); @@ -1274,6 +1305,38 @@ mod tests { ); } + #[test] + fn deletes_binary_files_without_text_diff() { + let cwd = temp_dir("binary-delete"); + fs::write(cwd.join("image.bin"), [0, 159, 146, 150]).unwrap(); + + let patch = "*** Begin Patch\n*** Delete File: image.bin\n*** End Patch"; + let result = apply_patch(&cwd, patch).unwrap(); + + assert!(result.success); + assert!(!cwd.join("image.bin").exists()); + assert_eq!(result.files[0].diff, ""); + assert_eq!(result.files[0].added_lines, 0); + assert_eq!(result.files[0].removed_lines, 0); + } + + #[test] + fn moves_binary_files_without_text_diff() { + let cwd = temp_dir("binary-move"); + fs::write(cwd.join("old.bin"), [0, 159, 146, 150]).unwrap(); + + let patch = + "*** Begin Patch\n*** Update File: old.bin\n*** Move to: new.bin\n*** End Patch"; + let result = apply_patch(&cwd, patch).unwrap(); + + assert!(result.success); + assert!(!cwd.join("old.bin").exists()); + assert_eq!(fs::read(cwd.join("new.bin")).unwrap(), [0, 159, 146, 150]); + assert!(result.files.iter().all(|file| { + file.diff.is_empty() && file.added_lines == 0 && file.removed_lines == 0 + })); + } + #[test] fn preserves_v4a_diff_compatibility() { let create_cases = [ From 27dd55a7fc461cb1ecd20b9fadcda24b0c81455d Mon Sep 17 00:00:00 2001 From: masteryyh Date: Wed, 26 Aug 2026 10:28:17 +0800 Subject: [PATCH 08/12] feat: add apply_patch file lock Signed-off-by: masteryyh --- packages/agenty-cli/src/cli/init.ts | 6 +- packages/agenty-cli/src/cli/model.ts | 1 + .../agenty-cli/src/components/FormPanel.tsx | 2 - .../src/components/WizardOverlay.tsx | 1 + .../src/components/wizardSetup.test.ts | 26 ++- .../agenty-cli/src/components/wizardSetup.ts | 49 +++- .../agenty-cli/src/consts/providerPresets.ts | 11 +- .../pkg/agentloop/builtin/apply_patch.go | 9 +- .../pkg/agentloop/builtin/apply_patch_test.go | 40 ++++ .../pkg/agentloop/builtin/common.go | 155 ++++++++++++- .../pkg/agentloop/builtin/common_test.go | 78 +++++++ .../agenty-core/pkg/agentloop/builtin/file.go | 4 +- .../pkg/agentloop/builtin/file_test.go | 101 +++++++++ .../pkg/agentloop/builtin/shell.go | 18 +- .../pkg/agentloop/builtin/shell_test.go | 18 +- .../agenty-core/pkg/agentloop/engine_test.go | 2 +- .../agenty-core/pkg/application/provider.go | 3 + .../pkg/application/provider_test.go | 8 + .../agenty-core/pkg/domain/agent/agent.go | 18 +- .../pkg/domain/agent/agent_test.go | 9 +- .../pkg/infra/modelcatalog/lister.go | 6 +- .../pkg/infra/modelcatalog/lister_test.go | 6 +- .../agenty-core/pkg/infra/storage/catalog.go | 209 ++++++++++++------ .../pkg/infra/storage/catalog_test.go | 76 +++++-- packages/patch-applier/src/lib.rs | 194 +++++++++++++++- 25 files changed, 927 insertions(+), 123 deletions(-) create mode 100644 packages/agenty-core/pkg/agentloop/builtin/common_test.go diff --git a/packages/agenty-cli/src/cli/init.ts b/packages/agenty-cli/src/cli/init.ts index f0dc1e0..e665633 100644 --- a/packages/agenty-cli/src/cli/init.ts +++ b/packages/agenty-cli/src/cli/init.ts @@ -5,6 +5,7 @@ import type { APIType } from "@/api/types"; import { CliError, flag, + hasFlag, outputFields, type ParsedArgs, render, @@ -20,13 +21,16 @@ export async function handleInit(client: AgentyClient, args: ParsedArgs): Promis const agentCode = flag(args, "agent")?.trim() || "default"; const contextWindow = positiveInteger(flag(args, "context-window") ?? "128000", "--context-window"); const apiKey = secret(args, "api-key", "api-key-env", "provider API key") ?? ""; + const apiKeyProvided = hasFlag(args, "api-key") || hasFlag(args, "api-key-env"); const providers = await client.listProviders(); const existingProvider = providers.find((provider) => provider.code === providerCode); const existingModel = existingProvider?.models.find((model) => model.code === modelCode); const effectiveContextWindow = existingModel?.contextWindow ?? contextWindow; if (existingProvider?.builtin) { - await client.updateProvider(providerCode, { apiKey }); + if (apiKeyProvided) { + await client.updateProvider(providerCode, { apiKey }); + } } else { await client.createProvider({ code: providerCode, diff --git a/packages/agenty-cli/src/cli/model.ts b/packages/agenty-cli/src/cli/model.ts index a5a9b12..c120d2d 100644 --- a/packages/agenty-cli/src/cli/model.ts +++ b/packages/agenty-cli/src/cli/model.ts @@ -73,6 +73,7 @@ export async function handleModel(client: AgentyClient, args: ParsedArgs): Promi const update: UpdateModelDto = { name: hasFlag(args, "name") ? requireFlag(args, "name") : current.name, contextWindow: hasFlag(args, "context-window") ? positiveInteger(requireFlag(args, "context-window"), "--context-window", true) : current.contextWindow, + maxOutputTokens: current.maxOutputTokens, multiModal: hasFlag(args, "multi-modal") ? booleanFlag(args, "multi-modal") : current.multiModal, light: hasFlag(args, "light") ? booleanFlag(args, "light") : current.light, isDefault: hasFlag(args, "default") ? booleanFlag(args, "default") : current.isDefault, diff --git a/packages/agenty-cli/src/components/FormPanel.tsx b/packages/agenty-cli/src/components/FormPanel.tsx index d6c0106..ad3f9fa 100644 --- a/packages/agenty-cli/src/components/FormPanel.tsx +++ b/packages/agenty-cli/src/components/FormPanel.tsx @@ -448,8 +448,6 @@ export function FormPanel({ const layout = fieldLayouts[visibleIndex]; const rowHeight = layout?.height ?? 1; const labelLines = layout?.labelLines ?? [field.label]; - const choiceOpen = choice.kind !== "idle" && choice.visibleIndex === visibleIndex; - return ( { + calls.push("provider.listModels"); + return []; + }, listAgents: async () => { calls.push("agent.list"); return agents; @@ -224,7 +228,6 @@ describe("first-run provider setup", () => { expect(client.calls).toEqual([ "provider.list", "agent.list", - "provider.update", "provider.addModel", "agent.update", "initialize.complete", @@ -244,7 +247,6 @@ describe("first-run provider setup", () => { expect(client.calls).toEqual([ "provider.list", "agent.list", - "provider.update", "provider.removeModel", "provider.addModel", "provider.addModel", @@ -282,9 +284,27 @@ describe("first-run provider setup", () => { expect(client.calls).toEqual([ "provider.list", "agent.list", - "provider.update", "agent.create", "initialize.complete", ]); }); + + test("does not persist an untouched discovered model cache", async () => { + const draft = createDraft(); + const provider = { ...createProvider(draft), modelsCached: true }; + const models = modelDraftsForProvider(draft, provider); + const client = fakeClient([provider], [createAgent("default", true)]); + + expect(models[0].source).toBe("cached"); + await persistWizardSetup(client, [draft], models, selectedModelId(models[0])); + + expect(client.createdModels).toEqual([]); + expect(client.deletedModels).toEqual([]); + expect(client.calls).toEqual([ + "provider.list", + "agent.list", + "agent.update", + "initialize.complete", + ]); + }); }); diff --git a/packages/agenty-cli/src/components/wizardSetup.ts b/packages/agenty-cli/src/components/wizardSetup.ts index afd9c0b..fb781fb 100644 --- a/packages/agenty-cli/src/components/wizardSetup.ts +++ b/packages/agenty-cli/src/components/wizardSetup.ts @@ -20,6 +20,7 @@ const DEFAULT_AGENT_SOUL = "Be helpful, concise, and accurate."; export interface WizardSetupClient { listProviders(): Promise; + listProviderModels(providerCode: string): Promise; listAgents(): Promise; createProvider(input: CreateModelProviderDto): Promise; updateProvider(code: string, input: UpdateModelProviderDto): Promise; @@ -128,16 +129,27 @@ export async function persistWizardSetup( for (const draft of drafts) { const providerCode = draft.code.trim(); const existing = existingProviders.find((provider) => provider.code === providerCode); + let providerChanged = false; if (draft.builtin) { - await client.updateProvider(providerCode, { apiKey: draft.apiKey.trim() }); + providerChanged = existing?.apiKey !== draft.apiKey.trim(); + if (providerChanged) { + await client.updateProvider(providerCode, { apiKey: draft.apiKey.trim() }); + } } else if (existing) { - await client.updateProvider(providerCode, { - name: draft.name.trim(), - type: draft.type, - baseUrl: draft.baseUrl.trim(), - apiKey: draft.apiKey.trim(), - freeFormTool: draft.type === "openai" && draft.freeFormTool, - }); + providerChanged = existing.name !== draft.name.trim() || + existing.type !== draft.type || + existing.baseUrl !== draft.baseUrl.trim() || + existing.apiKey !== draft.apiKey.trim() || + (existing.freeFormTool === true) !== (draft.type === "openai" && draft.freeFormTool); + if (providerChanged) { + await client.updateProvider(providerCode, { + name: draft.name.trim(), + type: draft.type, + baseUrl: draft.baseUrl.trim(), + apiKey: draft.apiKey.trim(), + freeFormTool: draft.type === "openai" && draft.freeFormTool, + }); + } } else { const providerInput: CreateModelProviderDto = { code: providerCode, @@ -150,16 +162,25 @@ export async function persistWizardSetup( await client.createProvider(providerInput); } + if (existing?.modelsCached === true && providerChanged) { + await client.listProviderModels(providerCode); + } + if (!draft.builtin) { const providerModels = modelsByProvider.get(draft.id) ?? []; - const desiredModelCodes = new Set(providerModels.map((model) => model.code.trim())); - for (const model of existing?.models ?? []) { - if (!desiredModelCodes.has(model.code)) { - await client.deleteModel(providerCode, model.code); + if (existing?.modelsCached !== true) { + const desiredModelCodes = new Set(providerModels.map((model) => model.code.trim())); + for (const model of existing?.models ?? []) { + if (!desiredModelCodes.has(model.code)) { + await client.deleteModel(providerCode, model.code); + } } } for (const model of providerModels) { + if (model.source === "cached") { + continue; + } await client.createModel({ providerCode, modelCode: model.code.trim(), @@ -173,6 +194,10 @@ export async function persistWizardSetup( isDefault: selectedModelId(model) === selectedId, }); } + + if (existing?.modelsCached === true && !providerChanged && providerModels.some((model) => model.source !== "cached")) { + await client.listProviderModels(providerCode); + } } } diff --git a/packages/agenty-cli/src/consts/providerPresets.ts b/packages/agenty-cli/src/consts/providerPresets.ts index 21141f0..b009917 100644 --- a/packages/agenty-cli/src/consts/providerPresets.ts +++ b/packages/agenty-cli/src/consts/providerPresets.ts @@ -26,6 +26,7 @@ export interface ProviderDraft { } export interface ModelDraft { + source: "configured" | "cached" | "new"; id: string; providerId: string; providerCode: string; @@ -84,8 +85,10 @@ export function createModelDraft( provider: ProviderDraft, id: string, existing?: CoreModelDto, + source: ModelDraft["source"] = existing === undefined ? "new" : "configured", ): ModelDraft { return { + source, id, providerId: provider.id, providerCode: provider.code, @@ -112,8 +115,14 @@ export function modelDraftsForProvider( provider: ProviderDraft, existing?: ModelProviderDto, ): ModelDraft[] { + const source: ModelDraft["source"] = existing?.modelsCached === true ? "cached" : "configured"; return (existing?.models ?? []).map((model) => - createModelDraft(provider, `${provider.id}:model:${model.code}`, model), + createModelDraft( + provider, + `${provider.id}:model:${model.code}`, + model, + source, + ), ); } diff --git a/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go b/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go index 5529c99..56d28bb 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go +++ b/packages/agenty-core/pkg/agentloop/builtin/apply_patch.go @@ -3,7 +3,6 @@ package builtin import ( "bytes" "context" - "errors" "fmt" "os/exec" "strings" @@ -65,7 +64,10 @@ func (tool *applyPatchTool) Execute( tool.fileSystem.mu.Lock() defer tool.fileSystem.mu.Unlock() - command := exec.CommandContext(ctx, "apply_patch") + if err := ctx.Err(); err != nil { + return nil, err + } + command := exec.Command("apply_patch") if strings.TrimSpace(callContext.Cwd) != "" { command.Dir = callContext.Cwd } @@ -75,9 +77,6 @@ func (tool *applyPatchTool) Execute( command.Stdout = &stdout command.Stderr = &stderr if err := command.Run(); err != nil { - if errors.Is(ctx.Err(), context.Canceled) { - return nil, ctx.Err() - } message := strings.TrimSpace(stderr.String()) if message == "" { message = strings.TrimSpace(stdout.String()) diff --git a/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go b/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go index 5083b9a..d63f892 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/apply_patch_test.go @@ -7,6 +7,7 @@ import ( "runtime" "strings" "testing" + "time" json "github.com/bytedance/sonic" @@ -70,6 +71,45 @@ func TestApplyPatchToolReportsHelperFailure(t *testing.T) { } } +func TestApplyPatchToolLetsStartedHelperFinishAfterCancellation(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("shell fixture uses a POSIX script") + } + + directory := t.TempDir() + started := filepath.Join(directory, "started") + installApplyPatchFixture(t, directory, "#!/bin/sh\ntouch started\nsleep 0.2\nprintf '%s\\n' '{\"success\":true,\"cwd\":\"/workspace\",\"files\":[]}'\n") + t.Setenv("PATH", directory+string(os.PathListSeparator)+os.Getenv("PATH")) + + tool := &applyPatchTool{fileSystem: &fileSystem{}} + input, err := json.Marshal(applyPatchArguments{Patch: "*** Begin Patch\n*** End Patch"}) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + result := make(chan error, 1) + go func() { + _, executeErr := tool.Execute(ctx, agentloop.CallContext{Cwd: directory}, input) + result <- executeErr + }() + + deadline := time.Now().Add(time.Second) + for { + if _, statErr := os.Stat(started); statErr == nil { + break + } + if time.Now().After(deadline) { + t.Fatal("apply_patch helper did not start") + } + time.Sleep(time.Millisecond) + } + cancel() + if executeErr := <-result; executeErr != nil { + t.Fatalf("started helper returned error after cancellation: %v", executeErr) + } +} + func TestApplyPatchToolRequiresPatch(t *testing.T) { tool := &applyPatchTool{fileSystem: &fileSystem{}} _, err := tool.Execute(context.Background(), agentloop.CallContext{}, []byte(`{}`)) diff --git a/packages/agenty-core/pkg/agentloop/builtin/common.go b/packages/agenty-core/pkg/agentloop/builtin/common.go index ecf4c2a..5c9a257 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/common.go +++ b/packages/agenty-core/pkg/agentloop/builtin/common.go @@ -48,13 +48,21 @@ func resolvePath(path, cwd string, allowEmpty bool) (string, error) { path = "." } - if filepath.IsAbs(path) { - return filepath.Clean(path), nil + path, err := expandEnvironmentVariables(path) + if err != nil { + return "", err + } + + absolutePath, isAbsolute, err := normalizeAbsolutePath(path) + if err != nil { + return "", err + } + if isAbsolute { + return filepath.Clean(absolutePath), nil } base := cwd if strings.TrimSpace(base) == "" { - var err error base, err = os.Getwd() if err != nil { return "", fmt.Errorf("resolve process working directory: %w", err) @@ -68,6 +76,147 @@ func resolvePath(path, cwd string, allowEmpty bool) (string, error) { return filepath.Clean(resolved), nil } +func normalizeAbsolutePath(path string) (string, bool, error) { + if strings.HasPrefix(path, "~") { + home, err := os.UserHomeDir() + if err != nil { + return "", false, fmt.Errorf("resolve user home directory: %w", err) + } + return filepath.Join(home, strings.TrimLeft(path[1:], `/\\`)), true, nil + } + + return path, filepath.IsAbs(path) || isWindowsAbsolutePath(path), nil +} + +func expandEnvironmentVariables(path string) (string, error) { + path, err := expandWindowsStyleEnvironmentVariables(path) + if err != nil { + return "", err + } + + path, err = expandPowerShellEnvironmentVariables(path) + if err != nil { + return "", err + } + + return expandPOSIXEnvironmentVariables(path) +} + +func expandWindowsStyleEnvironmentVariables(path string) (string, error) { + var builder strings.Builder + for offset := 0; offset < len(path); { + opening := strings.IndexByte(path[offset:], '%') + if opening < 0 { + builder.WriteString(path[offset:]) + break + } + + start := offset + opening + builder.WriteString(path[offset:start]) + closing := strings.IndexByte(path[start+1:], '%') + if closing < 0 { + builder.WriteString(path[start:]) + break + } + + end := start + closing + 1 + variable := path[start+1 : end] + if variable == "" { + builder.WriteString("%%") + } else { + value, err := environmentVariable(variable) + if err != nil { + return "", err + } + builder.WriteString(value) + } + offset = end + 1 + } + return builder.String(), nil +} + +func expandPowerShellEnvironmentVariables(path string) (string, error) { + const prefix = "$env:" + + var builder strings.Builder + for offset := 0; offset < len(path); { + start := powerShellEnvironmentVariableStart(path, offset, prefix) + if start < 0 { + builder.WriteString(path[offset:]) + break + } + + builder.WriteString(path[offset:start]) + variableStart := start + len(prefix) + variableEnd := variableStart + for variableEnd < len(path) && path[variableEnd] != '/' && path[variableEnd] != '\\' { + variableEnd++ + } + if variableEnd == variableStart { + return "", fmt.Errorf("resolve PowerShell environment variable: name must not be empty") + } + + value, err := environmentVariable(path[variableStart:variableEnd]) + if err != nil { + return "", err + } + builder.WriteString(value) + offset = variableEnd + } + return builder.String(), nil +} + +func powerShellEnvironmentVariableStart(path string, offset int, prefix string) int { + for index := offset; index+len(prefix) <= len(path); index++ { + if strings.EqualFold(path[index:index+len(prefix)], prefix) { + return index + } + } + return -1 +} + +func expandPOSIXEnvironmentVariables(path string) (string, error) { + var expandErr error + expanded := os.Expand(path, func(name string) string { + if expandErr != nil { + return "" + } + if name == "" { + expandErr = fmt.Errorf("resolve environment variable: name must not be empty") + return "" + } + + value, err := environmentVariable(name) + if err != nil { + expandErr = err + return "" + } + return value + }) + if expandErr != nil { + return "", expandErr + } + return expanded, nil +} + +func environmentVariable(name string) (string, error) { + value, found := os.LookupEnv(name) + if !found { + return "", fmt.Errorf("resolve environment variable %q: not set", name) + } + return value, nil +} + +func isWindowsAbsolutePath(path string) bool { + if strings.HasPrefix(path, `\\`) { + return true + } + + return len(path) >= 2 && + ((path[0] >= 'a' && path[0] <= 'z') || (path[0] >= 'A' && path[0] <= 'Z')) && + path[1] == ':' +} + func regularFileInfo(path string) (os.FileInfo, error) { info, err := os.Stat(path) if err != nil { diff --git a/packages/agenty-core/pkg/agentloop/builtin/common_test.go b/packages/agenty-core/pkg/agentloop/builtin/common_test.go new file mode 100644 index 0000000..a6538e6 --- /dev/null +++ b/packages/agenty-core/pkg/agentloop/builtin/common_test.go @@ -0,0 +1,78 @@ +package builtin + +import ( + "path/filepath" + "testing" +) + +func TestResolvePathExpandsEnvironmentVariablesBeforeCheckingAbsolutePaths(t *testing.T) { + cwd := filepath.Join(t.TempDir(), "cwd") + absoluteRoot := filepath.Join(t.TempDir(), "absolute") + t.Setenv("AGENTY_TEST_ABSOLUTE_ROOT", absoluteRoot) + t.Setenv("AGENTY_TEST_RELATIVE_ROOT", "nested") + t.Setenv("AGENTY_TEST_WINDOWS_ROOT", `C:\workspace`) + + tests := []struct { + name string + path string + want string + }{ + { + name: "drive letter with backslashes", + path: `C:\workspace\file.txt`, + want: filepath.Clean(`C:\workspace\file.txt`), + }, + { + name: "drive letter with slashes", + path: `C:/workspace/file.txt`, + want: filepath.Clean(`C:/workspace/file.txt`), + }, + { + name: "UNC path", + path: `\\server\share\file.txt`, + want: filepath.Clean(`\\server\share\file.txt`), + }, + { + name: "braced POSIX environment variable", + path: `${AGENTY_TEST_ABSOLUTE_ROOT}/file.txt`, + want: filepath.Join(absoluteRoot, "file.txt"), + }, + { + name: "POSIX environment variable", + path: `$AGENTY_TEST_ABSOLUTE_ROOT/file.txt`, + want: filepath.Join(absoluteRoot, "file.txt"), + }, + { + name: "Windows environment variable", + path: `%AGENTY_TEST_ABSOLUTE_ROOT%/file.txt`, + want: filepath.Join(absoluteRoot, "file.txt"), + }, + { + name: "PowerShell environment variable", + path: `$env:AGENTY_TEST_ABSOLUTE_ROOT/file.txt`, + want: filepath.Join(absoluteRoot, "file.txt"), + }, + { + name: "relative environment variable", + path: `${AGENTY_TEST_RELATIVE_ROOT}/file.txt`, + want: filepath.Join(cwd, "nested", "file.txt"), + }, + { + name: "Windows environment variable expands to a drive path", + path: `%AGENTY_TEST_WINDOWS_ROOT%\file.txt`, + want: filepath.Clean(`C:\workspace\file.txt`), + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := resolvePath(test.path, cwd, false) + if err != nil { + t.Fatal(err) + } + if got != test.want { + t.Errorf("resolvePath(%q) = %q, want %q", test.path, got, test.want) + } + }) + } +} diff --git a/packages/agenty-core/pkg/agentloop/builtin/file.go b/packages/agenty-core/pkg/agentloop/builtin/file.go index 9422c93..342a6da 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/file.go +++ b/packages/agenty-core/pkg/agentloop/builtin/file.go @@ -35,7 +35,9 @@ func (tool *readFileTool) Definition() agentloop.ToolDefinition { return agentloop.ToolDefinition{ Name: "read_file", Description: "Read a text file with optional inclusive 1-based line bounds. " + - "Relative paths resolve from the session working directory. The result contains numbered lines.", + "Environment variables are expanded before absolute-path detection. Absolute paths are read directly; " + + "relative paths resolve from the session working directory. " + + "The result contains numbered lines.", InputSchema: objectSchema( map[string]agentloop.JSONSchema{ "path": stringSchema("Absolute path or path relative to the session working directory."), diff --git a/packages/agenty-core/pkg/agentloop/builtin/file_test.go b/packages/agenty-core/pkg/agentloop/builtin/file_test.go index 0eaf858..65a04bb 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/file_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/file_test.go @@ -3,8 +3,11 @@ package builtin_test import ( "os" "path/filepath" + "runtime" "strings" "testing" + + json "github.com/bytedance/sonic" ) func TestReadFile(t *testing.T) { @@ -47,6 +50,104 @@ func TestReadFile(t *testing.T) { } } +func TestReadFileAbsolutePaths(t *testing.T) { + directory := t.TempDir() + + absolutePath := filepath.Join(t.TempDir(), "absolute.txt") + if err := os.WriteFile(absolutePath, []byte("absolute"), 0o644); err != nil { + t.Fatal(err) + } + + homeDirectory := t.TempDir() + homeVariable := "HOME" + if runtime.GOOS == "windows" { + homeVariable = "USERPROFILE" + } + t.Setenv(homeVariable, homeDirectory) + + homePath := filepath.Join(homeDirectory, "from-home.txt") + if err := os.WriteFile(homePath, []byte("home"), 0o644); err != nil { + t.Fatal(err) + } + + environmentDirectory := t.TempDir() + t.Setenv("AGENTY_READ_FILE_ROOT", environmentDirectory) + environmentPath := filepath.Join(environmentDirectory, "from-environment.txt") + if err := os.WriteFile(environmentPath, []byte("environment"), 0o644); err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + path string + wantPath string + wantContents string + }{ + { + name: "native absolute path", + path: absolutePath, + wantPath: absolutePath, + wantContents: "absolute", + }, + { + name: "home shorthand", + path: "~/from-home.txt", + wantPath: homePath, + wantContents: "home", + }, + { + name: "braced POSIX environment variable", + path: "${AGENTY_READ_FILE_ROOT}/from-environment.txt", + wantPath: environmentPath, + wantContents: "environment", + }, + { + name: "POSIX environment variable", + path: "$AGENTY_READ_FILE_ROOT/from-environment.txt", + wantPath: environmentPath, + wantContents: "environment", + }, + { + name: "Windows environment variable", + path: "%AGENTY_READ_FILE_ROOT%/from-environment.txt", + wantPath: environmentPath, + wantContents: "environment", + }, + { + name: "PowerShell environment variable", + path: "$env:AGENTY_READ_FILE_ROOT/from-environment.txt", + wantPath: environmentPath, + wantContents: "environment", + }, + } + + registry := newRegistry(t) + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + encodedPath, err := json.MarshalString(test.path) + if err != nil { + t.Fatal(err) + } + + encoded, err := executeTool(t, registry, "read_file", directory, `{"path":`+encodedPath+`}`) + if err != nil { + t.Fatal(err) + } + + result := decodeResult[struct { + Path string `json:"path"` + Content string `json:"content"` + }](t, encoded) + if result.Path != test.wantPath { + t.Errorf("path = %q, want %q", result.Path, test.wantPath) + } + if result.Content != "1: "+test.wantContents { + t.Errorf("content = %q, want %q", result.Content, "1: "+test.wantContents) + } + }) + } +} + func TestReadFileValidatesLineBounds(t *testing.T) { t.Parallel() diff --git a/packages/agenty-core/pkg/agentloop/builtin/shell.go b/packages/agenty-core/pkg/agentloop/builtin/shell.go index 42e50ed..6042eff 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/shell.go +++ b/packages/agenty-core/pkg/agentloop/builtin/shell.go @@ -27,6 +27,7 @@ const ( type shellArguments struct { Commands []string `json:"commands"` + Stdin *string `json:"stdin,omitempty"` TimeoutMs *int64 `json:"timeout_ms,omitempty"` MaxOutputLength *int64 `json:"max_output_length,omitempty"` } @@ -37,9 +38,9 @@ func (tool *shellTool) Definition() agentloop.ToolDefinition { return agentloop.ToolDefinition{ Type: agentloop.ToolTypeShell, Name: "shell", - Description: "Execute up to 4 shell commands in parallel. " + + Description: "Execute up to 4 independent, complete shell commands in parallel. The commands array contains separate commands, not fragments of one command. Do not run commands in parallel when they perform the same kind of operation or modify the same file; combine dependent steps into one command. " + "Uses zsh on macOS, bash on Linux, and sh as a fallback when the preferred shell is unavailable. " + - "On Windows, uses pwsh.exe or powershell.exe when available, then cmd.exe.", + "On Windows, uses pwsh.exe or powershell.exe when available, then cmd.exe. Use stdin only with one command when the command reads patch data.", InputSchema: objectSchema(map[string]agentloop.JSONSchema{ "commands": { Type: agentloop.JSONSchemaTypeArray, @@ -47,6 +48,10 @@ func (tool *shellTool) Definition() agentloop.ToolDefinition { MaxItems: new(uint64(maxShellCommands)), Items: &agentloop.JSONSchema{Type: agentloop.JSONSchemaTypeString}, }, + "stdin": { + Type: agentloop.JSONSchemaTypeString, + Description: "Optional standard input for the single command that reads it, such as cmd.exe apply_patch.", + }, "timeout_ms": { Type: agentloop.JSONSchemaTypeInteger, Description: "Maximum wall-clock time in milliseconds for each command.", @@ -76,6 +81,9 @@ func (tool *shellTool) Execute( if err != nil { return nil, err } + if arguments.Stdin != nil && len(arguments.Commands) != 1 { + return nil, fmt.Errorf("stdin requires exactly one command") + } results := make([]conversation.ShellCommandOutput, len(arguments.Commands)) jobs := make(chan shellJob) @@ -86,7 +94,7 @@ func (tool *shellTool) Execute( go func() { defer waitGroup.Done() for job := range jobs { - results[job.index] = executeShellCommand(ctx, callContext.Cwd, job.command, timeout, outputLimit) + results[job.index] = executeShellCommand(ctx, callContext.Cwd, job.command, arguments.Stdin, timeout, outputLimit) } }() } @@ -143,6 +151,7 @@ func executeShellCommand( parent context.Context, cwd string, command string, + stdin *string, timeout time.Duration, outputLimit int64, ) conversation.ShellCommandOutput { @@ -151,6 +160,9 @@ func executeShellCommand( process := newShellCommand(commandContext, command) prepareShellProcess(process) + if stdin != nil { + process.Stdin = strings.NewReader(*stdin) + } if strings.TrimSpace(cwd) != "" { process.Dir = cwd } diff --git a/packages/agenty-core/pkg/agentloop/builtin/shell_test.go b/packages/agenty-core/pkg/agentloop/builtin/shell_test.go index 7ff486d..83bd161 100644 --- a/packages/agenty-core/pkg/agentloop/builtin/shell_test.go +++ b/packages/agenty-core/pkg/agentloop/builtin/shell_test.go @@ -2,7 +2,9 @@ package builtin_test import ( "context" + "fmt" "os" + "runtime" "strings" "testing" "time" @@ -94,7 +96,7 @@ func TestShellDefinitionDocumentsRuntimeAndCommandLimit(t *testing.T) { if maxItems == nil || *maxItems != wantMaxItems { t.Fatalf("commands maxItems = %v, want %d", maxItems, wantMaxItems) } - for _, phrase := range []string{"zsh on macOS", "bash on Linux", "sh as a fallback"} { + for _, phrase := range []string{"independent, complete shell commands in parallel", "commands array contains separate commands", "same file", "zsh on macOS", "bash on Linux", "sh as a fallback"} { if !strings.Contains(definition.Description, phrase) { t.Errorf("description %q does not mention %q", definition.Description, phrase) } @@ -115,6 +117,7 @@ func TestShellRejectsInvalidArguments(t *testing.T) { `{"commands":["true"],"timeout_ms":0}`, `{"commands":["true"],"max_output_length":0}`, `{"commands":["true","true","true","true","true"]}`, + `{"commands":["true","true"],"stdin":"patch"}`, } { if _, err := tool.Execute(context.Background(), agentloop.CallContext{}, []byte(input)); err == nil { t.Errorf("Execute(%s) succeeded", input) @@ -122,6 +125,19 @@ func TestShellRejectsInvalidArguments(t *testing.T) { } } +func TestShellPassesStdinToSingleCommand(t *testing.T) { + t.Parallel() + + command := "cat" + if runtime.GOOS == "windows" { + command = "more" + } + output := executeShell(t, fmt.Sprintf(`{"commands":[%q],"stdin":"patch input"}`, command)) + if output.Output[0].Stdout != "patch input" { + t.Fatalf("stdin output = %q, want patch input", output.Output[0].Stdout) + } +} + func TestShellReportsProcessStartErrors(t *testing.T) { t.Parallel() diff --git a/packages/agenty-core/pkg/agentloop/engine_test.go b/packages/agenty-core/pkg/agentloop/engine_test.go index f4bc02e..3b9f419 100644 --- a/packages/agenty-core/pkg/agentloop/engine_test.go +++ b/packages/agenty-core/pkg/agentloop/engine_test.go @@ -458,7 +458,7 @@ func TestEngineProjectsApplyPatchByProviderCapability(t *testing.T) { if gotApplyPatch != test.wantApplyPatch { t.Errorf("apply_patch registered = %v, want %v", gotApplyPatch, test.wantApplyPatch) } - gotShellPrompt := strings.Contains(requests[0].SystemPrompt, "run the apply_patch command") + gotShellPrompt := strings.Contains(requests[0].SystemPrompt, "shell tool with one complete apply_patch command") if gotShellPrompt != test.wantShellPrompt { t.Errorf("shell fallback prompt present = %v, want %v", gotShellPrompt, test.wantShellPrompt) } diff --git a/packages/agenty-core/pkg/application/provider.go b/packages/agenty-core/pkg/application/provider.go index d3559d0..931bc8a 100644 --- a/packages/agenty-core/pkg/application/provider.go +++ b/packages/agenty-core/pkg/application/provider.go @@ -350,6 +350,9 @@ func (s *ProviderService) AddModel(ctx context.Context, providerCode, modelCode if maxOutputTokens <= 0 { maxOutputTokens = catalog.DefaultMaxOutputTokens } + if in.ContextWindow > 0 && maxOutputTokens >= int64(in.ContextWindow) { + return nil, Validation("max output tokens must be less than context window") + } reasoning := true if in.Reasoning != nil { reasoning = *in.Reasoning diff --git a/packages/agenty-core/pkg/application/provider_test.go b/packages/agenty-core/pkg/application/provider_test.go index dae0a18..465556f 100644 --- a/packages/agenty-core/pkg/application/provider_test.go +++ b/packages/agenty-core/pkg/application/provider_test.go @@ -320,6 +320,14 @@ func TestProviderAddModelAndRemoveModel(t *testing.T) { t.Fatal(err) } + if _, err := providerSvc.AddModel(ctx, "anthropic", "too-large", application.ModelInput{ + Name: "Too large", + ContextWindow: 32_000, + MaxOutputTokens: 32_000, + }); appErrorCode(err) != application.CodeValidation { + t.Fatalf("equal output/context error = %v, want validation", err) + } + p, err := providerSvc.AddModel(ctx, "anthropic", "claude-opus-4-8", application.ModelInput{ Name: "Claude Opus 4.8", ContextWindow: 200_000, diff --git a/packages/agenty-core/pkg/domain/agent/agent.go b/packages/agenty-core/pkg/domain/agent/agent.go index 6131c3e..b58be31 100644 --- a/packages/agenty-core/pkg/domain/agent/agent.go +++ b/packages/agenty-core/pkg/domain/agent/agent.go @@ -30,7 +30,23 @@ You will receive this at the very beginning of the session, and maybe more after {{ if .UseApplyPatchShell }} -The current provider does not support the free-form apply_patch tool. For every file modification, call the shell tool and run the apply_patch command with a complete V4A patch envelope passed through a heredoc on stdin. Do not use cat, sed, printf, or ad hoc scripts to edit files. +The current provider does not support the free-form apply_patch tool. For every file modification, call the shell tool with one complete apply_patch command and a complete V4A patch envelope. + +On macOS/Linux, pass the patch through a POSIX heredoc: +apply_patch <<'PATCH' +*** Begin Patch +... +*** End Patch +PATCH + +On PowerShell, pass it through a literal here-string: +@' +*** Begin Patch +... +*** End Patch +'@ | apply_patch + +If Windows falls back to cmd.exe, call the shell tool with the single command "apply_patch" and pass the complete patch in its stdin field. Do not use cat, sed, printf, echo, or ad hoc scripts to edit files. The shell tool runs commands in parallel, so never put dependent edits, the same kind of operation, or edits to the same file in parallel commands. {{ end }} diff --git a/packages/agenty-core/pkg/domain/agent/agent_test.go b/packages/agenty-core/pkg/domain/agent/agent_test.go index 53d59c1..b2fbecd 100644 --- a/packages/agenty-core/pkg/domain/agent/agent_test.go +++ b/packages/agenty-core/pkg/domain/agent/agent_test.go @@ -39,7 +39,7 @@ func TestAgent_ResolveSystemPrompt(t *testing.T) { if !strings.Contains(got, "\n"+tt.soul+"\n") { t.Errorf("ResolveSystemPrompt() soul = %q", got) } - gotApplyPatchShell := strings.Contains(got, "run the apply_patch command") + gotApplyPatchShell := strings.Contains(got, "shell tool with one complete apply_patch command") if gotApplyPatchShell != tt.useApplyPatchShell { t.Errorf( "ResolveSystemPrompt() apply_patch shell prompt = %v, want %v", @@ -47,6 +47,13 @@ func TestAgent_ResolveSystemPrompt(t *testing.T) { tt.useApplyPatchShell, ) } + if tt.useApplyPatchShell { + for _, phrase := range []string{"apply_patch <<'PATCH'", "@'", "'@ | apply_patch", "stdin field", "same file in parallel"} { + if !strings.Contains(got, phrase) { + t.Errorf("ResolveSystemPrompt() does not contain %q", phrase) + } + } + } if strings.Contains(got, "{{") || strings.Contains(got, "}}") { t.Errorf("ResolveSystemPrompt() contains unresolved template actions: %q", got) } diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister.go b/packages/agenty-core/pkg/infra/modelcatalog/lister.go index 63e5fe7..a5138fe 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister.go @@ -308,7 +308,11 @@ func (l *Lister) listGemini(ctx context.Context, provider catalog.Provider) ([]c } efforts := []shared.ReasoningEffort{} if item.Thinking != nil && *item.Thinking { - efforts = shared.StandardReasoningEfforts() + efforts = []shared.ReasoningEffort{ + shared.ReasoningLow, + shared.ReasoningMedium, + shared.ReasoningHigh, + } } model, err := normalizeModel( diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go index 9c6c78b..1ad134c 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go @@ -209,7 +209,11 @@ func TestListerGeminiPaginationAndThinking(t *testing.T) { if err != nil { t.Fatalf("List: %v", err) } - if len(models) != 2 || models[0].Code != "gemini-3-flash" || models[0].Name != "Gemini 3 Flash" || len(models[0].ReasoningEfforts) != len(shared.StandardReasoningEfforts()) { + if len(models) != 2 || models[0].Code != "gemini-3-flash" || models[0].Name != "Gemini 3 Flash" || !reflect.DeepEqual(models[0].ReasoningEfforts, []shared.ReasoningEffort{ + shared.ReasoningLow, + shared.ReasoningMedium, + shared.ReasoningHigh, + }) { t.Fatalf("models = %#v", models) } if models[1].Name != "gemini-3-pro" || len(models[1].ReasoningEfforts) != 0 { diff --git a/packages/agenty-core/pkg/infra/storage/catalog.go b/packages/agenty-core/pkg/infra/storage/catalog.go index f190b3e..ae6a395 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog.go +++ b/packages/agenty-core/pkg/infra/storage/catalog.go @@ -6,6 +6,7 @@ import ( "maps" "os" "path/filepath" + "reflect" "slices" "strings" "sync" @@ -24,11 +25,12 @@ type CatalogRepository struct { builtinProviders map[shared.Code]*catalog.Provider builtinOrder []shared.Code cacheMu sync.RWMutex + modelCache map[shared.Code]modelDiscoveryCache } type modelDiscoveryCache struct { - ExpiresAt time.Time `json:"expiresAt"` - Models []catalog.Model `json:"models"` + ExpiresAt time.Time + Models []catalog.Model } func NewCatalogRepository(providersDir string, builtinProviders ...*catalog.Provider) *CatalogRepository { @@ -43,7 +45,12 @@ func NewCatalogRepository(providersDir string, builtinProviders ...*catalog.Prov builtins[copy.Code] = copy order = append(order, copy.Code) } - return &CatalogRepository{providersDir: providersDir, builtinProviders: builtins, builtinOrder: order} + return &CatalogRepository{ + providersDir: providersDir, + builtinProviders: builtins, + builtinOrder: order, + modelCache: make(map[shared.Code]modelDiscoveryCache), + } } func (r *CatalogRepository) Get(_ context.Context, code shared.Code) (*catalog.Provider, error) { @@ -89,10 +96,10 @@ func (r *CatalogRepository) getLocked(code shared.Code, includeDiscoveryCache bo return &provider, nil } -// NeedsModelDiscovery reports whether a provider with no embedded/persisted -// models has no fresh discovery cache. Expired cache entries remain available -// through Get so callers can continue using the last known model metadata while -// a subsequent list operation refreshes it. +// NeedsModelDiscovery reports whether a provider has no fresh discovery cache +// or has no configured models. Expired cache entries remain available through +// Get so callers can continue using the last known model metadata while a +// subsequent list operation refreshes it. func (r *CatalogRepository) NeedsModelDiscovery(_ context.Context, code shared.Code) (bool, error) { r.cacheMu.RLock() defer r.cacheMu.RUnlock() @@ -101,15 +108,11 @@ func (r *CatalogRepository) NeedsModelDiscovery(_ context.Context, code shared.C if err != nil { return false, err } - if len(provider.Models) > 0 { - return false, nil - } - - cache, err := r.readModelDiscoveryCache(code) - if err != nil { - return true, nil + cache, ok := r.modelCache[code] + if ok { + return cache.ExpiresAt.IsZero() || !time.Now().UTC().Before(cache.ExpiresAt), nil } - return cache == nil || cache.ExpiresAt.IsZero() || !time.Now().UTC().Before(cache.ExpiresAt), nil + return len(provider.Models) == 0, nil } func (r *CatalogRepository) List(ctx context.Context) ([]*catalog.Provider, error) { @@ -161,20 +164,39 @@ func (r *CatalogRepository) Save(_ context.Context, provider *catalog.Provider) } r.cacheMu.Lock() defer r.cacheMu.Unlock() - if err := r.invalidateModelDiscoveryCacheLocked(provider.Code); err != nil { - return err - } if _, builtin := r.builtinProviders[provider.Code]; builtin { + delete(r.modelCache, provider.Code) return r.saveAPIKey(provider.Code, provider.APIKey) } if err := os.MkdirAll(r.providersDir, 0700); err != nil { return err } - // ModelsCached is a transient RPC hint, not part of provider config. - provider.ModelsCached = false - normalizeModels(provider) - providerData, err := json.MarshalIndent(provider, "", " ") + providerToPersist := cloneProvider(provider) + if provider.ModelsCached { + configured, err := r.getLocked(provider.Code, false) + if err != nil { + return err + } + cache, hasCache := r.modelCache[provider.Code] + if hasCache && providerConfigurationEqual(provider, configured) { + cache.Models = retainedCachedModels(cache.Models, provider.Models) + r.modelCache[provider.Code] = cache + providerToPersist.Models = mergeConfiguredAndEditedModels( + configured.Models, + provider.Models, + cache.Models, + ) + } else { + providerToPersist.Models = configured.Models + delete(r.modelCache, provider.Code) + } + } else { + delete(r.modelCache, provider.Code) + } + providerToPersist.ModelsCached = false + normalizeModels(providerToPersist) + providerData, err := json.MarshalIndent(providerToPersist, "", " ") if err != nil { return err } @@ -224,7 +246,10 @@ func (r *CatalogRepository) ReplaceModels( return err } - normalized := slices.Clone(models) + normalized := make([]catalog.Model, len(models)) + for index, model := range models { + normalized[index] = cloneModel(model) + } for index := range normalized { catalog.NormalizeReasoningCapabilities(&normalized[index]) if normalized[index].MaxOutputTokens <= 0 { @@ -235,75 +260,115 @@ func (r *CatalogRepository) ReplaceModels( normalized = make([]catalog.Model, 0) } - data, err := json.MarshalIndent(modelDiscoveryCache{ + r.modelCache[code] = modelDiscoveryCache{ ExpiresAt: expiresAt.UTC(), Models: normalized, - }, "", " ") - if err != nil { - return fmt.Errorf("storage: encode model discovery cache: %w", err) - } - cacheDir := filepath.Join(r.providersDir, ".models") - if err := os.MkdirAll(cacheDir, 0700); err != nil { - return fmt.Errorf("storage: create model discovery cache directory: %w", err) - } - cachePath := filepath.Join(cacheDir, code.String()+".json") - if err := os.WriteFile(cachePath, data, 0600); err != nil { - return fmt.Errorf("storage: write model discovery cache: %w", err) } return nil } func (r *CatalogRepository) applyModelDiscoveryCache(provider *catalog.Provider) { - if len(provider.Models) > 0 { + cache, ok := r.modelCache[provider.Code] + if !ok { return } - cache, err := r.readModelDiscoveryCache(provider.Code) - if err != nil { - return + configuredCodes := make(map[shared.ModelCode]struct{}, len(provider.Models)) + for _, model := range provider.Models { + configuredCodes[model.Code] = struct{}{} } - if cache == nil { - return + for _, model := range cache.Models { + if _, exists := configuredCodes[model.Code]; exists { + continue + } + provider.Models = append(provider.Models, cloneModel(model)) + provider.ModelsCached = true } - provider.Models = cache.Models - provider.ModelsCached = true } -func (r *CatalogRepository) readModelDiscoveryCache(code shared.Code) (*modelDiscoveryCache, error) { - data, err := os.ReadFile(r.modelDiscoveryCachePath(code)) - if err != nil { - return nil, err - } +func (r *CatalogRepository) invalidateModelDiscoveryCacheLocked(code shared.Code) error { + delete(r.modelCache, code) + return nil +} - var cache modelDiscoveryCache - if err := json.Unmarshal(data, &cache); err != nil { - return nil, err - } - normalizeModelsForCache(&cache) - return &cache, nil +func providerConfigurationEqual(left, right *catalog.Provider) bool { + return left.Code == right.Code && + left.Name == right.Name && + left.Type == right.Type && + left.BaseURL == right.BaseURL && + left.APIKey == right.APIKey && + left.Builtin == right.Builtin && + left.Official == right.Official && + left.FreeFormTool == right.FreeFormTool && + left.ModelsURL == right.ModelsURL && + left.TokenCountURL == right.TokenCountURL && + reflect.DeepEqual(left.Metadata, right.Metadata) } -func normalizeModelsForCache(cache *modelDiscoveryCache) { - if cache.Models == nil { - cache.Models = make([]catalog.Model, 0) +func mergeConfiguredAndEditedModels( + configured []catalog.Model, + incoming []catalog.Model, + cached []catalog.Model, +) []catalog.Model { + result := make([]catalog.Model, 0, len(configured)+len(incoming)) + indices := make(map[shared.ModelCode]int, len(configured)+len(incoming)) + incomingByCode := make(map[shared.ModelCode]catalog.Model, len(incoming)) + for _, model := range incoming { + incomingByCode[model.Code] = model + } + for _, model := range configured { + if _, exists := incomingByCode[model.Code]; !exists { + continue + } + indices[model.Code] = len(result) + result = append(result, cloneModel(model)) } - for index := range cache.Models { - catalog.NormalizeReasoningCapabilities(&cache.Models[index]) - if cache.Models[index].MaxOutputTokens <= 0 { - cache.Models[index].MaxOutputTokens = catalog.DefaultMaxOutputTokens + + cachedByCode := make(map[shared.ModelCode]catalog.Model, len(cached)) + for _, model := range cached { + cachedByCode[model.Code] = model + } + for _, model := range incoming { + cachedModel, wasCached := cachedByCode[model.Code] + if wasCached && modelsEqual(model, cachedModel) { + continue } + if index, exists := indices[model.Code]; exists { + result[index] = cloneModel(model) + continue + } + indices[model.Code] = len(result) + result = append(result, cloneModel(model)) } + return result } -func (r *CatalogRepository) modelDiscoveryCachePath(code shared.Code) string { - return filepath.Join(r.providersDir, ".models", code.String()+".json") +func retainedCachedModels(cached, incoming []catalog.Model) []catalog.Model { + incomingByCode := make(map[shared.ModelCode]catalog.Model, len(incoming)) + for _, model := range incoming { + incomingByCode[model.Code] = model + } + retained := make([]catalog.Model, 0, len(cached)) + for _, model := range cached { + incomingModel, exists := incomingByCode[model.Code] + if exists && modelsEqual(model, incomingModel) { + retained = append(retained, cloneModel(model)) + } + } + return retained } -func (r *CatalogRepository) invalidateModelDiscoveryCacheLocked(code shared.Code) error { - err := os.Remove(r.modelDiscoveryCachePath(code)) - if err != nil && !os.IsNotExist(err) { - return fmt.Errorf("storage: remove model discovery cache: %w", err) - } - return nil +func modelsEqual(left, right catalog.Model) bool { + return left.Code == right.Code && + left.Name == right.Name && + left.ContextWindow == right.ContextWindow && + left.MaxOutputTokens == right.MaxOutputTokens && + left.MultiModal == right.MultiModal && + left.Light == right.Light && + left.Reasoning == right.Reasoning && + slices.Equal(left.ReasoningEfforts, right.ReasoningEfforts) && + left.IsDefault == right.IsDefault && + left.CreatedAt.Equal(right.CreatedAt) && + left.UpdatedAt.Equal(right.UpdatedAt) } func normalizeModels(provider *catalog.Provider) { @@ -356,8 +421,14 @@ func cloneProvider(provider *catalog.Provider) *catalog.Provider { copy.Models = make([]catalog.Model, len(provider.Models)) copy.Models = append(copy.Models[:0], provider.Models...) for index := range copy.Models { + copy.Models[index] = cloneModel(provider.Models[index]) catalog.NormalizeReasoningCapabilities(©.Models[index]) } copy.Metadata = maps.Clone(provider.Metadata) return © } + +func cloneModel(model catalog.Model) catalog.Model { + model.ReasoningEfforts = slices.Clone(model.ReasoningEfforts) + return model +} diff --git a/packages/agenty-core/pkg/infra/storage/catalog_test.go b/packages/agenty-core/pkg/infra/storage/catalog_test.go index d0685f1..5c919b3 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog_test.go +++ b/packages/agenty-core/pkg/infra/storage/catalog_test.go @@ -58,7 +58,7 @@ func TestCatalogSaveAndGet(t *testing.T) { UpdatedAt: time.Now().UTC(), } provider.Models = []catalog.Model{model1, model2} - provider.ModelsCached = true + provider.ModelsCached = false if err := repo.Save(ctx, provider); err != nil { t.Fatalf("Save: %v", err) @@ -199,7 +199,7 @@ func TestCatalogBuiltinProviderPersistsOnlyAPIKey(t *testing.T) { } } -func TestCatalogModelDiscoveryCacheUsesExpirationAndSurvivesRestart(t *testing.T) { +func TestCatalogModelDiscoveryCacheUsesExpirationAndStaysInMemory(t *testing.T) { dir := t.TempDir() builtins, err := catalogdata.LoadProviders() if err != nil { @@ -232,16 +232,8 @@ func TestCatalogModelDiscoveryCacheUsesExpirationAndSurvivesRestart(t *testing.T } cachePath := filepath.Join(repo.providersDir, ".models", "openrouter.json") - cacheData, err := os.ReadFile(cachePath) - if err != nil { - t.Fatalf("read cache: %v", err) - } - var cache modelDiscoveryCache - if err := json.Unmarshal(cacheData, &cache); err != nil { - t.Fatalf("decode cache: %v", err) - } - if cache.ExpiresAt.Before(time.Now().UTC()) || len(cache.Models) != 1 { - t.Fatalf("cache = %+v", cache) + if _, err := os.Stat(cachePath); !os.IsNotExist(err) { + t.Fatalf("model discovery cache file exists: %v", err) } restarted := NewCatalogRepository(repo.providersDir, builtins...) @@ -249,11 +241,11 @@ func TestCatalogModelDiscoveryCacheUsesExpirationAndSurvivesRestart(t *testing.T if err != nil { t.Fatalf("Get after restart: %v", err) } - if len(restartedProvider.Models) != 1 { - t.Fatalf("restarted models = %#v", restartedProvider.Models) + if len(restartedProvider.Models) != 0 { + t.Fatalf("restarted models = %#v, want no in-memory cache", restartedProvider.Models) } - if !restartedProvider.ModelsCached { - t.Fatal("restarted provider did not expose the transient cache marker") + if restartedProvider.ModelsCached { + t.Fatal("restarted provider exposed a cache that should not survive restart") } if err := repo.ReplaceModels(ctx, code, models, time.Now().UTC().Add(-time.Minute)); err != nil { @@ -275,6 +267,58 @@ func TestCatalogModelDiscoveryCacheUsesExpirationAndSurvivesRestart(t *testing.T } } +func TestCatalogSaveDoesNotPersistDiscoveredModels(t *testing.T) { + repo := newCatalogRepo(t) + ctx := context.Background() + provider, err := catalog.NewProvider("custom", "Custom", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + if err := repo.Save(ctx, provider); err != nil { + t.Fatal(err) + } + + model := catalog.Model{ + Code: mustCatalogModelCode("gateway/model"), + Name: "Gateway model", + ContextWindow: 128_000, + MaxOutputTokens: 8_192, + ReasoningEfforts: []shared.ReasoningEffort{}, + } + if err := repo.ReplaceModels(ctx, provider.Code, []catalog.Model{model}, time.Now().UTC().Add(time.Hour)); err != nil { + t.Fatal(err) + } + cached, err := repo.Get(ctx, provider.Code) + if err != nil { + t.Fatal(err) + } + if !cached.ModelsCached || len(cached.Models) != 1 { + t.Fatalf("cached provider = %+v", cached) + } + + if err := repo.Save(ctx, cached); err != nil { + t.Fatal(err) + } + data, err := os.ReadFile(filepath.Join(repo.providersDir, "custom.json")) + if err != nil { + t.Fatal(err) + } + var persisted catalog.Provider + if err := json.Unmarshal(data, &persisted); err != nil { + t.Fatal(err) + } + if len(persisted.Models) != 0 { + t.Fatalf("persisted discovered models = %#v, want none", persisted.Models) + } + reloaded, err := repo.Get(ctx, provider.Code) + if err != nil { + t.Fatal(err) + } + if !reloaded.ModelsCached || len(reloaded.Models) != 1 { + t.Fatalf("cache after Save = %+v", reloaded) + } +} + func TestCatalogList(t *testing.T) { repo := newCatalogRepo(t) ctx := context.Background() diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index 7dc44a1..4964142 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -1,4 +1,5 @@ use std::collections::{HashMap, HashSet}; +use std::ffi::OsString; use std::fmt::{self, Display, Formatter}; use std::fs::{self, File, OpenOptions}; use std::io::{self, Write}; @@ -142,7 +143,12 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { } let results = transaction.prepare_results()?; - transaction.commit()?; + let lock_paths = transaction.lock_paths()?; + let locks = FileLocks::acquire(&lock_paths, std::process::id())?; + let commit_result = transaction.commit(); + let release_result = locks.release(); + commit_result?; + release_result?; Ok(PatchResult { success: true, @@ -151,6 +157,145 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { }) } +struct FileLocks { + paths: Vec, + created_dirs: Vec, +} + +impl FileLocks { + fn acquire(targets: &[PathBuf], pid: u32) -> Result { + let mut targets = targets.to_vec(); + targets.sort(); + targets.dedup(); + + let mut locks = Self { + paths: Vec::with_capacity(targets.len()), + created_dirs: Vec::new(), + }; + for target in targets { + let parent = target.parent().unwrap_or_else(|| Path::new(".")); + match create_missing_dirs(parent) { + Ok(created) => locks.created_dirs.extend(created), + Err(error) => { + locks.cleanup(); + return Err(error); + } + } + + let lock_path = match lock_path(&target) { + Ok(path) => path, + Err(error) => { + locks.cleanup(); + return Err(error); + } + }; + let result = (|| -> Result<(), PatchError> { + let mut lock = OpenOptions::new() + .write(true) + .create_new(true) + .open(&lock_path) + .map_err(|error| { + if error.kind() == io::ErrorKind::AlreadyExists { + let owner = read_lock_owner(&lock_path); + return PatchError::Conflict(format!( + "file {} is locked by apply_patch pid {owner} ({})", + target.display(), + lock_path.display(), + )); + } + PatchError::Io(error) + })?; + writeln!(lock, "{pid}")?; + lock.sync_all()?; + Ok(()) + })(); + if let Err(error) = result { + locks.cleanup(); + return Err(error); + } + locks.paths.push(lock_path); + } + + Ok(locks) + } + + fn release(self) -> Result<(), PatchError> { + let mut first_error = None; + for path in self.paths.iter().rev() { + if let Err(error) = fs::remove_file(path) { + if error.kind() != io::ErrorKind::NotFound { + first_error.get_or_insert(PatchError::Io(error)); + } + } + } + if let Err(error) = sync_lock_parent_directories(&self.paths) { + first_error.get_or_insert(error); + } + for directory in self.created_dirs.iter().rev() { + if let Err(error) = fs::remove_dir(directory) { + if error.kind() != io::ErrorKind::NotFound + && error.kind() != io::ErrorKind::DirectoryNotEmpty + { + first_error.get_or_insert(PatchError::Io(error)); + } + } + } + first_error.map_or(Ok(()), Err) + } + + fn cleanup(&mut self) { + for path in self.paths.drain(..).rev() { + let _ = fs::remove_file(path); + } + for directory in self.created_dirs.drain(..).rev() { + let _ = fs::remove_dir(directory); + } + } +} + +fn lock_path(target: &Path) -> Result { + let file_name = target.file_name().ok_or_else(|| { + PatchError::Invalid(format!( + "cannot create a file lock for {}", + target.display() + )) + })?; + let mut lock_name = OsString::from("."); + lock_name.push(file_name); + lock_name.push(".lock"); + Ok(target + .parent() + .unwrap_or_else(|| Path::new(".")) + .join(lock_name)) +} + +fn read_lock_owner(path: &Path) -> String { + fs::read_to_string(path) + .map(|contents| contents.trim().to_string()) + .ok() + .filter(|contents| !contents.is_empty()) + .unwrap_or_else(|| "unknown".to_string()) +} + +#[cfg(unix)] +fn sync_lock_parent_directories(paths: &[PathBuf]) -> Result<(), PatchError> { + let mut parents = HashSet::new(); + for path in paths { + if let Some(parent) = path.parent() { + parents.insert(parent); + } + } + for parent in parents { + File::open(parent)?.sync_all()?; + } + Ok(()) +} + +#[cfg(not(unix))] +fn sync_lock_parent_directories(_paths: &[PathBuf]) -> Result<(), PatchError> { + Ok(()) +} + fn parse_envelope(cwd: &Path, patch: &str) -> Result, PatchError> { let lines = normalized_lines(patch); if lines.len() < 3 || lines.first().map(String::as_str) != Some(BEGIN_MARKER) { @@ -515,6 +660,14 @@ impl Transaction { Ok(results) } + fn lock_paths(&self) -> Result, PatchError> { + Ok(self + .changes()? + .into_iter() + .map(|change| change.path) + .collect()) + } + fn commit(&self) -> Result<(), PatchError> { let changes = self.changes()?; if changes.is_empty() { @@ -1241,6 +1394,45 @@ mod tests { ); } + #[test] + fn creates_pid_lock_for_each_transaction_target_and_removes_it() { + let cwd = temp_dir("file-lock"); + let target = cwd.join("notes.txt"); + let locks = FileLocks::acquire(std::slice::from_ref(&target), 4242).unwrap(); + let lock = cwd.join(".notes.txt.lock"); + assert_eq!(fs::read_to_string(&lock).unwrap(), "4242\n"); + locks.release().unwrap(); + assert!(!lock.exists()); + } + + #[test] + fn rejects_a_target_with_an_existing_pid_lock_without_writing() { + let cwd = temp_dir("locked-target"); + let target = cwd.join("notes.txt"); + fs::write(&target, "one\n").unwrap(); + fs::write(cwd.join(".notes.txt.lock"), "9876\n").unwrap(); + + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err().to_string(); + assert!(error.contains("locked by apply_patch pid 9876")); + assert_eq!(fs::read_to_string(target).unwrap(), "one\n"); + } + + #[test] + fn move_requires_locks_for_source_and_destination() { + let cwd = temp_dir("locked-move"); + let source = cwd.join("old.txt"); + let destination = cwd.join("new.txt"); + fs::write(&source, "one\n").unwrap(); + fs::write(cwd.join(".new.txt.lock"), "1111\n").unwrap(); + + let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err().to_string(); + assert!(error.contains("new.txt")); + assert_eq!(fs::read_to_string(source).unwrap(), "one\n"); + assert!(!destination.exists()); + } + #[cfg(unix)] #[test] fn restores_backups_when_staging_fails() { From 66a01ede5ddbc3abdf7d25a8887b9526c99d2c8c Mon Sep 17 00:00:00 2001 From: masteryyh Date: Thu, 27 Aug 2026 15:38:02 +0800 Subject: [PATCH 09/12] fix: compaction, file locking and env expanding Signed-off-by: masteryyh --- README.md | 4 +- README.zh-CN.md | 4 +- package.json | 4 +- packages/agenty-cli/src/api/types.ts | 1 + packages/agenty-cli/src/components/List.tsx | 7 +- .../src/components/ProviderOverlay.test.ts | 35 +++- .../src/components/ProviderOverlay.tsx | 24 ++- .../src/components/SelectOverlay.test.tsx | 129 +++++++++++++ .../src/components/SelectOverlay.tsx | 1 + .../src/components/wizardSetup.test.ts | 54 +++++- .../agenty-cli/src/components/wizardSetup.ts | 17 +- .../agenty-cli/src/consts/providerPresets.ts | 3 +- packages/agenty-cli/src/localCore.test.ts | 22 ++- packages/agenty-cli/src/localCore.ts | 20 +- packages/agenty-core/README-CN.md | 11 +- packages/agenty-core/README.md | 13 +- packages/agenty-core/cmd/main.go | 12 +- packages/agenty-core/package.json | 4 +- .../agenty-core/pkg/agentloop/compaction.go | 2 +- .../agenty-core/pkg/agentloop/engine_test.go | 4 +- .../agenty-core/pkg/domain/catalog/model.go | 1 + .../pkg/domain/catalog/provider.go | 7 + .../pkg/domain/catalog/provider_test.go | 28 +++ .../agenty-core/pkg/infra/config/config.go | 2 + .../pkg/infra/config/config_test.go | 5 +- .../agenty-core/pkg/infra/config/types.go | 3 + .../agenty-core/pkg/infra/storage/catalog.go | 12 +- .../pkg/infra/storage/catalog_test.go | 66 +++++++ packages/agenty-core/scripts/build.mjs | 107 +++++++++++ packages/agenty-core/scripts/build.test.mjs | 37 ++++ packages/patch-applier/Cargo.lock | 89 +++++++++ packages/patch-applier/Cargo.toml | 1 + packages/patch-applier/src/lib.rs | 178 ++++++++++++------ packages/patch-applier/tests/cli.rs | 4 + scripts/dev.mjs | 77 ++++++++ scripts/dev.test.mjs | 21 +++ scripts/platform.mjs | 44 +++++ 37 files changed, 944 insertions(+), 109 deletions(-) create mode 100644 packages/agenty-cli/src/components/SelectOverlay.test.tsx create mode 100644 packages/agenty-core/scripts/build.mjs create mode 100644 packages/agenty-core/scripts/build.test.mjs create mode 100644 scripts/dev.mjs create mode 100644 scripts/dev.test.mjs create mode 100644 scripts/platform.mjs diff --git a/README.md b/README.md index d9b3626..ace7cb9 100644 --- a/README.md +++ b/README.md @@ -63,8 +63,10 @@ Core stores data under `~/.agenty` by default. Pass `--data-dir ` to the C | Configuration | `~/.agenty/config.json` | | Session transcripts | `~/.agenty/sessions///
/.jsonl` | | Session index | `~/.agenty/agenty.sqlite` | -| Providers and models | Built-in catalog is embedded in the core binary; custom providers use `~/.agenty/providers/.json`, built-in provider files store only API keys; core automatically discovers empty configured catalogs when listing providers/models and caches them under `~/.agenty/providers/.models/` for 8 hours | +| Providers and models | Built-in catalog is embedded in the core binary; custom providers use `~/.agenty/providers/.json`, while built-in provider files store only API keys | +| Model discovery cache | Kept only in the running core process for 8 hours; it is refreshed on demand and does not survive a core restart | | Agents | `~/.agenty/agents/` | +| Patch transaction locks | `~/.agenty/locks/` | | Logs | `~/.agenty/logs///
/core.log` | `AGENTY_LOG_LEVEL` accepts `debug`, `info`, `warn`, or `error`. diff --git a/README.zh-CN.md b/README.zh-CN.md index 655a645..e086970 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -57,8 +57,10 @@ core 默认把数据保存在 `~/.agenty`。可向 CLI 传入 `--data-dir | 配置 | `~/.agenty/config.json` | | 会话 transcript | `~/.agenty/sessions///
/.jsonl` | | 会话索引 | `~/.agenty/agenty.sqlite` | -| Providers 和 models | 内置 catalog 固化在 core 二进制中;自定义 provider 使用 `~/.agenty/providers/.json`,内置 provider 文件仅保存 API key;core 在获取 provider/model 列表时自动发现已配置的空模型 catalog,结果缓存于 `~/.agenty/providers/.models/`,有效期 8 小时 | +| Providers 和 models | 内置 catalog 固化在 core 二进制中;自定义 provider 使用 `~/.agenty/providers/.json`,内置 provider 文件仅保存 API key | +| 模型发现缓存 | 仅保存在运行中的 core 进程内,有效期 8 小时;按需刷新,core 重启后不会保留 | | Agents | `~/.agenty/agents/` | +| Patch 事务锁 | `~/.agenty/locks/` | | 日志 | `~/.agenty/logs///
/core.log` | `AGENTY_LOG_LEVEL` 接受 `debug`、`info`、`warn` 或 `error`; diff --git a/package.json b/package.json index c577efa..42d9e62 100644 --- a/package.json +++ b/package.json @@ -4,11 +4,11 @@ "private": true, "scripts": { "build": "turbo run build && mkdir -p dist && (cp packages/agenty-bootstrap/bin/* dist/ 2>/dev/null || true)", - "test": "turbo run test", + "test": "node --test scripts/dev.test.mjs && turbo run test", "lint": "eslint .", "lint:fix": "eslint . --fix", "clean": "rm -rf dist && turbo run clean", - "dev": "turbo run build --filter=agenty-bootstrap && pnpm --filter agenty-bootstrap exec bun -e \"const { join } = await import('node:path'); const { resolveArch, resolveOS } = await import('./scripts/target.ts'); const os = resolveOS(); const executable = join('bin', 'agenty-' + os + '-' + resolveArch() + (os === 'windows' ? '.exe' : '')); const child = Bun.spawn([executable, ...process.argv.slice(1)], { stdin: 'inherit', stdout: 'inherit', stderr: 'inherit', env: process.env }); process.exit(await child.exited);\"", + "dev": "node scripts/dev.mjs", "deepclean": "pnpm clean && rm -rf node_modules packages/*/node_modules", "core:build": "turbo run build --filter=agenty-core", "core:test": "turbo run test --filter=agenty-core", diff --git a/packages/agenty-cli/src/api/types.ts b/packages/agenty-cli/src/api/types.ts index ec1cc65..3628f6e 100644 --- a/packages/agenty-cli/src/api/types.ts +++ b/packages/agenty-cli/src/api/types.ts @@ -54,6 +54,7 @@ export interface ModelDto { reasoning?: boolean; reasoningEfforts?: ReasoningEffort[]; isDefault: boolean; + cached?: boolean; createdAt?: string; updatedAt?: string; } diff --git a/packages/agenty-cli/src/components/List.tsx b/packages/agenty-cli/src/components/List.tsx index 2942852..9315bdd 100644 --- a/packages/agenty-cli/src/components/List.tsx +++ b/packages/agenty-cli/src/components/List.tsx @@ -105,6 +105,7 @@ export interface ListNavigationOptions { item: T | undefined, ) => void; active?: boolean; + closeWhenInactive?: boolean; } export function useListNavigation({ @@ -115,12 +116,16 @@ export function useListNavigation({ onClose, onInput, active = true, + closeWhenInactive = false, }: ListNavigationOptions) { useInput((input, key, event) => { if (key.escape && onClose) { onClose(); return; } + if (!active) { + return; + } if (key.upArrow) { event.preventDefault(); onCursor(Math.max(cursor - 1, 0)); @@ -138,7 +143,7 @@ export function useListNavigation({ return; } onInput?.(input, key, event, item); - }, { isActive: active }); + }, { isActive: active || (closeWhenInactive && onClose !== undefined) }); } export interface KeyValueRow { diff --git a/packages/agenty-cli/src/components/ProviderOverlay.test.ts b/packages/agenty-cli/src/components/ProviderOverlay.test.ts index 59bf0cf..5245832 100644 --- a/packages/agenty-cli/src/components/ProviderOverlay.test.ts +++ b/packages/agenty-cli/src/components/ProviderOverlay.test.ts @@ -1,9 +1,10 @@ import { describe, expect, test } from "bun:test"; -import type { ModelProviderDto } from "../api/types"; +import type { CoreModelDto, ModelProviderDto } from "../api/types"; import { buildBuiltinProviderUpdate, buildCreateModelFields, + buildModelUpdate, buildProviderFields, parseModelValues, } from "./ProviderOverlay"; @@ -84,4 +85,36 @@ describe("provider overlay model advanced options", () => { expect(buildCreateModelFields(false).find((field) => field.key === "light")?.label) .toBe("Light model"); }); + + test("preserves a model default flag while editing unrelated fields", () => { + const target: CoreModelDto = { + code: "model", + name: "Model", + contextWindow: 128_000, + maxOutputTokens: 8_192, + multiModal: false, + light: false, + isDefault: true, + }; + const parsed = parseModelValues({ + code: "model", + name: "Renamed model", + contextWindow: "128000", + maxOutputTokens: "16384", + multiModal: "false", + light: "false", + reasoning: "true", + reasoningEfforts: "[\"low\", \"high\"]", + }); + + expect(typeof parsed).not.toBe("string"); + if (typeof parsed === "string") { + throw new Error(parsed); + } + expect(buildModelUpdate(target, parsed)).toMatchObject({ + isDefault: true, + maxOutputTokens: 16_384, + name: "Renamed model", + }); + }); }); diff --git a/packages/agenty-cli/src/components/ProviderOverlay.tsx b/packages/agenty-cli/src/components/ProviderOverlay.tsx index 69c0110..d8a8e13 100644 --- a/packages/agenty-cli/src/components/ProviderOverlay.tsx +++ b/packages/agenty-cli/src/components/ProviderOverlay.tsx @@ -6,6 +6,7 @@ import type { CreateModelDto, ModelProviderDto, ReasoningEffort, + UpdateModelDto, UpdateModelProviderDto, } from "../api/types"; import { STANDARD_REASONING_EFFORTS } from "../api/types"; @@ -212,6 +213,19 @@ export function parseModelValues(values: Record): CreateModelDto }; } +export function buildModelUpdate(target: CoreModelDto, input: CreateModelDto): UpdateModelDto { + return { + name: input.name, + contextWindow: input.contextWindow, + maxOutputTokens: input.maxOutputTokens, + multiModal: input.multiModal, + light: input.light, + reasoning: input.reasoning, + reasoningEfforts: input.reasoningEfforts, + isDefault: target.isDefault, + }; +} + type Mode = | { kind: "list" } | { kind: "create-provider" } @@ -452,15 +466,7 @@ export function ProviderOverlay() { return; } try { - await client.updateModel(provider.code, target.code, { - name: parsed.name, - contextWindow: parsed.contextWindow, - maxOutputTokens: parsed.maxOutputTokens, - multiModal: parsed.multiModal, - light: parsed.light, - reasoning: parsed.reasoning, - reasoningEfforts: parsed.reasoningEfforts, - }); + await client.updateModel(provider.code, target.code, buildModelUpdate(target, parsed)); setToast(`Model updated: ${parsed.name}`); await reload(); returnToList(); diff --git a/packages/agenty-cli/src/components/SelectOverlay.test.tsx b/packages/agenty-cli/src/components/SelectOverlay.test.tsx new file mode 100644 index 0000000..7e2bbbb --- /dev/null +++ b/packages/agenty-cli/src/components/SelectOverlay.test.tsx @@ -0,0 +1,129 @@ +import { testRender } from "@opentui/react/test-utils"; +import { describe, expect, test } from "bun:test"; +import { act } from "react"; + +import { SelectOverlay } from "./SelectOverlay"; + +async function settleEffects(): Promise { + await new Promise((resolve) => { + setTimeout(resolve, 50); + }); +} + +describe("SelectOverlay", () => { + test("closes on Escape while entries are unavailable", async () => { + let closeCount = 0; + let resolveLoad: ((entries: Array<{ label: string; data: string }>) => void) | undefined; + const load = () => new Promise>((resolve) => { + resolveLoad = resolve; + }); + const setup = await testRender( + undefined} + onClose={() => { + closeCount++; + }} + />, + { width: 48, height: 8 }, + ); + + try { + await act(async () => { + await settleEffects(); + await setup.flush(); + }); + await act(async () => { + setup.mockInput.pressEscape(); + await settleEffects(); + await setup.flush(); + }); + + expect(closeCount).toBe(1); + } finally { + resolveLoad?.([]); + act(() => setup.renderer.destroy()); + } + }); + + test("closes on Escape after loading an empty result", async () => { + let closeCount = 0; + let resolveLoad: ((entries: Array<{ label: string; data: string }>) => void) | undefined; + const load = () => new Promise>((resolve) => { + resolveLoad = resolve; + }); + const setup = await testRender( + undefined} + onClose={() => { + closeCount++; + }} + />, + { width: 48, height: 8 }, + ); + + try { + await act(async () => { + await settleEffects(); + }); + await act(async () => { + resolveLoad?.([]); + await settleEffects(); + await setup.flush(); + }); + expect(setup.captureCharFrame()).toContain("No items"); + await act(async () => { + setup.mockInput.pressEscape(); + await settleEffects(); + await setup.flush(); + }); + + expect(closeCount).toBe(1); + } finally { + act(() => setup.renderer.destroy()); + } + }); + + test("closes on Escape after a load failure", async () => { + let closeCount = 0; + let rejectLoad: ((reason?: unknown) => void) | undefined; + const load = () => new Promise>((_resolve, reject) => { + rejectLoad = reject; + }); + const setup = await testRender( + undefined} + onClose={() => { + closeCount++; + }} + />, + { width: 48, height: 8 }, + ); + + try { + await act(async () => { + await settleEffects(); + }); + await act(async () => { + rejectLoad?.(new Error("load failed")); + await settleEffects(); + await setup.flush(); + }); + expect(setup.captureCharFrame()).toContain("Failed: load failed"); + await act(async () => { + setup.mockInput.pressEscape(); + await settleEffects(); + await setup.flush(); + }); + + expect(closeCount).toBe(1); + } finally { + act(() => setup.renderer.destroy()); + } + }); +}); diff --git a/packages/agenty-cli/src/components/SelectOverlay.tsx b/packages/agenty-cli/src/components/SelectOverlay.tsx index 7f182b1..5fddc2f 100644 --- a/packages/agenty-cli/src/components/SelectOverlay.tsx +++ b/packages/agenty-cli/src/components/SelectOverlay.tsx @@ -60,6 +60,7 @@ export function SelectOverlay({ onActivate: (entry) => onSelect(entry.data), onClose, active: entries !== null && entries.length > 0, + closeWhenInactive: true, }); return ( diff --git a/packages/agenty-cli/src/components/wizardSetup.test.ts b/packages/agenty-cli/src/components/wizardSetup.test.ts index bd6b21a..8f54b4f 100644 --- a/packages/agenty-cli/src/components/wizardSetup.test.ts +++ b/packages/agenty-cli/src/components/wizardSetup.test.ts @@ -291,7 +291,11 @@ describe("first-run provider setup", () => { test("does not persist an untouched discovered model cache", async () => { const draft = createDraft(); - const provider = { ...createProvider(draft), modelsCached: true }; + const provider = { + ...createProvider(draft), + models: [{ ...createProvider(draft).models[0], cached: true }], + modelsCached: true, + }; const models = modelDraftsForProvider(draft, provider); const client = fakeClient([provider], [createAgent("default", true)]); @@ -307,4 +311,52 @@ describe("first-run provider setup", () => { "initialize.complete", ]); }); + + test("removes configured models while preserving discovered cache entries", async () => { + const draft = createDraft(); + const provider = { + ...createProvider(draft), + models: [ + { ...createProvider(draft).models[0], code: "configured", name: "Configured", isDefault: false }, + { ...createProvider(draft).models[0], code: "cached", name: "Cached", cached: true, isDefault: true }, + ], + modelsCached: true, + }; + const models = modelDraftsForProvider(draft, provider); + const cached = models.find((model) => model.code === "cached"); + const client = fakeClient([provider], [createAgent("default", true)]); + + expect(models.map((model) => ({ code: model.code, source: model.source }))).toEqual([ + { code: "configured", source: "configured" }, + { code: "cached", source: "cached" }, + ]); + expect(cached).toBeDefined(); + await persistWizardSetup(client, [draft], [cached!], selectedModelId(cached!)); + + expect(client.deletedModels).toEqual([{ providerCode: "custom", modelCode: "configured" }]); + expect(client.createdModels).toEqual([]); + }); + + test("preserves a default model for each provider", async () => { + const firstDraft = createDraft(); + const secondDraft = { + ...createDraft(), + id: "custom:1", + code: "second", + name: "Second", + }; + const first = { ...createModel(firstDraft), id: "custom:0:first", code: "first", isDefault: true }; + const second = { ...createModel(secondDraft), id: "custom:1:second", code: "second", isDefault: true }; + const client = fakeClient( + [createProvider(firstDraft, first), createProvider(secondDraft, second)], + [createAgent("default", true)], + ); + + await persistWizardSetup(client, [firstDraft, secondDraft], [first, second], selectedModelId(first)); + + expect(client.createdModels.map(({ modelCode, isDefault }) => ({ modelCode, isDefault }))).toEqual([ + { modelCode: "first", isDefault: true }, + { modelCode: "second", isDefault: true }, + ]); + }); }); diff --git a/packages/agenty-cli/src/components/wizardSetup.ts b/packages/agenty-cli/src/components/wizardSetup.ts index fb781fb..200c1a5 100644 --- a/packages/agenty-cli/src/components/wizardSetup.ts +++ b/packages/agenty-cli/src/components/wizardSetup.ts @@ -168,15 +168,16 @@ export async function persistWizardSetup( if (!draft.builtin) { const providerModels = modelsByProvider.get(draft.id) ?? []; - if (existing?.modelsCached !== true) { - const desiredModelCodes = new Set(providerModels.map((model) => model.code.trim())); - for (const model of existing?.models ?? []) { - if (!desiredModelCodes.has(model.code)) { - await client.deleteModel(providerCode, model.code); - } + const desiredModelCodes = new Set(providerModels.map((model) => model.code.trim())); + for (const model of existing?.models ?? []) { + if (model.cached !== true && !desiredModelCodes.has(model.code)) { + await client.deleteModel(providerCode, model.code); } } + const selectedProviderModel = providerModels.find((model) => selectedModelId(model) === selectedId); + const persistSelectedProviderDefault = selectedProviderModel !== undefined && selectedProviderModel.source !== "cached"; + for (const model of providerModels) { if (model.source === "cached") { continue; @@ -191,7 +192,9 @@ export async function persistWizardSetup( light: model.light, reasoning: model.reasoning !== false && (model.reasoning === true || model.reasoningEfforts.length > 0), reasoningEfforts: model.reasoningEfforts, - isDefault: selectedModelId(model) === selectedId, + isDefault: persistSelectedProviderDefault + ? selectedModelId(model) === selectedId + : model.isDefault, }); } diff --git a/packages/agenty-cli/src/consts/providerPresets.ts b/packages/agenty-cli/src/consts/providerPresets.ts index b009917..3374389 100644 --- a/packages/agenty-cli/src/consts/providerPresets.ts +++ b/packages/agenty-cli/src/consts/providerPresets.ts @@ -115,13 +115,12 @@ export function modelDraftsForProvider( provider: ProviderDraft, existing?: ModelProviderDto, ): ModelDraft[] { - const source: ModelDraft["source"] = existing?.modelsCached === true ? "cached" : "configured"; return (existing?.models ?? []).map((model) => createModelDraft( provider, `${provider.id}:model:${model.code}`, model, - source, + model.cached === true ? "cached" : "configured", ), ); } diff --git a/packages/agenty-cli/src/localCore.test.ts b/packages/agenty-cli/src/localCore.test.ts index a243f68..aab009f 100644 --- a/packages/agenty-cli/src/localCore.test.ts +++ b/packages/agenty-cli/src/localCore.test.ts @@ -2,7 +2,7 @@ import { delimiter } from "node:path"; import { describe, expect, test } from "bun:test"; -import { pickCorePath, prependCoreDirectoryToPath } from "./localCore"; +import { MANAGED_BIN_DIR, pickCorePath, prependCoreDirectoryToPath } from "./localCore"; const candidates = { repoBin: "/repo/packages/agenty-core/bin/agenty-core", @@ -30,7 +30,25 @@ describe("pickCorePath", () => { describe("prependCoreDirectoryToPath", () => { test("places the core directory before the inherited PATH", () => { expect(prependCoreDirectoryToPath("/managed/bin/core", "/usr/bin")).toBe( - ["/managed/bin", "/usr/bin"].join(delimiter), + ["/managed/bin", MANAGED_BIN_DIR, "/usr/bin"].join(delimiter), ); }); + + test("keeps the managed helper directory reachable for an override", () => { + expect(prependCoreDirectoryToPath( + "/custom/bin/core", + "/usr/bin:/managed/bin", + "/managed/bin", + ":", + )).toBe("/custom/bin:/managed/bin:/usr/bin"); + }); + + test("uses the Windows path delimiter when assembling an override PATH", () => { + expect(prependCoreDirectoryToPath( + "C:/custom/core.exe", + "C:/Windows/System32;C:/managed/bin", + "C:/managed/bin", + ";", + )).toBe("C:/custom;C:/managed/bin;C:/Windows/System32"); + }); }); diff --git a/packages/agenty-cli/src/localCore.ts b/packages/agenty-cli/src/localCore.ts index 0dfbbf8..4a04489 100644 --- a/packages/agenty-cli/src/localCore.ts +++ b/packages/agenty-cli/src/localCore.ts @@ -4,10 +4,10 @@ import { delimiter, dirname, join, resolve } from "node:path"; import { StdioRPCClient } from "./core/rpc"; +export const MANAGED_BIN_DIR = join(homedir(), ".agenty", "bin"); + export const MANAGED_CORE_PATH = join( - homedir(), - ".agenty", - "bin", + MANAGED_BIN_DIR, process.platform === "win32" ? "core.exe" : "core", ); @@ -38,8 +38,18 @@ export function pickCorePath( return null; } -export function prependCoreDirectoryToPath(binary: string, currentPath?: string): string { - return [dirname(binary), currentPath].filter(Boolean).join(delimiter); +export function prependCoreDirectoryToPath( + binary: string, + currentPath?: string, + managedBinDirectory = MANAGED_BIN_DIR, + pathDelimiter = delimiter, +): string { + const entries = [ + dirname(binary), + managedBinDirectory, + ...(currentPath?.split(pathDelimiter) ?? []), + ].filter(Boolean); + return Array.from(new Set(entries)).join(pathDelimiter); } export interface LocalCore { diff --git a/packages/agenty-core/README-CN.md b/packages/agenty-core/README-CN.md index decd515..cc64d9b 100644 --- a/packages/agenty-core/README-CN.md +++ b/packages/agenty-core/README-CN.md @@ -177,11 +177,12 @@ Methods 使用 `resource.action` 命名: `provider.list` 可选接收 `{providerCode}`。不传时,core 会并行获取所有已配置且 catalog 为空的 provider;传入时只会获取指定 provider。`provider.listModels` 接收同样的 -`{providerCode}`,供直接调用者使用 core 内置的发现流程。成功结果会缓存到 -`~/.agenty/providers/.models/.json`,有效期为 8 小时,JSON 中保存 `expiresAt` -和标准化模型列表。过期缓存仍作为旧数据返回,下一次 list 时再刷新。它会兼容常见的 `id`、 -名称和 token 限制字段,自动跟随 provider 分页;上下文窗口或最大输出 token 缺失或不为正数时 -分别使用 `256000` 和 `65536`,缺少 reasoning 能力时返回空的 `reasoningEfforts` 数组。 +`{providerCode}`,供直接调用者使用 core 内置的发现流程。成功结果仅缓存于运行中的 core 进程, +有效期为 8 小时,不写入磁盘,core 重启后不会保留。过期缓存仍作为旧数据返回,下一次 list 时 +再刷新。它会兼容常见的 `id`、名称和 token 限制字段,自动跟随 provider 分页;上下文窗口或最大 +输出 token 缺失或不为正数时分别使用 `256000` 和 `65536`,缺少 reasoning 能力时返回空的 +`reasoningEfforts` 数组。事务性 `apply_patch` 锁位于 `~/.agenty/locks/`,每个锁都会记录 helper +进程 PID 和完整目标路径。 `session.start` 接收 `{id, content}`,持久化 running round 后立即返回 round 标识和 `running` 状态,完整 agent turn 由引擎异步继续执行。执行期间,core 会写出 diff --git a/packages/agenty-core/README.md b/packages/agenty-core/README.md index bea45b2..737900a 100644 --- a/packages/agenty-core/README.md +++ b/packages/agenty-core/README.md @@ -199,12 +199,13 @@ Methods follow a `resource.action` naming: `provider.list` accepts an optional `{providerCode}`. Without it, core discovers all configured providers whose catalog is empty in parallel; with it, only that provider is eligible for discovery. `provider.listModels` accepts `{providerCode}` and exposes the same -core-owned discovery path for direct callers. Successful discovery is cached under -`~/.agenty/providers/.models/.json` for 8 hours; the JSON stores an `expiresAt` -timestamp and the normalized models. Expired entries remain available as stale data while a -subsequent list refreshes them. It maps common `id`/name/token-limit fields, follows provider -pagination, defaults missing or non-positive context/output limits to `256000` and `65536`, -and represents missing reasoning capability as an empty `reasoningEfforts` array. +core-owned discovery path for direct callers. Successful discovery is cached in the running +core process for 8 hours. It is not written to disk and does not survive a core +restart. Expired entries remain available as stale data while a subsequent list refreshes them. +It maps common `id`/name/token-limit fields, follows provider pagination, defaults missing or +non-positive context/output limits to `256000` and `65536`, and represents missing reasoning +capability as an empty `reasoningEfforts` array. Transactional `apply_patch` locks live under +`~/.agenty/locks/`; each lock records its helper PID and full target path. `session.start` accepts `{id, content}` and returns the persisted round's identifiers and `running` status immediately; the engine continues the full agent turn diff --git a/packages/agenty-core/cmd/main.go b/packages/agenty-core/cmd/main.go index f8b1fc8..b14190a 100644 --- a/packages/agenty-core/cmd/main.go +++ b/packages/agenty-core/cmd/main.go @@ -6,6 +6,7 @@ import ( "fmt" "log/slog" "os" + "path/filepath" "time" "github.com/masteryyh/agenty-core/pkg/agentloop" @@ -30,6 +31,15 @@ func run() (exitCode int) { fmt.Fprintln(os.Stderr, "agenty-core: failed to initialize config:", err) return 1 } + dataDir, err := filepath.Abs(config.Get().Paths().DataDir) + if err != nil { + fmt.Fprintln(os.Stderr, "agenty-core: failed to resolve data directory:", err) + return 1 + } + if err := os.Setenv(config.EnvDataDir, dataDir); err != nil { + fmt.Fprintln(os.Stderr, "agenty-core: failed to set data directory environment:", err) + return 1 + } logger, err := logging.Open() if err != nil { @@ -59,7 +69,7 @@ func run() (exitCode int) { } }() - slog.InfoContext(ctx, "agenty-core started", "dataDir", config.Get().Paths().DataDir) + slog.InfoContext(ctx, "agenty-core started", "dataDir", dataDir) toolRegistry := agentloop.NewRegistry() if err := builtin.RegisterAll(toolRegistry); err != nil { slog.ErrorContext(ctx, "failed to register built-in tools", "error", err) diff --git a/packages/agenty-core/package.json b/packages/agenty-core/package.json index be4fbc4..978c3ab 100644 --- a/packages/agenty-core/package.json +++ b/packages/agenty-core/package.json @@ -3,8 +3,8 @@ "version": "0.1.0", "private": true, "scripts": { - "build": "mkdir -p \"${PACKAGE_DIR:-bin}\" && go build -o \"${PACKAGE_DIR:-bin}/${BIN_NAME:-agenty-core}\" ./cmd && if [ \"${GOOS:-}\" = \"windows\" ]; then cp ../patch-applier/target/release/apply_patch.exe \"${PACKAGE_DIR:-bin}/apply_patch.exe\"; else cp ../patch-applier/target/release/apply_patch \"${PACKAGE_DIR:-bin}/apply_patch\"; fi", - "test": "go test ./...", + "build": "node scripts/build.mjs", + "test": "node --test scripts/build.test.mjs && go test ./...", "test:integration": "go test -tags=integration ./...", "test:e2e": "go test -tags=e2e -count=1 -parallel=8 ./test/e2e", "test:e2e:race": "go test -race -tags=e2e -count=1 -parallel=4 ./test/e2e", diff --git a/packages/agenty-core/pkg/agentloop/compaction.go b/packages/agenty-core/pkg/agentloop/compaction.go index 072f0da..d6aca3e 100644 --- a/packages/agenty-core/pkg/agentloop/compaction.go +++ b/packages/agenty-core/pkg/agentloop/compaction.go @@ -121,7 +121,7 @@ func (engine *Engine) compactPreparedForWindow( SystemPrompt: prepared.systemPrompt, Messages: baseMessages, Tools: engine.toolDefinitions(prepared.freeFormTool), - MaxOutputTokens: maxOutputTokens, + MaxOutputTokens: prepared.maxOutputTokens, ReasoningEffort: preparedReasoningEffort(prepared), } contextTokensBefore := estimateRequestTokens(baseRequest) diff --git a/packages/agenty-core/pkg/agentloop/engine_test.go b/packages/agenty-core/pkg/agentloop/engine_test.go index 3b9f419..13e9205 100644 --- a/packages/agenty-core/pkg/agentloop/engine_test.go +++ b/packages/agenty-core/pkg/agentloop/engine_test.go @@ -724,8 +724,8 @@ func TestModelSwitchCompactsWithCurrentModelBeforePersistingTarget(t *testing.T) if len(caller.Requests()) != 1 { t.Fatalf("model switch LLM requests = %d, want 1 compaction request", len(caller.Requests())) } - if caller.Requests()[0].MaxOutputTokens != 1_024 { - t.Fatalf("model switch compaction max output = %d, want 1024", caller.Requests()[0].MaxOutputTokens) + if caller.Requests()[0].MaxOutputTokens != 8_192 { + t.Fatalf("model switch compaction max output = %d, want 8192", caller.Requests()[0].MaxOutputTokens) } events := fixture.sessions.events[session.ID] diff --git a/packages/agenty-core/pkg/domain/catalog/model.go b/packages/agenty-core/pkg/domain/catalog/model.go index acc2279..e87518e 100644 --- a/packages/agenty-core/pkg/domain/catalog/model.go +++ b/packages/agenty-core/pkg/domain/catalog/model.go @@ -18,6 +18,7 @@ type Model struct { Reasoning bool `json:"reasoning"` ReasoningEfforts []shared.ReasoningEffort `json:"reasoningEfforts"` IsDefault bool `json:"isDefault"` + Cached bool `json:"cached,omitempty"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt"` } diff --git a/packages/agenty-core/pkg/domain/catalog/provider.go b/packages/agenty-core/pkg/domain/catalog/provider.go index d578d1f..8dacf31 100644 --- a/packages/agenty-core/pkg/domain/catalog/provider.go +++ b/packages/agenty-core/pkg/domain/catalog/provider.go @@ -64,6 +64,13 @@ func (p *Provider) AddModel(m Model) { if m.MaxOutputTokens <= 0 { m.MaxOutputTokens = DefaultMaxOutputTokens } + if m.IsDefault { + for index := range p.Models { + if p.Models[index].Code != m.Code { + p.Models[index].IsDefault = false + } + } + } for i := range p.Models { if p.Models[i].Code == m.Code { p.Models[i] = m diff --git a/packages/agenty-core/pkg/domain/catalog/provider_test.go b/packages/agenty-core/pkg/domain/catalog/provider_test.go index b96e7d4..dd60800 100644 --- a/packages/agenty-core/pkg/domain/catalog/provider_test.go +++ b/packages/agenty-core/pkg/domain/catalog/provider_test.go @@ -52,3 +52,31 @@ func TestProvider_ModelLifecycle(t *testing.T) { t.Error("DefaultModel found a model after the default was removed") } } + +func TestProvider_AddModelSetsSingleDefault(t *testing.T) { + t.Parallel() + + provider := &Provider{Models: []Model{ + {Code: "model-a", Name: "A", IsDefault: true}, + {Code: "model-b", Name: "B"}, + }} + + provider.AddModel(Model{Code: "model-b", Name: "B", IsDefault: true}) + + modelA, err := provider.Model("model-a") + if err != nil { + t.Fatal(err) + } + modelB, err := provider.Model("model-b") + if err != nil { + t.Fatal(err) + } + if modelA.IsDefault || !modelB.IsDefault { + t.Fatalf("default flags = model-a:%t model-b:%t, want false/true", modelA.IsDefault, modelB.IsDefault) + } + + defaultModel, ok := provider.DefaultModel() + if !ok || defaultModel.Code != "model-b" { + t.Fatalf("default model = %+v, %t", defaultModel, ok) + } +} diff --git a/packages/agenty-core/pkg/infra/config/config.go b/packages/agenty-core/pkg/infra/config/config.go index b3a7ad1..63746bf 100644 --- a/packages/agenty-core/pkg/infra/config/config.go +++ b/packages/agenty-core/pkg/infra/config/config.go @@ -89,6 +89,7 @@ func InitializeDataDir() error { paths.SessionsDir, paths.AgentsDir, paths.ProvidersDir, + paths.LocksDir, } { if err := os.MkdirAll(dir, 0755); err != nil { return err @@ -163,6 +164,7 @@ func ResolvePaths() (*Paths, error) { SessionsDir: filepath.Join(dataDir, "sessions"), AgentsDir: filepath.Join(dataDir, "agents"), ProvidersDir: filepath.Join(dataDir, "providers"), + LocksDir: filepath.Join(dataDir, "locks"), DatabaseFile: filepath.Join(dataDir, "agenty.sqlite"), }, nil } diff --git a/packages/agenty-core/pkg/infra/config/config_test.go b/packages/agenty-core/pkg/infra/config/config_test.go index 18d5b6e..a8d4ce2 100644 --- a/packages/agenty-core/pkg/infra/config/config_test.go +++ b/packages/agenty-core/pkg/infra/config/config_test.go @@ -20,7 +20,7 @@ func TestInitializeDataDirCreatesStructure(t *testing.T) { t.Errorf("ResolvePaths: %v", err) } - for _, dir := range []string{paths.SessionsDir, paths.AgentsDir, paths.ProvidersDir} { + for _, dir := range []string{paths.SessionsDir, paths.AgentsDir, paths.ProvidersDir, paths.LocksDir} { if _, err := os.Stat(dir); os.IsNotExist(err) { t.Errorf("expected directory %s to exist", dir) } @@ -102,6 +102,9 @@ func TestResolvePathsUsesEnvVar(t *testing.T) { if paths.ConfigFile != filepath.Join(custom, "config.json") { t.Errorf("ConfigFile = %s", paths.ConfigFile) } + if paths.LocksDir != filepath.Join(custom, "locks") { + t.Errorf("LocksDir = %s", paths.LocksDir) + } } func TestLoadYAML(t *testing.T) { diff --git a/packages/agenty-core/pkg/infra/config/types.go b/packages/agenty-core/pkg/infra/config/types.go index b020965..ceed7db 100644 --- a/packages/agenty-core/pkg/infra/config/types.go +++ b/packages/agenty-core/pkg/infra/config/types.go @@ -43,6 +43,9 @@ type Paths struct { // ProvidersDir is DataDir/providers, where provider directories live. ProvidersDir string + // LocksDir is DataDir/locks, where apply_patch transaction locks live. + LocksDir string + // DatabaseFile is DataDir/agenty.sqlite. DatabaseFile string } diff --git a/packages/agenty-core/pkg/infra/storage/catalog.go b/packages/agenty-core/pkg/infra/storage/catalog.go index ae6a395..6f05bad 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog.go +++ b/packages/agenty-core/pkg/infra/storage/catalog.go @@ -88,6 +88,7 @@ func (r *CatalogRepository) getLocked(code shared.Code, includeDiscoveryCache bo return nil, err } provider.ModelsCached = false + clearCachedModelMarkers(&provider) normalizeModels(&provider) if includeDiscoveryCache { r.applyModelDiscoveryCache(&provider) @@ -195,6 +196,7 @@ func (r *CatalogRepository) Save(_ context.Context, provider *catalog.Provider) delete(r.modelCache, provider.Code) } providerToPersist.ModelsCached = false + clearCachedModelMarkers(providerToPersist) normalizeModels(providerToPersist) providerData, err := json.MarshalIndent(providerToPersist, "", " ") if err != nil { @@ -280,7 +282,9 @@ func (r *CatalogRepository) applyModelDiscoveryCache(provider *catalog.Provider) if _, exists := configuredCodes[model.Code]; exists { continue } - provider.Models = append(provider.Models, cloneModel(model)) + cachedModel := cloneModel(model) + cachedModel.Cached = true + provider.Models = append(provider.Models, cachedModel) provider.ModelsCached = true } } @@ -383,6 +387,12 @@ func normalizeModels(provider *catalog.Provider) { } } +func clearCachedModelMarkers(provider *catalog.Provider) { + for index := range provider.Models { + provider.Models[index].Cached = false + } +} + type providerCredentials struct { APIKey string `json:"apiKey"` } diff --git a/packages/agenty-core/pkg/infra/storage/catalog_test.go b/packages/agenty-core/pkg/infra/storage/catalog_test.go index 5c919b3..7240aee 100644 --- a/packages/agenty-core/pkg/infra/storage/catalog_test.go +++ b/packages/agenty-core/pkg/infra/storage/catalog_test.go @@ -230,6 +230,9 @@ func TestCatalogModelDiscoveryCacheUsesExpirationAndStaysInMemory(t *testing.T) if !loaded.ModelsCached { t.Fatal("cached provider did not expose the transient cache marker") } + if !loaded.Models[0].Cached { + t.Fatal("discovered model did not expose the transient cache marker") + } cachePath := filepath.Join(repo.providersDir, ".models", "openrouter.json") if _, err := os.Stat(cachePath); !os.IsNotExist(err) { @@ -295,6 +298,9 @@ func TestCatalogSaveDoesNotPersistDiscoveredModels(t *testing.T) { if !cached.ModelsCached || len(cached.Models) != 1 { t.Fatalf("cached provider = %+v", cached) } + if !cached.Models[0].Cached { + t.Fatal("discovered model did not expose the transient cache marker") + } if err := repo.Save(ctx, cached); err != nil { t.Fatal(err) @@ -319,6 +325,66 @@ func TestCatalogSaveDoesNotPersistDiscoveredModels(t *testing.T) { } } +func TestCatalogModelDiscoveryCacheMarksOnlyDiscoveredModels(t *testing.T) { + repo := newCatalogRepo(t) + ctx := context.Background() + provider, err := catalog.NewProvider("custom", "Custom", catalog.APIOpenAI) + if err != nil { + t.Fatal(err) + } + provider.Models = []catalog.Model{{ + Code: mustCatalogModelCode("configured"), + Name: "Configured", + ContextWindow: 128_000, + MaxOutputTokens: 8_192, + ReasoningEfforts: []shared.ReasoningEffort{}, + }} + if err := repo.Save(ctx, provider); err != nil { + t.Fatal(err) + } + + if err := repo.ReplaceModels(ctx, provider.Code, []catalog.Model{{ + Code: mustCatalogModelCode("discovered"), + Name: "Discovered", + ContextWindow: 128_000, + MaxOutputTokens: 8_192, + ReasoningEfforts: []shared.ReasoningEffort{}, + }}, time.Now().UTC().Add(time.Hour)); err != nil { + t.Fatal(err) + } + + loaded, err := repo.Get(ctx, provider.Code) + if err != nil { + t.Fatal(err) + } + configured, err := loaded.Model(mustCatalogModelCode("configured")) + if err != nil { + t.Fatal(err) + } + discovered, err := loaded.Model(mustCatalogModelCode("discovered")) + if err != nil { + t.Fatal(err) + } + if configured.Cached || !discovered.Cached { + t.Fatalf("cached flags = configured:%t discovered:%t, want false/true", configured.Cached, discovered.Cached) + } + + if err := repo.Save(ctx, loaded); err != nil { + t.Fatal(err) + } + data, err := os.ReadFile(filepath.Join(repo.providersDir, "custom.json")) + if err != nil { + t.Fatal(err) + } + var persisted catalog.Provider + if err := json.Unmarshal(data, &persisted); err != nil { + t.Fatal(err) + } + if len(persisted.Models) != 1 || persisted.Models[0].Code != configured.Code || persisted.Models[0].Cached { + t.Fatalf("persisted models = %#v", persisted.Models) + } +} + func TestCatalogList(t *testing.T) { repo := newCatalogRepo(t) ctx := context.Background() diff --git a/packages/agenty-core/scripts/build.mjs b/packages/agenty-core/scripts/build.mjs new file mode 100644 index 0000000..5c9afe8 --- /dev/null +++ b/packages/agenty-core/scripts/build.mjs @@ -0,0 +1,107 @@ +import { spawnSync } from "node:child_process"; +import { copyFileSync, existsSync, mkdirSync } from "node:fs"; +import { isAbsolute, join, resolve } from "node:path"; +import { fileURLToPath } from "node:url"; + +const PACKAGE_ROOT = resolve(import.meta.dirname, ".."); + +function targetForGoOS(rawOS) { + const lowerOS = rawOS.toLowerCase(); + if (lowerOS === "darwin" || lowerOS === "macos") { + return { artifactOS: "macos", extension: "", goOS: "darwin" }; + } + if (lowerOS.startsWith("win")) { + return { artifactOS: "windows", extension: ".exe", goOS: "windows" }; + } + if (lowerOS === "linux") { + return { artifactOS: "linux", extension: "", goOS: "linux" }; + } + throw new Error(`unsupported Go operating system: ${rawOS}`); +} + +function hostGoOS(hostPlatform) { + if (hostPlatform === "darwin") { + return "darwin"; + } + if (hostPlatform === "win32") { + return "windows"; + } + if (hostPlatform === "linux") { + return "linux"; + } + throw new Error(`unsupported host operating system: ${hostPlatform}`); +} + +function outputPath(packageRoot, path) { + return isAbsolute(path) ? path : resolve(packageRoot, path); +} + +function executableName(name, extension) { + if (extension === "" || name.toLowerCase().endsWith(extension)) { + return name; + } + return `${name}${extension}`; +} + +export function resolveCoreBuildPlan( + environment = process.env, + hostPlatform = process.platform, + packageRoot = PACKAGE_ROOT, +) { + const requestedGoOS = environment.GOOS?.trim() || hostGoOS(hostPlatform); + const target = targetForGoOS(requestedGoOS); + const outputDirectory = outputPath(packageRoot, environment.PACKAGE_DIR?.trim() || "bin"); + const coreName = executableName(environment.BIN_NAME?.trim() || "agenty-core", target.extension); + const helperName = `apply_patch${target.extension}`; + const repositoryRoot = resolve(packageRoot, "../.."); + return { + corePath: join(outputDirectory, coreName), + goArgs: ["build", "-o", join(outputDirectory, coreName), "./cmd"], + helperDestination: join(outputDirectory, helperName), + helperSource: join(repositoryRoot, "packages/patch-applier/target/release", helperName), + outputDirectory, + packageRoot, + target, + }; +} + +function exitCode(label, result) { + if (result.error) { + throw result.error; + } + if (result.status === null) { + const detail = result.signal ? ` (signal ${result.signal})` : ""; + throw new Error(`${label} did not return an exit code${detail}`); + } + return result.status; +} + +function run() { + const plan = resolveCoreBuildPlan(); + mkdirSync(plan.outputDirectory, { recursive: true }); + const build = spawnSync("go", plan.goArgs, { + cwd: plan.packageRoot, + env: process.env, + stdio: "inherit", + }); + const buildExitCode = exitCode("go build", build); + if (buildExitCode !== 0) { + return buildExitCode; + } + if (!existsSync(plan.helperSource)) { + throw new Error(`apply_patch binary not found at ${plan.helperSource}`); + } + copyFileSync(plan.helperSource, plan.helperDestination); + console.log(`agenty-core built -> ${plan.corePath}`); + return 0; +} + +const currentFile = fileURLToPath(import.meta.url); +if (process.argv[1] && resolve(process.argv[1]) === currentFile) { + try { + process.exitCode = run(); + } catch (error) { + console.error(error instanceof Error ? error.message : error); + process.exitCode = 1; + } +} diff --git a/packages/agenty-core/scripts/build.test.mjs b/packages/agenty-core/scripts/build.test.mjs new file mode 100644 index 0000000..c57029d --- /dev/null +++ b/packages/agenty-core/scripts/build.test.mjs @@ -0,0 +1,37 @@ +import assert from "node:assert/strict"; +import { join, resolve } from "node:path"; +import test from "node:test"; + +import { resolveCoreBuildPlan } from "./build.mjs"; + +const repositoryRoot = resolve("/repo"); +const packageRoot = join(repositoryRoot, "packages/agenty-core"); + +test("uses host defaults for a macOS core build", () => { + const plan = resolveCoreBuildPlan({}, "darwin", packageRoot); + + assert.deepEqual(plan.target, { artifactOS: "macos", extension: "", goOS: "darwin" }); + assert.equal(plan.corePath, join(packageRoot, "bin/agenty-core")); + assert.equal(plan.helperSource, join(repositoryRoot, "packages/patch-applier/target/release/apply_patch")); + assert.equal(plan.helperDestination, join(packageRoot, "bin/apply_patch")); +}); + +test("uses the GOOS target and Windows extensions", () => { + const plan = resolveCoreBuildPlan({ + BIN_NAME: "core", + GOOS: "windows", + PACKAGE_DIR: "bin/windows_amd64", + }, "darwin", packageRoot); + + assert.deepEqual(plan.target, { artifactOS: "windows", extension: ".exe", goOS: "windows" }); + assert.equal(plan.corePath, join(packageRoot, "bin/windows_amd64/core.exe")); + assert.equal(plan.helperSource, join(repositoryRoot, "packages/patch-applier/target/release/apply_patch.exe")); + assert.equal(plan.helperDestination, join(packageRoot, "bin/windows_amd64/apply_patch.exe")); +}); + +test("detects a Windows host when GOOS is not set", () => { + const plan = resolveCoreBuildPlan({}, "win32", packageRoot); + + assert.equal(plan.target.goOS, "windows"); + assert.equal(plan.corePath, join(packageRoot, "bin/agenty-core.exe")); +}); diff --git a/packages/patch-applier/Cargo.lock b/packages/patch-applier/Cargo.lock index 22061e7..b6456fe 100644 --- a/packages/patch-applier/Cargo.lock +++ b/packages/patch-applier/Cargo.lock @@ -2,12 +2,77 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "hybrid-array", +] + +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "const-oid", + "crypto-common", +] + +[[package]] +name = "hybrid-array" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "707114b52a152fa7bdb290cd7cd5912d9467273b6d74e21b8d81aca1f8533f6b" +dependencies = [ + "typenum", +] + [[package]] name = "itoa" version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "keccak" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d8f198d1db720e4940b5a493201d199d9f24f568f8f746bd13706243a2f71598" +dependencies = [ + "cfg-if", + "cpufeatures", +] + +[[package]] +name = "libc" +version = "0.2.189" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" + [[package]] name = "memchr" version = "2.8.3" @@ -20,6 +85,7 @@ version = "0.1.0" dependencies = [ "serde", "serde_json", + "sha3", "similar", ] @@ -84,12 +150,29 @@ dependencies = [ "zmij", ] +[[package]] +name = "sha3" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc9bad02c26382724b2d2692c6f179285e4b54eeecd7968f52a50059c3c11759" +dependencies = [ + "digest", + "keccak", + "sponge-cursor", +] + [[package]] name = "similar" version = "2.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" +[[package]] +name = "sponge-cursor" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a0219bd7d979d58245a4f41f695e1ac9f8befdffadd7f61f1bae9e39abc6620" + [[package]] name = "syn" version = "3.0.4" @@ -101,6 +184,12 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + [[package]] name = "unicode-ident" version = "1.0.24" diff --git a/packages/patch-applier/Cargo.toml b/packages/patch-applier/Cargo.toml index 6cd75c4..6bfcd27 100644 --- a/packages/patch-applier/Cargo.toml +++ b/packages/patch-applier/Cargo.toml @@ -12,6 +12,7 @@ path = "src/main.rs" [dependencies] serde = { version = "1", features = ["derive"] } serde_json = "1" +sha3 = "0.12" similar = "2" [profile.release] diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index 4964142..a2eef5a 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -1,12 +1,12 @@ use std::collections::{HashMap, HashSet}; -use std::ffi::OsString; use std::fmt::{self, Display, Formatter}; use std::fs::{self, File, OpenOptions}; use std::io::{self, Write}; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicU64, Ordering}; -use serde::Serialize; +use serde::{Deserialize, Serialize}; +use sha3::{Digest, Sha3_256}; use similar::{ChangeTag, TextDiff}; const BEGIN_MARKER: &str = "*** Begin Patch"; @@ -144,7 +144,8 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { let results = transaction.prepare_results()?; let lock_paths = transaction.lock_paths()?; - let locks = FileLocks::acquire(&lock_paths, std::process::id())?; + let lock_directory = lock_directory()?; + let locks = FileLocks::acquire(&lock_directory, &lock_paths, std::process::id())?; let commit_result = transaction.commit(); let release_result = locks.release(); commit_result?; @@ -159,36 +160,27 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { struct FileLocks { paths: Vec, - created_dirs: Vec, +} + +#[derive(Debug, Deserialize, Serialize)] +struct LockOwner { + pid: u32, + path: String, } impl FileLocks { - fn acquire(targets: &[PathBuf], pid: u32) -> Result { + fn acquire(lock_directory: &Path, targets: &[PathBuf], pid: u32) -> Result { + fs::create_dir_all(lock_directory)?; + let mut targets = targets.to_vec(); targets.sort(); targets.dedup(); let mut locks = Self { paths: Vec::with_capacity(targets.len()), - created_dirs: Vec::new(), }; for target in targets { - let parent = target.parent().unwrap_or_else(|| Path::new(".")); - match create_missing_dirs(parent) { - Ok(created) => locks.created_dirs.extend(created), - Err(error) => { - locks.cleanup(); - return Err(error); - } - } - - let lock_path = match lock_path(&target) { - Ok(path) => path, - Err(error) => { - locks.cleanup(); - return Err(error); - } - }; + let lock_path = lock_path(lock_directory, &target); let result = (|| -> Result<(), PatchError> { let mut lock = OpenOptions::new() .write(true) @@ -198,14 +190,22 @@ impl FileLocks { if error.kind() == io::ErrorKind::AlreadyExists { let owner = read_lock_owner(&lock_path); return PatchError::Conflict(format!( - "file {} is locked by apply_patch pid {owner} ({})", + "file {} is locked by apply_patch {owner} ({})", target.display(), lock_path.display(), )); } PatchError::Io(error) })?; - writeln!(lock, "{pid}")?; + let owner = LockOwner { + pid, + path: target.display().to_string(), + }; + let data = serde_json::to_vec(&owner).map_err(|error| { + PatchError::Invalid(format!("serialize lock owner: {error}")) + })?; + lock.write_all(&data)?; + lock.write_all(b"\n")?; lock.sync_all()?; Ok(()) })(); @@ -231,15 +231,6 @@ impl FileLocks { if let Err(error) = sync_lock_parent_directories(&self.paths) { first_error.get_or_insert(error); } - for directory in self.created_dirs.iter().rev() { - if let Err(error) = fs::remove_dir(directory) { - if error.kind() != io::ErrorKind::NotFound - && error.kind() != io::ErrorKind::DirectoryNotEmpty - { - first_error.get_or_insert(PatchError::Io(error)); - } - } - } first_error.map_or(Ok(()), Err) } @@ -247,33 +238,43 @@ impl FileLocks { for path in self.paths.drain(..).rev() { let _ = fs::remove_file(path); } - for directory in self.created_dirs.drain(..).rev() { - let _ = fs::remove_dir(directory); - } } } -fn lock_path(target: &Path) -> Result { - let file_name = target.file_name().ok_or_else(|| { - PatchError::Invalid(format!( - "cannot create a file lock for {}", - target.display() - )) - })?; - let mut lock_name = OsString::from("."); - lock_name.push(file_name); - lock_name.push(".lock"); - Ok(target - .parent() - .unwrap_or_else(|| Path::new(".")) - .join(lock_name)) +fn lock_directory() -> Result { + let data_directory = std::env::var_os("AGENTY_DATA_DIR") + .filter(|directory| !directory.is_empty()) + .ok_or_else(|| { + PatchError::Invalid("AGENTY_DATA_DIR is required for apply_patch locks".to_string()) + })?; + let data_directory = PathBuf::from(data_directory); + if !data_directory.is_absolute() { + return Err(PatchError::Invalid( + "AGENTY_DATA_DIR must be an absolute path for apply_patch locks".to_string(), + )); + } + Ok(data_directory.join("locks")) +} + +fn lock_path(lock_directory: &Path, target: &Path) -> PathBuf { + let mut digest = Sha3_256::new(); + digest.update(target.as_os_str().as_encoded_bytes()); + let digest = digest.finalize(); + let mut name = String::with_capacity(digest.len() * 2 + ".lock".len()); + for byte in digest { + const HEX: &[u8; 16] = b"0123456789abcdef"; + name.push(HEX[(byte >> 4) as usize] as char); + name.push(HEX[(byte & 0x0f) as usize] as char); + } + name.push_str(".lock"); + lock_directory.join(name) } fn read_lock_owner(path: &Path) -> String { - fs::read_to_string(path) - .map(|contents| contents.trim().to_string()) + fs::read(path) .ok() - .filter(|contents| !contents.is_empty()) + .and_then(|contents| serde_json::from_slice::(&contents).ok()) + .map(|owner| format!("pid {} ({})", owner.pid, owner.path)) .unwrap_or_else(|| "unknown".to_string()) } @@ -1262,9 +1263,26 @@ impl MetadataMode for fs::Metadata { mod tests { use super::*; use std::fs; + use std::sync::OnceLock; use std::time::{SystemTime, UNIX_EPOCH}; + fn test_data_dir() -> &'static PathBuf { + static DATA_DIR: OnceLock = OnceLock::new(); + + DATA_DIR.get_or_init(|| { + let suffix = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos(); + let path = std::env::temp_dir().join(format!("agenty-patch-data-{suffix}")); + fs::create_dir_all(&path).unwrap(); + std::env::set_var("AGENTY_DATA_DIR", &path); + path + }) + } + fn temp_dir(name: &str) -> PathBuf { + let _ = test_data_dir(); let suffix = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() @@ -1398,9 +1416,13 @@ mod tests { fn creates_pid_lock_for_each_transaction_target_and_removes_it() { let cwd = temp_dir("file-lock"); let target = cwd.join("notes.txt"); - let locks = FileLocks::acquire(std::slice::from_ref(&target), 4242).unwrap(); - let lock = cwd.join(".notes.txt.lock"); - assert_eq!(fs::read_to_string(&lock).unwrap(), "4242\n"); + let lock_directory = test_data_dir().join("locks"); + let locks = + FileLocks::acquire(&lock_directory, std::slice::from_ref(&target), 4242).unwrap(); + let lock = lock_path(&lock_directory, &target); + let owner: LockOwner = serde_json::from_slice(&fs::read(&lock).unwrap()).unwrap(); + assert_eq!(owner.pid, 4242); + assert_eq!(owner.path, target.display().to_string()); locks.release().unwrap(); assert!(!lock.exists()); } @@ -1410,7 +1432,18 @@ mod tests { let cwd = temp_dir("locked-target"); let target = cwd.join("notes.txt"); fs::write(&target, "one\n").unwrap(); - fs::write(cwd.join(".notes.txt.lock"), "9876\n").unwrap(); + let lock_directory = test_data_dir().join("locks"); + fs::create_dir_all(&lock_directory).unwrap(); + let lock = lock_path(&lock_directory, &target); + fs::write( + &lock, + serde_json::to_vec(&LockOwner { + pid: 9876, + path: target.display().to_string(), + }) + .unwrap(), + ) + .unwrap(); let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** End Patch"; let error = apply_patch(&cwd, patch).unwrap_err().to_string(); @@ -1424,7 +1457,18 @@ mod tests { let source = cwd.join("old.txt"); let destination = cwd.join("new.txt"); fs::write(&source, "one\n").unwrap(); - fs::write(cwd.join(".new.txt.lock"), "1111\n").unwrap(); + let lock_directory = test_data_dir().join("locks"); + fs::create_dir_all(&lock_directory).unwrap(); + let lock = lock_path(&lock_directory, &destination); + fs::write( + &lock, + serde_json::to_vec(&LockOwner { + pid: 1111, + path: destination.display().to_string(), + }) + .unwrap(), + ) + .unwrap(); let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** End Patch"; let error = apply_patch(&cwd, patch).unwrap_err().to_string(); @@ -1433,6 +1477,24 @@ mod tests { assert!(!destination.exists()); } + #[test] + fn allows_repository_files_with_the_old_lock_name() { + let cwd = temp_dir("repository-lock-name"); + let target = cwd.join("notes.txt"); + let repository_file = cwd.join(".notes.txt.lock"); + fs::write(&target, "one\n").unwrap(); + fs::write(&repository_file, "repository data\n").unwrap(); + + let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** End Patch"; + apply_patch(&cwd, patch).unwrap(); + + assert_eq!(fs::read_to_string(target).unwrap(), "two\n"); + assert_eq!( + fs::read_to_string(repository_file).unwrap(), + "repository data\n" + ); + } + #[cfg(unix)] #[test] fn restores_backups_when_staging_fails() { diff --git a/packages/patch-applier/tests/cli.rs b/packages/patch-applier/tests/cli.rs index 3d48e62..3f9e26d 100644 --- a/packages/patch-applier/tests/cli.rs +++ b/packages/patch-applier/tests/cli.rs @@ -18,8 +18,10 @@ fn temp_dir(name: &str) -> std::path::PathBuf { #[test] fn prints_complete_success_json() { let cwd = temp_dir("success"); + let data_dir = cwd.join("data"); let mut child = Command::new(env!("CARGO_BIN_EXE_apply_patch")) .current_dir(&cwd) + .env("AGENTY_DATA_DIR", &data_dir) .stdin(Stdio::piped()) .stdout(Stdio::piped()) .stderr(Stdio::piped()) @@ -56,8 +58,10 @@ fn prints_complete_success_json() { #[test] fn prints_complete_error_json() { let cwd = temp_dir("error"); + let data_dir = cwd.join("data"); let mut child = Command::new(env!("CARGO_BIN_EXE_apply_patch")) .current_dir(&cwd) + .env("AGENTY_DATA_DIR", &data_dir) .stdin(Stdio::piped()) .stdout(Stdio::piped()) .stderr(Stdio::piped()) diff --git a/scripts/dev.mjs b/scripts/dev.mjs new file mode 100644 index 0000000..ee5919f --- /dev/null +++ b/scripts/dev.mjs @@ -0,0 +1,77 @@ +import { spawnSync } from "node:child_process"; +import { existsSync } from "node:fs"; +import { resolve } from "node:path"; +import { fileURLToPath } from "node:url"; + +import { executableExtension, resolveArch, resolveOS } from "./platform.mjs"; + +const REPOSITORY_ROOT = resolve(import.meta.dirname, ".."); + +export function packageManagerCommand(hostPlatform = process.platform) { + return hostPlatform === "win32" ? "pnpm.cmd" : "pnpm"; +} + +export function resolveDevPlan( + args, + environment = process.env, + hostPlatform = process.platform, + hostArch = process.arch, + repositoryRoot = REPOSITORY_ROOT, +) { + const os = resolveOS(environment, hostPlatform); + const arch = resolveArch(environment, hostArch); + const launcher = resolve( + repositoryRoot, + "packages/agenty-bootstrap/bin", + `agenty-${os}-${arch}${executableExtension(os)}`, + ); + return { + buildArgs: ["exec", "turbo", "run", "build", "--filter=agenty-bootstrap"], + launcher, + launcherArgs: args, + }; +} + +function exitCode(label, result) { + if (result.error) { + throw result.error; + } + if (result.status === null) { + const detail = result.signal ? ` (signal ${result.signal})` : ""; + throw new Error(`${label} did not return an exit code${detail}`); + } + return result.status; +} + +function run() { + const plan = resolveDevPlan(process.argv.slice(2)); + const build = spawnSync(packageManagerCommand(), plan.buildArgs, { + cwd: REPOSITORY_ROOT, + env: process.env, + stdio: "inherit", + }); + const buildExitCode = exitCode("bootstrap build", build); + if (buildExitCode !== 0) { + return buildExitCode; + } + if (!existsSync(plan.launcher)) { + throw new Error(`bootstrap launcher not found at ${plan.launcher}`); + } + + const launcher = spawnSync(plan.launcher, plan.launcherArgs, { + cwd: REPOSITORY_ROOT, + env: process.env, + stdio: "inherit", + }); + return exitCode("bootstrap launcher", launcher); +} + +const currentFile = fileURLToPath(import.meta.url); +if (process.argv[1] && resolve(process.argv[1]) === currentFile) { + try { + process.exitCode = run(); + } catch (error) { + console.error(error instanceof Error ? error.message : error); + process.exitCode = 1; + } +} diff --git a/scripts/dev.test.mjs b/scripts/dev.test.mjs new file mode 100644 index 0000000..0b12ba8 --- /dev/null +++ b/scripts/dev.test.mjs @@ -0,0 +1,21 @@ +import assert from "node:assert/strict"; +import { resolve } from "node:path"; +import test from "node:test"; + +import { packageManagerCommand, resolveDevPlan } from "./dev.mjs"; + +test("selects the launcher for the host target and forwards arguments", () => { + const plan = resolveDevPlan(["--version"], {}, "darwin", "arm64", "/repo"); + + assert.deepEqual(plan.buildArgs, ["exec", "turbo", "run", "build", "--filter=agenty-bootstrap"]); + assert.equal(plan.launcher, resolve("/repo/packages/agenty-bootstrap/bin/agenty-macos-arm64")); + assert.deepEqual(plan.launcherArgs, ["--version"]); +}); + +test("uses a Windows launcher extension and package manager command", () => { + const plan = resolveDevPlan([], { ARCH: "amd64", OS: "Windows_NT" }, "win32", "x64", "/repo"); + + assert.equal(plan.launcher, resolve("/repo/packages/agenty-bootstrap/bin/agenty-windows-amd64.exe")); + assert.equal(packageManagerCommand("win32"), "pnpm.cmd"); + assert.equal(packageManagerCommand("linux"), "pnpm"); +}); diff --git a/scripts/platform.mjs b/scripts/platform.mjs new file mode 100644 index 0000000..f5538a8 --- /dev/null +++ b/scripts/platform.mjs @@ -0,0 +1,44 @@ +/** @typedef {"linux" | "macos" | "windows"} TargetOS */ +/** @typedef {"amd64" | "arm64"} TargetArch */ + +/** + * @param {Record} environment + * @param {string} hostArch + * @returns {TargetArch} + */ +export function resolveArch(environment = process.env, hostArch = process.arch) { + const rawArch = environment.ARCH?.trim() || hostArch; + const lowerArch = rawArch.toLowerCase(); + if (lowerArch === "x64" || lowerArch === "x86_64" || lowerArch === "amd64") { + return "amd64"; + } + if (lowerArch === "arm64" || lowerArch === "aarch64") { + return "arm64"; + } + throw new Error(`unsupported architecture: ${rawArch}`); +} + +/** + * @param {Record} environment + * @param {string} hostPlatform + * @returns {TargetOS} + */ +export function resolveOS(environment = process.env, hostPlatform = process.platform) { + const rawOS = environment.OS?.trim() || hostPlatform; + const lowerOS = rawOS.toLowerCase(); + if (lowerOS === "darwin" || lowerOS === "macos") { + return "macos"; + } + if (lowerOS.startsWith("win")) { + return "windows"; + } + if (lowerOS === "linux") { + return "linux"; + } + throw new Error(`unsupported operating system: ${rawOS}`); +} + +/** @param {TargetOS} os */ +export function executableExtension(os) { + return os === "windows" ? ".exe" : ""; +} From 9337378e3063ed95f9d81757f8b848d98a6d53e5 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Thu, 27 Aug 2026 20:29:18 +0800 Subject: [PATCH 10/12] fix: optimize apply_patch tool code structure Signed-off-by: masteryyh --- packages/patch-applier/Cargo.toml | 1 + packages/patch-applier/README.md | 7 + packages/patch-applier/src/file_lock.rs | 213 +++++++++ packages/patch-applier/src/lib.rs | 559 ++-------------------- packages/patch-applier/src/transaction.rs | 350 ++++++++++++++ 5 files changed, 600 insertions(+), 530 deletions(-) create mode 100644 packages/patch-applier/src/file_lock.rs create mode 100644 packages/patch-applier/src/transaction.rs diff --git a/packages/patch-applier/Cargo.toml b/packages/patch-applier/Cargo.toml index 6bfcd27..83ea4ef 100644 --- a/packages/patch-applier/Cargo.toml +++ b/packages/patch-applier/Cargo.toml @@ -2,6 +2,7 @@ name = "patch-applier" version = "0.1.0" edition = "2021" +rust-version = "1.89" license = "Apache-2.0" description = "Atomic V4A patch applier for Agenty" diff --git a/packages/patch-applier/README.md b/packages/patch-applier/README.md index 6d711a1..f023c58 100644 --- a/packages/patch-applier/README.md +++ b/packages/patch-applier/README.md @@ -10,6 +10,13 @@ snapshot, and rejects incompatible state transitions or path ownership conflicts commit phase stages all new contents before replacing or deleting targets and rolls back completed replacements if a later filesystem operation fails. +Filesystem coordination uses persistent lock files under `AGENTY_DATA_DIR/locks`. Each patch +holds shared advisory locks while reading and preparing its snapshot, then acquires exclusive +locks for the paths it will change before revalidating and committing. Lock files are coordination +artifacts and are intentionally retained after the process exits. Paths are acquired in sorted +order, and conflicting acquisitions wait until the existing holder releases its lock. The +operating system releases the advisory lock when the owning `File` handle closes. + Success and failure both write one JSON object to stdout. A successful result includes the final unified diff and added/removed line counts for every changed path: diff --git a/packages/patch-applier/src/file_lock.rs b/packages/patch-applier/src/file_lock.rs new file mode 100644 index 0000000..d8129c5 --- /dev/null +++ b/packages/patch-applier/src/file_lock.rs @@ -0,0 +1,213 @@ +use std::fs::{self, File, OpenOptions}; +use std::io; +use std::path::{Path, PathBuf}; + +use sha3::{Digest, Sha3_256}; + +use super::{FileGroup, PatchError}; + +#[derive(Debug, Clone, Copy)] +pub(super) enum LockMode { + Shared, + Exclusive, +} + +pub(super) struct FileLocks { + locks: Vec, +} + +impl FileLocks { + pub(super) fn acquire( + lock_directory: &Path, + targets: &[PathBuf], + mode: LockMode, + ) -> Result { + fs::create_dir_all(lock_directory)?; + + let mut targets = targets.to_vec(); + targets.sort(); + targets.dedup(); + + let mut locks = Self { + locks: Vec::with_capacity(targets.len()), + }; + for target in targets { + let lock_path = lock_path(lock_directory, &target); + let lock = OpenOptions::new() + .read(true) + .write(true) + .create(true) + .truncate(false) + .open(&lock_path) + .map_err(|error| lock_error(&target, "open", error))?; + let result = match mode { + LockMode::Shared => lock.lock_shared(), + LockMode::Exclusive => lock.lock(), + }; + result.map_err(|error| lock_error(&target, "acquire", error))?; + locks.locks.push(FileLock { + path: lock_path, + file: lock, + }); + } + + Ok(locks) + } + + pub(super) fn release(self) -> Result<(), PatchError> { + let mut first_error = None; + for lock in self.locks.iter().rev() { + if let Err(error) = lock.file.unlock() { + first_error.get_or_insert(lock_error(&lock.path, "release", error)); + } + } + first_error.map_or(Ok(()), Err) + } +} + +struct FileLock { + path: PathBuf, + file: File, +} + +fn lock_error(path: &Path, action: &str, error: io::Error) -> PatchError { + PatchError::Io(io::Error::new( + error.kind(), + format!("{action} lock for {}: {error}", path.display()), + )) +} + +pub(super) fn operation_lock_paths(groups: &[FileGroup]) -> Vec { + groups + .iter() + .flat_map(|group| { + group.operations.iter().flat_map(|operation| { + std::iter::once(operation.path.clone()).chain(operation.move_to.clone()) + }) + }) + .collect() +} + +pub(super) fn lock_directory() -> Result { + let data_directory = std::env::var_os("AGENTY_DATA_DIR") + .filter(|directory| !directory.is_empty()) + .ok_or_else(|| { + PatchError::Invalid("AGENTY_DATA_DIR is required for apply_patch locks".to_string()) + })?; + let data_directory = PathBuf::from(data_directory); + if !data_directory.is_absolute() { + return Err(PatchError::Invalid( + "AGENTY_DATA_DIR must be an absolute path for apply_patch locks".to_string(), + )); + } + Ok(data_directory.join("locks")) +} + +pub(super) fn lock_path(lock_directory: &Path, target: &Path) -> PathBuf { + let mut digest = Sha3_256::new(); + digest.update(target.as_os_str().as_encoded_bytes()); + let digest = digest.finalize(); + let mut name = String::with_capacity(digest.len() * 2 + ".lock".len()); + for byte in digest { + const HEX: &[u8; 16] = b"0123456789abcdef"; + name.push(HEX[(byte >> 4) as usize] as char); + name.push(HEX[(byte & 0x0f) as usize] as char); + } + name.push_str(".lock"); + lock_directory.join(name) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::time::{SystemTime, UNIX_EPOCH}; + + fn temp_lock_root(name: &str) -> PathBuf { + let suffix = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock must follow the Unix epoch") + .as_nanos(); + let root = std::env::temp_dir().join(format!("agenty-patch-lock-{name}-{suffix}")); + fs::create_dir_all(&root).expect("create lock test directory"); + root + } + + #[test] + fn holds_an_exclusive_advisory_lock_in_a_persistent_file() { + let root = temp_lock_root("exclusive"); + let target = root.join("notes.txt"); + let lock_directory = root.join("locks"); + let locks = FileLocks::acquire( + &lock_directory, + std::slice::from_ref(&target), + LockMode::Exclusive, + ) + .unwrap(); + let lock = lock_path(&lock_directory, &target); + assert!(lock.exists()); + + let probe = OpenOptions::new() + .read(true) + .write(true) + .open(&lock) + .unwrap(); + assert!(matches!( + probe.try_lock(), + Err(std::fs::TryLockError::WouldBlock) + )); + + locks.release().unwrap(); + assert!(lock.exists()); + probe.try_lock().unwrap(); + probe.unlock().unwrap(); + } + + #[test] + fn shared_locks_can_coexist_but_exclude_writers() { + let root = temp_lock_root("shared"); + let target = root.join("notes.txt"); + let lock_directory = root.join("locks"); + let shared_a = FileLocks::acquire( + &lock_directory, + std::slice::from_ref(&target), + LockMode::Shared, + ) + .unwrap(); + let shared_b = FileLocks::acquire( + &lock_directory, + std::slice::from_ref(&target), + LockMode::Shared, + ) + .unwrap(); + let lock = lock_path(&lock_directory, &target); + let probe = OpenOptions::new() + .read(true) + .write(true) + .open(&lock) + .unwrap(); + + assert!(matches!( + probe.try_lock(), + Err(std::fs::TryLockError::WouldBlock) + )); + shared_a.release().unwrap(); + shared_b.release().unwrap(); + + probe.try_lock().unwrap(); + probe.unlock().unwrap(); + + let exclusive = FileLocks::acquire( + &lock_directory, + std::slice::from_ref(&target), + LockMode::Exclusive, + ) + .unwrap(); + assert!(matches!( + probe.try_lock_shared(), + Err(std::fs::TryLockError::WouldBlock) + )); + exclusive.release().unwrap(); + probe.try_lock_shared().unwrap(); + probe.unlock().unwrap(); + } +} diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index a2eef5a..2b0a89b 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -1,13 +1,18 @@ -use std::collections::{HashMap, HashSet}; +use std::collections::HashMap; use std::fmt::{self, Display, Formatter}; -use std::fs::{self, File, OpenOptions}; -use std::io::{self, Write}; +use std::fs::{self, File}; +use std::io; use std::path::{Path, PathBuf}; -use std::sync::atomic::{AtomicU64, Ordering}; -use serde::{Deserialize, Serialize}; -use sha3::{Digest, Sha3_256}; +mod file_lock; +mod transaction; + +#[cfg(test)] +use file_lock::lock_path; +use file_lock::{lock_directory, operation_lock_paths, FileLocks, LockMode}; +use serde::Serialize; use similar::{ChangeTag, TextDiff}; +use transaction::Transaction; const BEGIN_MARKER: &str = "*** Begin Patch"; const END_MARKER: &str = "*** End Patch"; @@ -17,8 +22,6 @@ const ADD_MARKER: &str = "*** Add File:"; const MOVE_MARKER: &str = "*** Move to:"; const END_OF_FILE_MARKER: &str = "*** End of File"; -static TRANSACTION_COUNTER: AtomicU64 = AtomicU64::new(0); - #[derive(Debug)] pub enum PatchError { Io(io::Error), @@ -134,6 +137,9 @@ pub struct FileResult { pub fn apply_patch(cwd: &Path, patch: &str) -> Result { let operations = parse_envelope(cwd, patch)?; let groups = classify_operations(operations)?; + let read_lock_paths = operation_lock_paths(&groups); + let lock_directory = lock_directory()?; + let read_locks = FileLocks::acquire(&lock_directory, &read_lock_paths, LockMode::Shared)?; let mut transaction = Transaction::new(cwd.to_path_buf()); for group in groups { @@ -143,11 +149,12 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { } let results = transaction.prepare_results()?; + read_locks.release()?; + let lock_paths = transaction.lock_paths()?; - let lock_directory = lock_directory()?; - let locks = FileLocks::acquire(&lock_directory, &lock_paths, std::process::id())?; + let write_locks = FileLocks::acquire(&lock_directory, &lock_paths, LockMode::Exclusive)?; let commit_result = transaction.commit(); - let release_result = locks.release(); + let release_result = write_locks.release(); commit_result?; release_result?; @@ -158,145 +165,6 @@ pub fn apply_patch(cwd: &Path, patch: &str) -> Result { }) } -struct FileLocks { - paths: Vec, -} - -#[derive(Debug, Deserialize, Serialize)] -struct LockOwner { - pid: u32, - path: String, -} - -impl FileLocks { - fn acquire(lock_directory: &Path, targets: &[PathBuf], pid: u32) -> Result { - fs::create_dir_all(lock_directory)?; - - let mut targets = targets.to_vec(); - targets.sort(); - targets.dedup(); - - let mut locks = Self { - paths: Vec::with_capacity(targets.len()), - }; - for target in targets { - let lock_path = lock_path(lock_directory, &target); - let result = (|| -> Result<(), PatchError> { - let mut lock = OpenOptions::new() - .write(true) - .create_new(true) - .open(&lock_path) - .map_err(|error| { - if error.kind() == io::ErrorKind::AlreadyExists { - let owner = read_lock_owner(&lock_path); - return PatchError::Conflict(format!( - "file {} is locked by apply_patch {owner} ({})", - target.display(), - lock_path.display(), - )); - } - PatchError::Io(error) - })?; - let owner = LockOwner { - pid, - path: target.display().to_string(), - }; - let data = serde_json::to_vec(&owner).map_err(|error| { - PatchError::Invalid(format!("serialize lock owner: {error}")) - })?; - lock.write_all(&data)?; - lock.write_all(b"\n")?; - lock.sync_all()?; - Ok(()) - })(); - if let Err(error) = result { - locks.cleanup(); - return Err(error); - } - locks.paths.push(lock_path); - } - - Ok(locks) - } - - fn release(self) -> Result<(), PatchError> { - let mut first_error = None; - for path in self.paths.iter().rev() { - if let Err(error) = fs::remove_file(path) { - if error.kind() != io::ErrorKind::NotFound { - first_error.get_or_insert(PatchError::Io(error)); - } - } - } - if let Err(error) = sync_lock_parent_directories(&self.paths) { - first_error.get_or_insert(error); - } - first_error.map_or(Ok(()), Err) - } - - fn cleanup(&mut self) { - for path in self.paths.drain(..).rev() { - let _ = fs::remove_file(path); - } - } -} - -fn lock_directory() -> Result { - let data_directory = std::env::var_os("AGENTY_DATA_DIR") - .filter(|directory| !directory.is_empty()) - .ok_or_else(|| { - PatchError::Invalid("AGENTY_DATA_DIR is required for apply_patch locks".to_string()) - })?; - let data_directory = PathBuf::from(data_directory); - if !data_directory.is_absolute() { - return Err(PatchError::Invalid( - "AGENTY_DATA_DIR must be an absolute path for apply_patch locks".to_string(), - )); - } - Ok(data_directory.join("locks")) -} - -fn lock_path(lock_directory: &Path, target: &Path) -> PathBuf { - let mut digest = Sha3_256::new(); - digest.update(target.as_os_str().as_encoded_bytes()); - let digest = digest.finalize(); - let mut name = String::with_capacity(digest.len() * 2 + ".lock".len()); - for byte in digest { - const HEX: &[u8; 16] = b"0123456789abcdef"; - name.push(HEX[(byte >> 4) as usize] as char); - name.push(HEX[(byte & 0x0f) as usize] as char); - } - name.push_str(".lock"); - lock_directory.join(name) -} - -fn read_lock_owner(path: &Path) -> String { - fs::read(path) - .ok() - .and_then(|contents| serde_json::from_slice::(&contents).ok()) - .map(|owner| format!("pid {} ({})", owner.pid, owner.path)) - .unwrap_or_else(|| "unknown".to_string()) -} - -#[cfg(unix)] -fn sync_lock_parent_directories(paths: &[PathBuf]) -> Result<(), PatchError> { - let mut parents = HashSet::new(); - for path in paths { - if let Some(parent) = path.parent() { - parents.insert(parent); - } - } - for parent in parents { - File::open(parent)?.sync_all()?; - } - Ok(()) -} - -#[cfg(not(unix))] -fn sync_lock_parent_directories(_paths: &[PathBuf]) -> Result<(), PatchError> { - Ok(()) -} - fn parse_envelope(cwd: &Path, patch: &str) -> Result, PatchError> { let lines = normalized_lines(patch); if lines.len() < 3 || lines.first().map(String::as_str) != Some(BEGIN_MARKER) { @@ -485,342 +353,6 @@ fn classify_operations(operations: Vec) -> Result, Pat Ok(groups) } -struct Transaction { - cwd: PathBuf, - original: HashMap, - state: HashMap, - touched_order: Vec, - deleted_paths: HashSet, -} - -impl Transaction { - fn new(cwd: PathBuf) -> Self { - Self { - cwd, - original: HashMap::new(), - state: HashMap::new(), - touched_order: Vec::new(), - deleted_paths: HashSet::new(), - } - } - - fn apply(&mut self, operation: Operation) -> Result<(), PatchError> { - self.touch(&operation.path)?; - let current = self.current(&operation.path)?.clone(); - match operation.kind { - OperationKind::Add => { - if current.kind != EntryKind::Missing - || self.deleted_paths.contains(&operation.path) - { - return Err(PatchError::Conflict(format!( - "operation at line {} creates an existing or previously deleted file {}", - operation.line, - operation.path.display() - ))); - } - let data = parse_create_diff(&operation.diff)?; - self.state - .insert(operation.path.clone(), VirtualFile::regular(data, 0o644)); - } - OperationKind::Update => { - if current.kind != EntryKind::Regular { - return Err(PatchError::Conflict(format!( - "operation at line {} updates a non-existent or non-regular file {}", - operation.line, - operation.path.display() - ))); - } - let data = if is_binary_data(¤t.data) { - if operation.diff.is_empty() { - current.data.clone() - } else { - return Err(PatchError::Invalid( - "binary file updates must not contain a text diff".to_string(), - )); - } - } else { - apply_update_diff(¤t.data, &operation.diff)? - }; - self.state.insert( - operation.path.clone(), - VirtualFile::regular(data, current.mode), - ); - if let Some(destination) = operation.move_to { - self.move_file(&operation.path, &destination, operation.line)?; - } - } - OperationKind::Delete => { - if !matches!(current.kind, EntryKind::Regular | EntryKind::Symlink) { - return Err(PatchError::Conflict(format!( - "operation at line {} deletes a missing or non-regular file {}", - operation.line, - operation.path.display() - ))); - } - self.state - .insert(operation.path.clone(), VirtualFile::missing()); - self.deleted_paths.insert(operation.path); - } - } - Ok(()) - } - - fn touch(&mut self, path: &Path) -> Result<(), PatchError> { - if self.original.contains_key(path) { - return Ok(()); - } - let snapshot = read_snapshot(path)?; - let state = VirtualFile { - kind: snapshot.kind, - data: snapshot.data.clone(), - mode: snapshot.mode, - }; - self.original.insert(path.to_path_buf(), snapshot); - self.state.insert(path.to_path_buf(), state); - self.touched_order.push(path.to_path_buf()); - Ok(()) - } - - fn current(&mut self, path: &Path) -> Result<&VirtualFile, PatchError> { - self.touch(path)?; - self.state - .get(path) - .ok_or_else(|| PatchError::Invalid(format!("missing state for {}", path.display()))) - } - - fn move_file( - &mut self, - source: &Path, - destination: &Path, - line: usize, - ) -> Result<(), PatchError> { - if source == destination { - return Ok(()); - } - self.touch(destination)?; - let destination_state = self.current(destination)?.clone(); - if destination_state.kind != EntryKind::Missing || self.deleted_paths.contains(destination) - { - return Err(PatchError::Conflict(format!( - "operation at line {line} moves {} onto an occupied path {}", - source.display(), - destination.display() - ))); - } - let source_state = self.current(source)?.clone(); - self.state.insert(destination.to_path_buf(), source_state); - self.state - .insert(source.to_path_buf(), VirtualFile::missing()); - self.deleted_paths.insert(source.to_path_buf()); - Ok(()) - } - - fn prepare_results(&self) -> Result, PatchError> { - let mut results = Vec::new(); - for path in &self.touched_order { - let before = self.original.get(path).ok_or_else(|| { - PatchError::Invalid(format!("missing original snapshot for {}", path.display())) - })?; - let after = self.state.get(path).ok_or_else(|| { - PatchError::Invalid(format!("missing final state for {}", path.display())) - })?; - if before.kind == after.kind && before.data == after.data { - continue; - } - if is_binary_snapshot(before) || is_binary_virtual(after) { - results.push(FileResult { - path: relative_display(&self.cwd, path), - diff: String::new(), - added_lines: 0, - removed_lines: 0, - }); - continue; - } - let old_text = text_for_diff(before)?; - let new_text = text_for_diff_virtual(after)?; - let relative = relative_display(&self.cwd, path); - let old_header = if before.kind == EntryKind::Missing { - "/dev/null".to_string() - } else { - format!("a/{relative}") - }; - let new_header = if after.kind == EntryKind::Missing { - "/dev/null".to_string() - } else { - format!("b/{relative}") - }; - let (diff, added_lines, removed_lines) = - unified_diff(&old_text, &new_text, &old_header, &new_header); - results.push(FileResult { - path: relative, - diff, - added_lines, - removed_lines, - }); - } - Ok(results) - } - - fn lock_paths(&self) -> Result, PatchError> { - Ok(self - .changes()? - .into_iter() - .map(|change| change.path) - .collect()) - } - - fn commit(&self) -> Result<(), PatchError> { - let changes = self.changes()?; - if changes.is_empty() { - return Ok(()); - } - - for change in &changes { - let expected = self.original.get(&change.path).ok_or_else(|| { - PatchError::Invalid(format!( - "missing original snapshot for {}", - change.path.display() - )) - })?; - if &read_snapshot(&change.path)? != expected { - return Err(PatchError::Conflict(format!( - "file changed while the patch was being prepared: {}", - change.path.display() - ))); - } - } - - let transaction_id = format!( - ".agenty-apply-patch-{}-{}", - std::process::id(), - TRANSACTION_COUNTER.fetch_add(1, Ordering::Relaxed) - ); - let mut staged = Vec::new(); - let mut installed = Vec::new(); - let mut backups = Vec::new(); - let mut created_dirs = Vec::new(); - - let result = (|| -> Result<(), PatchError> { - for change in &changes { - if path_exists(&change.path)? { - let backup = backup_path(&change.path, &transaction_id)?; - fs::rename(&change.path, &backup)?; - backups.push((change.path.clone(), backup)); - } - } - - for change in &changes { - if let Some(file) = &change.after { - let parent = change.path.parent().unwrap_or(&self.cwd); - created_dirs.extend(create_missing_dirs(parent)?); - let temp = parent.join(format!(".{transaction_id}.stage")); - let temp = unique_path(&temp)?; - let mut output = OpenOptions::new() - .write(true) - .create_new(true) - .open(&temp)?; - output.write_all(&file.data)?; - output.flush()?; - set_mode(&output, file.mode)?; - output.sync_all()?; - staged.push((change.path.clone(), temp)); - } - } - - for (path, temp) in &staged { - fs::rename(temp, path)?; - installed.push(path.clone()); - } - sync_parent_directories(&changes)?; - Ok(()) - })(); - - if result.is_err() { - for path in installed.iter().rev() { - let _ = fs::remove_file(path); - } - for (path, backup) in backups.iter().rev() { - let _ = fs::rename(backup, path); - } - for temp in staged.iter().map(|(_, temp)| temp) { - let _ = fs::remove_file(temp); - } - for directory in &created_dirs { - let _ = fs::remove_dir(directory); - } - let _ = sync_parent_directories(&changes); - } else { - for (_, backup) in backups { - let _ = fs::remove_file(backup); - } - } - result - } - - fn changes(&self) -> Result, PatchError> { - let mut changes = Vec::new(); - for path in &self.touched_order { - let before = self.original.get(path).ok_or_else(|| { - PatchError::Invalid(format!("missing original snapshot for {}", path.display())) - })?; - let after = self.state.get(path).ok_or_else(|| { - PatchError::Invalid(format!("missing final state for {}", path.display())) - })?; - let before_exists = before.kind != EntryKind::Missing; - let after_file = if after.kind == EntryKind::Regular { - Some(after.clone()) - } else { - None - }; - if before_exists && after.kind == EntryKind::Missing { - changes.push(Change { - path: path.clone(), - after: None, - }); - } else if (!before_exists && after_file.is_some()) - || (before.kind == EntryKind::Regular - && after.kind == EntryKind::Regular - && (before.data != after.data || before.mode != after.mode)) - { - changes.push(Change { - path: path.clone(), - after: after_file, - }); - } else if before.kind != after.kind { - return Err(PatchError::Conflict(format!( - "unsupported final state transition for {}", - path.display() - ))); - } - } - Ok(changes) - } -} - -struct Change { - path: PathBuf, - after: Option, -} - -#[cfg(unix)] -fn sync_parent_directories(changes: &[Change]) -> Result<(), PatchError> { - let mut parents = HashSet::new(); - for change in changes { - if let Some(parent) = change.path.parent() { - parents.insert(parent); - } - } - for parent in parents { - File::open(parent)?.sync_all()?; - } - Ok(()) -} - -#[cfg(not(unix))] -fn sync_parent_directories(_changes: &[Change]) -> Result<(), PatchError> { - Ok(()) -} - fn read_snapshot(path: &Path) -> Result { let metadata = match fs::symlink_metadata(path) { Ok(metadata) => metadata, @@ -1413,68 +945,35 @@ mod tests { } #[test] - fn creates_pid_lock_for_each_transaction_target_and_removes_it() { - let cwd = temp_dir("file-lock"); - let target = cwd.join("notes.txt"); - let lock_directory = test_data_dir().join("locks"); - let locks = - FileLocks::acquire(&lock_directory, std::slice::from_ref(&target), 4242).unwrap(); - let lock = lock_path(&lock_directory, &target); - let owner: LockOwner = serde_json::from_slice(&fs::read(&lock).unwrap()).unwrap(); - assert_eq!(owner.pid, 4242); - assert_eq!(owner.path, target.display().to_string()); - locks.release().unwrap(); - assert!(!lock.exists()); - } - - #[test] - fn rejects_a_target_with_an_existing_pid_lock_without_writing() { + fn reuses_an_existing_lock_file_as_advisory_lock_storage() { let cwd = temp_dir("locked-target"); let target = cwd.join("notes.txt"); fs::write(&target, "one\n").unwrap(); let lock_directory = test_data_dir().join("locks"); fs::create_dir_all(&lock_directory).unwrap(); let lock = lock_path(&lock_directory, &target); - fs::write( - &lock, - serde_json::to_vec(&LockOwner { - pid: 9876, - path: target.display().to_string(), - }) - .unwrap(), - ) - .unwrap(); + fs::write(&lock, b"stale metadata from an older apply_patch\n").unwrap(); let patch = "*** Begin Patch\n*** Update File: notes.txt\n@@\n-one\n+two\n*** End Patch"; - let error = apply_patch(&cwd, patch).unwrap_err().to_string(); - assert!(error.contains("locked by apply_patch pid 9876")); - assert_eq!(fs::read_to_string(target).unwrap(), "one\n"); + apply_patch(&cwd, patch).unwrap(); + assert_eq!(fs::read_to_string(target).unwrap(), "two\n"); + assert!(lock.exists()); } #[test] - fn move_requires_locks_for_source_and_destination() { + fn move_creates_locks_for_source_and_destination() { let cwd = temp_dir("locked-move"); let source = cwd.join("old.txt"); let destination = cwd.join("new.txt"); fs::write(&source, "one\n").unwrap(); let lock_directory = test_data_dir().join("locks"); - fs::create_dir_all(&lock_directory).unwrap(); - let lock = lock_path(&lock_directory, &destination); - fs::write( - &lock, - serde_json::to_vec(&LockOwner { - pid: 1111, - path: destination.display().to_string(), - }) - .unwrap(), - ) - .unwrap(); let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** End Patch"; - let error = apply_patch(&cwd, patch).unwrap_err().to_string(); - assert!(error.contains("new.txt")); - assert_eq!(fs::read_to_string(source).unwrap(), "one\n"); - assert!(!destination.exists()); + apply_patch(&cwd, patch).unwrap(); + assert_eq!(fs::read_to_string(&destination).unwrap(), "two\n"); + assert!(!source.exists()); + assert!(lock_path(&lock_directory, &source).exists()); + assert!(lock_path(&lock_directory, &destination).exists()); } #[test] diff --git a/packages/patch-applier/src/transaction.rs b/packages/patch-applier/src/transaction.rs new file mode 100644 index 0000000..f07efbc --- /dev/null +++ b/packages/patch-applier/src/transaction.rs @@ -0,0 +1,350 @@ +use std::collections::{HashMap, HashSet}; +use std::fs::{self, File, OpenOptions}; +use std::io::Write; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; + +use super::{ + apply_update_diff, backup_path, create_missing_dirs, is_binary_data, is_binary_snapshot, + is_binary_virtual, parse_create_diff, path_exists, read_snapshot, relative_display, set_mode, + text_for_diff, text_for_diff_virtual, unified_diff, unique_path, EntryKind, FileResult, + FileSnapshot, Operation, OperationKind, PatchError, VirtualFile, +}; + +static TRANSACTION_COUNTER: AtomicU64 = AtomicU64::new(0); + +pub(super) struct Transaction { + cwd: PathBuf, + original: HashMap, + state: HashMap, + touched_order: Vec, + deleted_paths: HashSet, +} + +impl Transaction { + pub(super) fn new(cwd: PathBuf) -> Self { + Self { + cwd, + original: HashMap::new(), + state: HashMap::new(), + touched_order: Vec::new(), + deleted_paths: HashSet::new(), + } + } + + pub(super) fn apply(&mut self, operation: Operation) -> Result<(), PatchError> { + self.touch(&operation.path)?; + let current = self.current(&operation.path)?.clone(); + match operation.kind { + OperationKind::Add => { + if current.kind != EntryKind::Missing + || self.deleted_paths.contains(&operation.path) + { + return Err(PatchError::Conflict(format!( + "operation at line {} creates an existing or previously deleted file {}", + operation.line, + operation.path.display() + ))); + } + let data = parse_create_diff(&operation.diff)?; + self.state + .insert(operation.path.clone(), VirtualFile::regular(data, 0o644)); + } + OperationKind::Update => { + if current.kind != EntryKind::Regular { + return Err(PatchError::Conflict(format!( + "operation at line {} updates a non-existent or non-regular file {}", + operation.line, + operation.path.display() + ))); + } + let data = if is_binary_data(¤t.data) { + if operation.diff.is_empty() { + current.data.clone() + } else { + return Err(PatchError::Invalid( + "binary file updates must not contain a text diff".to_string(), + )); + } + } else { + apply_update_diff(¤t.data, &operation.diff)? + }; + self.state.insert( + operation.path.clone(), + VirtualFile::regular(data, current.mode), + ); + if let Some(destination) = operation.move_to { + self.move_file(&operation.path, &destination, operation.line)?; + } + } + OperationKind::Delete => { + if !matches!(current.kind, EntryKind::Regular | EntryKind::Symlink) { + return Err(PatchError::Conflict(format!( + "operation at line {} deletes a missing or non-regular file {}", + operation.line, + operation.path.display() + ))); + } + self.state + .insert(operation.path.clone(), VirtualFile::missing()); + self.deleted_paths.insert(operation.path); + } + } + Ok(()) + } + + fn touch(&mut self, path: &Path) -> Result<(), PatchError> { + if self.original.contains_key(path) { + return Ok(()); + } + let snapshot = read_snapshot(path)?; + let state = VirtualFile { + kind: snapshot.kind, + data: snapshot.data.clone(), + mode: snapshot.mode, + }; + self.original.insert(path.to_path_buf(), snapshot); + self.state.insert(path.to_path_buf(), state); + self.touched_order.push(path.to_path_buf()); + Ok(()) + } + + fn current(&mut self, path: &Path) -> Result<&VirtualFile, PatchError> { + self.touch(path)?; + self.state + .get(path) + .ok_or_else(|| PatchError::Invalid(format!("missing state for {}", path.display()))) + } + + fn move_file( + &mut self, + source: &Path, + destination: &Path, + line: usize, + ) -> Result<(), PatchError> { + if source == destination { + return Ok(()); + } + self.touch(destination)?; + let destination_state = self.current(destination)?.clone(); + if destination_state.kind != EntryKind::Missing || self.deleted_paths.contains(destination) + { + return Err(PatchError::Conflict(format!( + "operation at line {line} moves {} onto an occupied path {}", + source.display(), + destination.display() + ))); + } + let source_state = self.current(source)?.clone(); + self.state.insert(destination.to_path_buf(), source_state); + self.state + .insert(source.to_path_buf(), VirtualFile::missing()); + self.deleted_paths.insert(source.to_path_buf()); + Ok(()) + } + + pub(super) fn prepare_results(&self) -> Result, PatchError> { + let mut results = Vec::new(); + for path in &self.touched_order { + let before = self.original.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing original snapshot for {}", path.display())) + })?; + let after = self.state.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing final state for {}", path.display())) + })?; + if before.kind == after.kind && before.data == after.data { + continue; + } + if is_binary_snapshot(before) || is_binary_virtual(after) { + results.push(FileResult { + path: relative_display(&self.cwd, path), + diff: String::new(), + added_lines: 0, + removed_lines: 0, + }); + continue; + } + let old_text = text_for_diff(before)?; + let new_text = text_for_diff_virtual(after)?; + let relative = relative_display(&self.cwd, path); + let old_header = if before.kind == EntryKind::Missing { + "/dev/null".to_string() + } else { + format!("a/{relative}") + }; + let new_header = if after.kind == EntryKind::Missing { + "/dev/null".to_string() + } else { + format!("b/{relative}") + }; + let (diff, added_lines, removed_lines) = + unified_diff(&old_text, &new_text, &old_header, &new_header); + results.push(FileResult { + path: relative, + diff, + added_lines, + removed_lines, + }); + } + Ok(results) + } + + pub(super) fn lock_paths(&self) -> Result, PatchError> { + Ok(self + .changes()? + .into_iter() + .map(|change| change.path) + .collect()) + } + + pub(super) fn commit(&self) -> Result<(), PatchError> { + let changes = self.changes()?; + if changes.is_empty() { + return Ok(()); + } + + for change in &changes { + let expected = self.original.get(&change.path).ok_or_else(|| { + PatchError::Invalid(format!( + "missing original snapshot for {}", + change.path.display() + )) + })?; + if &read_snapshot(&change.path)? != expected { + return Err(PatchError::Conflict(format!( + "file changed while the patch was being prepared: {}", + change.path.display() + ))); + } + } + + let transaction_id = format!( + ".agenty-apply-patch-{}-{}", + std::process::id(), + TRANSACTION_COUNTER.fetch_add(1, Ordering::Relaxed) + ); + let mut staged = Vec::new(); + let mut installed = Vec::new(); + let mut backups = Vec::new(); + let mut created_dirs = Vec::new(); + + let result = (|| -> Result<(), PatchError> { + for change in &changes { + if path_exists(&change.path)? { + let backup = backup_path(&change.path, &transaction_id)?; + fs::rename(&change.path, &backup)?; + backups.push((change.path.clone(), backup)); + } + } + + for change in &changes { + if let Some(file) = &change.after { + let parent = change.path.parent().unwrap_or(&self.cwd); + created_dirs.extend(create_missing_dirs(parent)?); + let temp = parent.join(format!(".{transaction_id}.stage")); + let temp = unique_path(&temp)?; + let mut output = OpenOptions::new() + .write(true) + .create_new(true) + .open(&temp)?; + output.write_all(&file.data)?; + output.flush()?; + set_mode(&output, file.mode)?; + output.sync_all()?; + staged.push((change.path.clone(), temp)); + } + } + + for (path, temp) in &staged { + fs::rename(temp, path)?; + installed.push(path.clone()); + } + sync_parent_directories(&changes)?; + Ok(()) + })(); + + if result.is_err() { + for path in installed.iter().rev() { + let _ = fs::remove_file(path); + } + for (path, backup) in backups.iter().rev() { + let _ = fs::rename(backup, path); + } + for temp in staged.iter().map(|(_, temp)| temp) { + let _ = fs::remove_file(temp); + } + for directory in &created_dirs { + let _ = fs::remove_dir(directory); + } + let _ = sync_parent_directories(&changes); + } else { + for (_, backup) in backups { + let _ = fs::remove_file(backup); + } + } + result + } + + fn changes(&self) -> Result, PatchError> { + let mut changes = Vec::new(); + for path in &self.touched_order { + let before = self.original.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing original snapshot for {}", path.display())) + })?; + let after = self.state.get(path).ok_or_else(|| { + PatchError::Invalid(format!("missing final state for {}", path.display())) + })?; + let before_exists = before.kind != EntryKind::Missing; + let after_file = if after.kind == EntryKind::Regular { + Some(after.clone()) + } else { + None + }; + if before_exists && after.kind == EntryKind::Missing { + changes.push(Change { + path: path.clone(), + after: None, + }); + } else if (!before_exists && after_file.is_some()) + || (before.kind == EntryKind::Regular + && after.kind == EntryKind::Regular + && (before.data != after.data || before.mode != after.mode)) + { + changes.push(Change { + path: path.clone(), + after: after_file, + }); + } else if before.kind != after.kind { + return Err(PatchError::Conflict(format!( + "unsupported final state transition for {}", + path.display() + ))); + } + } + Ok(changes) + } +} + +struct Change { + path: PathBuf, + after: Option, +} + +#[cfg(unix)] +fn sync_parent_directories(changes: &[Change]) -> Result<(), PatchError> { + let mut parents = HashSet::new(); + for change in changes { + if let Some(parent) = change.path.parent() { + parents.insert(parent); + } + } + for parent in parents { + File::open(parent)?.sync_all()?; + } + Ok(()) +} + +#[cfg(not(unix))] +fn sync_parent_directories(_changes: &[Change]) -> Result<(), PatchError> { + Ok(()) +} From b695bd5ff7f36d202113723f7d2be24f087c99a1 Mon Sep 17 00:00:00 2001 From: masteryyh Date: Fri, 28 Aug 2026 11:08:49 +0800 Subject: [PATCH 11/12] fix: preserve file metadata when moving Signed-off-by: masteryyh --- packages/patch-applier/src/lib.rs | 47 ++++++ packages/patch-applier/src/transaction.rs | 176 +++++++++++++++++++--- 2 files changed, 200 insertions(+), 23 deletions(-) diff --git a/packages/patch-applier/src/lib.rs b/packages/patch-applier/src/lib.rs index 2b0a89b..5d4a0ac 100644 --- a/packages/patch-applier/src/lib.rs +++ b/packages/patch-applier/src/lib.rs @@ -91,6 +91,7 @@ struct VirtualFile { kind: EntryKind, data: Vec, mode: u32, + origin: Option, } impl VirtualFile { @@ -99,6 +100,7 @@ impl VirtualFile { kind: EntryKind::Missing, data: Vec::new(), mode: 0, + origin: None, } } @@ -107,6 +109,7 @@ impl VirtualFile { kind: EntryKind::Regular, data, mode, + origin: None, } } } @@ -855,6 +858,50 @@ mod tests { assert_eq!(fs::read_to_string(cwd.join("new.txt")).unwrap(), "three\n"); } + #[cfg(unix)] + #[test] + fn moves_file_by_reusing_the_source_inode_and_metadata() { + use std::os::unix::fs::{MetadataExt, PermissionsExt}; + + let cwd = temp_dir("move-metadata"); + let source = cwd.join("old.txt"); + let destination = cwd.join("new.txt"); + fs::write(&source, "one\n").unwrap(); + fs::set_permissions(&source, fs::Permissions::from_mode(0o751)).unwrap(); + let source_metadata = fs::metadata(&source).unwrap(); + + let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** End Patch"; + apply_patch(&cwd, patch).unwrap(); + + let destination_metadata = fs::metadata(&destination).unwrap(); + assert_eq!(destination_metadata.ino(), source_metadata.ino()); + assert_eq!(MetadataExt::mode(&destination_metadata) & 0o7777, 0o751); + assert_eq!(fs::read_to_string(destination).unwrap(), "two\n"); + } + + #[cfg(unix)] + #[test] + fn restores_a_moved_inode_when_later_staging_fails() { + use std::os::unix::fs::PermissionsExt; + + let cwd = temp_dir("move-staging-failure"); + let source = cwd.join("old.txt"); + let destination = cwd.join("new.txt"); + fs::write(&source, "one").unwrap(); + let blocked = cwd.join("blocked"); + fs::create_dir(&blocked).unwrap(); + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o555)).unwrap(); + + let patch = "*** Begin Patch\n*** Update File: old.txt\n*** Move to: new.txt\n@@\n-one\n+two\n*** Add File: blocked/child.txt\n+child\n*** End Patch"; + let error = apply_patch(&cwd, patch).unwrap_err(); + + fs::set_permissions(&blocked, fs::Permissions::from_mode(0o755)).unwrap(); + assert!(error.to_string().contains("Permission denied")); + assert_eq!(fs::read_to_string(source).unwrap(), "one"); + assert!(!destination.exists()); + assert!(!blocked.join("child.txt").exists()); + } + #[test] fn malformed_patch_has_no_side_effect() { let cwd = temp_dir("malformed"); diff --git a/packages/patch-applier/src/transaction.rs b/packages/patch-applier/src/transaction.rs index f07efbc..0542cca 100644 --- a/packages/patch-applier/src/transaction.rs +++ b/packages/patch-applier/src/transaction.rs @@ -69,10 +69,9 @@ impl Transaction { } else { apply_update_diff(¤t.data, &operation.diff)? }; - self.state.insert( - operation.path.clone(), - VirtualFile::regular(data, current.mode), - ); + let mut updated = VirtualFile::regular(data, current.mode); + updated.origin = current.origin.clone(); + self.state.insert(operation.path.clone(), updated); if let Some(destination) = operation.move_to { self.move_file(&operation.path, &destination, operation.line)?; } @@ -102,6 +101,7 @@ impl Transaction { kind: snapshot.kind, data: snapshot.data.clone(), mode: snapshot.mode, + origin: (snapshot.kind == EntryKind::Regular).then(|| path.to_path_buf()), }; self.original.insert(path.to_path_buf(), snapshot); self.state.insert(path.to_path_buf(), state); @@ -226,6 +226,8 @@ impl Transaction { let mut staged = Vec::new(); let mut installed = Vec::new(); let mut backups = Vec::new(); + let mut backup_by_path = HashMap::new(); + let mut consumed_backups = HashSet::new(); let mut created_dirs = Vec::new(); let result = (|| -> Result<(), PatchError> { @@ -233,6 +235,7 @@ impl Transaction { if path_exists(&change.path)? { let backup = backup_path(&change.path, &transaction_id)?; fs::rename(&change.path, &backup)?; + backup_by_path.insert(change.path.clone(), backup.clone()); backups.push((change.path.clone(), backup)); } } @@ -241,38 +244,112 @@ impl Transaction { if let Some(file) = &change.after { let parent = change.path.parent().unwrap_or(&self.cwd); created_dirs.extend(create_missing_dirs(parent)?); - let temp = parent.join(format!(".{transaction_id}.stage")); - let temp = unique_path(&temp)?; - let mut output = OpenOptions::new() - .write(true) - .create_new(true) - .open(&temp)?; - output.write_all(&file.data)?; - output.flush()?; - set_mode(&output, file.mode)?; - output.sync_all()?; - staged.push((change.path.clone(), temp)); + // Reuse the source inode for moves so ACLs, xattrs, ownership, + // and other filesystem metadata follow the moved file. + if let Some(origin) = file.origin.as_ref().filter(|origin| { + origin.as_path() != change.path.as_path() + && backup_by_path.contains_key(*origin) + }) { + if !consumed_backups.insert(origin.clone()) { + return Err(PatchError::Conflict(format!( + "multiple files reuse the backup for {}", + origin.display() + ))); + } + let backup = backup_by_path.get(origin).ok_or_else(|| { + PatchError::Invalid(format!( + "missing backup for metadata source {}", + origin.display() + )) + })?; + let original = self.original.get(origin).ok_or_else(|| { + PatchError::Invalid(format!( + "missing original snapshot for metadata source {}", + origin.display() + )) + })?; + let rewritten = file.data != original.data || file.mode != original.mode; + staged.push(StagedFile::Backup { + path: change.path.clone(), + origin: origin.clone(), + backup: backup.clone(), + rewritten, + }); + if rewritten { + rewrite_file(backup, file)?; + } + } else { + let temp = parent.join(format!(".{transaction_id}.stage")); + let temp = unique_path(&temp)?; + staged.push(StagedFile::Temporary { + path: change.path.clone(), + temp: temp.clone(), + }); + let mut output = OpenOptions::new() + .write(true) + .create_new(true) + .open(&temp)?; + output.write_all(&file.data)?; + output.flush()?; + set_mode(&output, file.mode)?; + output.sync_all()?; + } } } - for (path, temp) in &staged { - fs::rename(temp, path)?; - installed.push(path.clone()); + for stage in &staged { + match stage { + StagedFile::Temporary { path, temp } => { + fs::rename(temp, path)?; + installed.push(InstalledFile::Temporary { path: path.clone() }); + } + StagedFile::Backup { path, backup, .. } => { + fs::rename(backup, path)?; + installed.push(InstalledFile::Backup { + path: path.clone(), + backup: backup.clone(), + }); + } + } } sync_parent_directories(&changes)?; Ok(()) })(); if result.is_err() { - for path in installed.iter().rev() { - let _ = fs::remove_file(path); + for installed_file in installed.iter().rev() { + match installed_file { + InstalledFile::Temporary { path } => { + let _ = fs::remove_file(path); + } + InstalledFile::Backup { path, backup } => { + let _ = fs::rename(path, backup); + } + } + } + for stage in &staged { + match stage { + StagedFile::Temporary { temp, .. } => { + let _ = fs::remove_file(temp); + } + StagedFile::Backup { + origin, + backup, + rewritten, + .. + } => { + if *rewritten { + if let Some(original) = self.original.get(origin) { + let _ = + rewrite_file_contents(backup, &original.data, original.mode); + } + } + } + } } for (path, backup) in backups.iter().rev() { let _ = fs::rename(backup, path); } - for temp in staged.iter().map(|(_, temp)| temp) { - let _ = fs::remove_file(temp); - } for directory in &created_dirs { let _ = fs::remove_dir(directory); } @@ -330,6 +407,59 @@ struct Change { after: Option, } +enum StagedFile { + Temporary { + path: PathBuf, + temp: PathBuf, + }, + Backup { + path: PathBuf, + origin: PathBuf, + backup: PathBuf, + rewritten: bool, + }, +} + +enum InstalledFile { + Temporary { path: PathBuf }, + Backup { path: PathBuf, backup: PathBuf }, +} + +fn rewrite_file(path: &Path, file: &VirtualFile) -> Result<(), PatchError> { + rewrite_file_contents(path, &file.data, file.mode) +} + +fn rewrite_file_contents(path: &Path, data: &[u8], mode: u32) -> Result<(), PatchError> { + let mut output = match OpenOptions::new().write(true).open(path) { + Ok(output) => output, + Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { + make_writable(path, mode)?; + OpenOptions::new().write(true).open(path)? + } + Err(error) => return Err(error.into()), + }; + output.set_len(0)?; + output.write_all(data)?; + output.flush()?; + set_mode(&output, mode)?; + output.sync_all()?; + Ok(()) +} + +#[cfg(unix)] +fn make_writable(path: &Path, mode: u32) -> Result<(), PatchError> { + use std::os::unix::fs::PermissionsExt; + if mode & 0o200 == 0 { + fs::set_permissions(path, fs::Permissions::from_mode((mode | 0o600) & 0o7777))?; + } + Ok(()) +} + +#[cfg(not(unix))] +fn make_writable(_path: &Path, _mode: u32) -> Result<(), PatchError> { + Ok(()) +} + #[cfg(unix)] fn sync_parent_directories(changes: &[Change]) -> Result<(), PatchError> { let mut parents = HashSet::new(); From ba57b6b8de3f78eba0324497ec95a889df73ee9f Mon Sep 17 00:00:00 2001 From: masteryyh Date: Fri, 28 Aug 2026 11:09:16 +0800 Subject: [PATCH 12/12] fix: filter non text models when fetching model list Signed-off-by: masteryyh --- .../pkg/infra/modelcatalog/lister.go | 8 ++++- .../pkg/infra/modelcatalog/lister_test.go | 29 +++++++++++++------ 2 files changed, 27 insertions(+), 10 deletions(-) diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister.go b/packages/agenty-core/pkg/infra/modelcatalog/lister.go index a5138fe..7cdc811 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister.go @@ -210,7 +210,8 @@ type openRouterModel struct { } type openRouterArchitecture struct { - InputModalities []string `json:"input_modalities"` + InputModalities []string `json:"input_modalities"` + OutputModalities []string `json:"output_modalities"` } type openRouterTopProvider struct { @@ -228,8 +229,13 @@ func (l *Lister) listOpenAICompatible(ctx context.Context, provider catalog.Prov return nil, err } + isOpenRouter := provider.Code.String() == "openrouter" models := make([]catalog.AvailableModel, 0, len(response.Data)) for index, item := range response.Data { + if isOpenRouter && !slices.Contains(item.Architecture.OutputModalities, "text") { + continue + } + contextWindow := item.ContextLength if contextWindow <= 0 { contextWindow = item.TopProvider.ContextLength diff --git a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go index 1ad134c..1467798 100644 --- a/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go +++ b/packages/agenty-core/pkg/infra/modelcatalog/lister_test.go @@ -131,11 +131,25 @@ func TestListerOpenRouterFieldsAndReasoning(t *testing.T) { "id": "openai/gpt-test", "name": "", "context_length": 128000, - "architecture": map[string]any{"input_modalities": []string{"text", "image"}}, - "top_provider": map[string]any{"max_completion_tokens": 32768}, - "reasoning": map[string]any{"supported_efforts": []string{"high", "minimal", "low"}}, + "architecture": map[string]any{ + "input_modalities": []string{"text", "image"}, + "output_modalities": []string{"text"}, + }, + "top_provider": map[string]any{"max_completion_tokens": 32768}, + "reasoning": map[string]any{"supported_efforts": []string{"high", "minimal", "low"}}, + }, + { + "id": "embed-model", + "architecture": map[string]any{ + "output_modalities": []string{"embeddings"}, + }, + }, + { + "id": "image-model", + "architecture": map[string]any{ + "output_modalities": []string{"image"}, + }, }, - {"id": "plain-model"}, }, } if err := json.NewEncoder(w).Encode(payload); err != nil { @@ -153,7 +167,7 @@ func TestListerOpenRouterFieldsAndReasoning(t *testing.T) { if err != nil { t.Fatalf("List: %v", err) } - if len(models) != 2 { + if len(models) != 1 { t.Fatalf("models = %#v", models) } if models[0].Name != "openai/gpt-test" || models[0].ContextWindow != 128000 || models[0].MaxOutputTokens != 32768 || !models[0].MultiModal { @@ -162,14 +176,11 @@ func TestListerOpenRouterFieldsAndReasoning(t *testing.T) { if !reflect.DeepEqual(models[0].ReasoningEfforts, []shared.ReasoningEffort{shared.ReasoningLow, shared.ReasoningHigh}) { t.Errorf("first reasoning efforts = %#v", models[0].ReasoningEfforts) } - if models[1].ContextWindow != catalog.DefaultAvailableModelContextWindow || models[1].MaxOutputTokens != catalog.DefaultAvailableModelMaxOutputTokens || len(models[1].ReasoningEfforts) != 0 { - t.Errorf("second model = %#v", models[1]) - } } func TestListerOpenRouterNullReasoningMeansAllStandardEfforts(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, _ = w.Write([]byte(`{"data":[{"id":"reasoning-model","reasoning":{"supported_efforts":null}}]}`)) + _, _ = w.Write([]byte(`{"data":[{"id":"reasoning-model","architecture":{"output_modalities":["text"]},"reasoning":{"supported_efforts":null}}]}`)) })) defer server.Close()