From d64a31ebdb598aab806e24e18a9a422f71ec5e0b Mon Sep 17 00:00:00 2001 From: Goon Date: Sun, 24 May 2026 10:42:26 +0700 Subject: [PATCH] fix(usage): enforce caps on auxiliary llm calls --- cmd/gateway.go | 11 +- cmd/gateway_consumer.go | 4 +- cmd/gateway_consumer_deps.go | 2 + cmd/gateway_consumer_normal.go | 6 +- cmd/gateway_deps.go | 8 +- cmd/gateway_hooks.go | 6 +- cmd/gateway_http_handlers.go | 7 +- cmd/gateway_http_wiring.go | 4 +- cmd/gateway_lifecycle.go | 2 +- cmd/gateway_managed.go | 7 +- cmd/gateway_methods.go | 4 +- internal/agent/intent_classify.go | 12 +- internal/agent/title_generate.go | 16 ++- internal/channels/history_compaction.go | 13 +- internal/channels/instance_loader.go | 11 +- internal/consolidation/dreaming_worker.go | 16 ++- internal/consolidation/episodic_worker.go | 18 ++- internal/consolidation/workers.go | 4 + internal/gateway/methods/chat.go | 11 +- internal/hooks/handlers/prompt.go | 34 ++++- internal/http/knowledge_graph.go | 10 +- internal/http/knowledge_graph_handlers.go | 17 ++- internal/http/pending_messages.go | 8 +- internal/http/provider_verify.go | 11 +- internal/http/providers.go | 6 + internal/http/summoner.go | 6 +- internal/http/summoner_regenerate.go | 10 +- internal/http/usage_caps.go | 53 ++++---- internal/i18n/catalog_en.go | 23 +++- internal/i18n/catalog_vi.go | 23 +++- internal/i18n/catalog_zh.go | 23 +++- internal/i18n/keys.go | 111 +++++++++------- internal/knowledgegraph/extractor.go | 20 ++- internal/store/pg/usage_caps.go | 9 +- internal/usage/caps/chat_call.go | 150 ++++++++++++++++++++++ internal/usage/caps/service.go | 7 + internal/usage/caps/service_test.go | 111 +++++++++++++++- internal/vault/enrich_classify.go | 9 ++ internal/vault/enrich_worker.go | 55 +++++--- 39 files changed, 701 insertions(+), 157 deletions(-) create mode 100644 internal/usage/caps/chat_call.go diff --git a/cmd/gateway.go b/cmd/gateway.go index 31428735..5a2e6006 100644 --- a/cmd/gateway.go +++ b/cmd/gateway.go @@ -223,6 +223,7 @@ func runGateway() { } } setupMemoryEmbeddings(pgStores, providerRegistry) + usageCapSvc := usagecaps.NewService(pgStores.UsageCaps, pgStores.Providers) // Resolve background provider for consolidation + vault enrichment. // Fallback: background.provider → agent.default_provider → first registered provider. @@ -234,6 +235,7 @@ func runGateway() { var kgExtractor *kg.Extractor if pgStores.KnowledgeGraph != nil { kgExtractor = kg.NewExtractor(bgProvider, bgModel, 0) + kgExtractor.SetUsageCapService(usageCapSvc) } cleanupConsolidation := consolidation.Register(consolidation.ConsolidationDeps{ EpisodicStore: pgStores.Episodic, @@ -245,6 +247,7 @@ func runGateway() { Registry: providerRegistry, Extractor: kgExtractor, AlertDeps: bgalert.AlertDeps{SystemConfigs: pgStores.SystemConfigs, MsgBus: msgBus}, + UsageCaps: usageCapSvc, AgentStore: pgStores.Agents, }) defer cleanupConsolidation() @@ -267,6 +270,7 @@ func runGateway() { MsgBus: msgBus, TeamStore: pgStores.Teams, AlertDeps: bgalert.AlertDeps{SystemConfigs: pgStores.SystemConfigs, MsgBus: msgBus}, + UsageCaps: usageCapSvc, }) enrichProgress = ep enrichWorker = ew @@ -283,7 +287,6 @@ func runGateway() { slog.Info("bootstrap: capabilities backfill complete", "agents", count) } - usageCapSvc := usagecaps.NewService(pgStores.UsageCaps, pgStores.Providers) if readImage, ok := toolsReg.Get("read_image"); ok { if t, ok := readImage.(*tools.ReadImageTool); ok { t.SetUsageCapService(usageCapSvc) @@ -364,6 +367,7 @@ func runGateway() { workspace: workspace, dataDir: dataDir, domainBus: domainBus, + usageCapSvc: usageCapSvc, audioMgr: audioMgr, } @@ -376,7 +380,7 @@ func runGateway() { httpapi.InitGatewayNoAuthFallbackAllowed(config.GatewayNoAuthFallbackAllowed(cfg.Gateway)) exportTokenStore := httpapi.InitExportTokenStore() defer exportTokenStore.Stop() - agentsH, skillsH, tracesH, mcpH, channelInstancesH, providersH, builtinToolsH, pendingMessagesH, teamEventsH, secureCLIH, secureCLIGrantH, mcpUserCredsH := wireHTTP(pgStores, cfg.Agents.Defaults.Workspace, dataDir, bundledSkillsDir, msgBus, toolsReg, providerRegistry, modelReg, permPE.IsOwner, gatewayAddr, mcpToolLister) + agentsH, skillsH, tracesH, mcpH, channelInstancesH, providersH, builtinToolsH, pendingMessagesH, teamEventsH, secureCLIH, secureCLIGrantH, mcpUserCredsH := wireHTTP(pgStores, cfg.Agents.Defaults.Workspace, dataDir, bundledSkillsDir, msgBus, toolsReg, providerRegistry, modelReg, permPE.IsOwner, gatewayAddr, mcpToolLister, usageCapSvc) // Wire dependencies for system prompt preview parity. if agentsH != nil { @@ -433,7 +437,7 @@ func runGateway() { // Register all RPC methods server.SetLogTee(logTee) server.SetRuntimeLogsHandler(httpapi.NewRuntimeLogsHandler(logTee)) - pairingMethods, heartbeatMethods, chatMethods, cfgPermsMethods := registerAllMethods(server, agentRouter, pgStores.Sessions, pgStores.Cron, pgStores.Pairing, cfg, cfgPath, workspace, dataDir, msgBus, execApprovalMgr, pgStores.Agents, pgStores.Skills, pgStores.ConfigSecrets, pgStores.Teams, contextFileInterceptor, logTee, pgStores.Heartbeats, pgStores.ConfigPermissions, pgStores.SystemConfigs, pgStores.Tenants, pgStores.SkillTenantCfgs, audioMgr) + pairingMethods, heartbeatMethods, chatMethods, cfgPermsMethods := registerAllMethods(server, agentRouter, pgStores.Sessions, pgStores.Cron, pgStores.Pairing, cfg, cfgPath, workspace, dataDir, msgBus, execApprovalMgr, pgStores.Agents, pgStores.Skills, pgStores.ConfigSecrets, pgStores.Teams, contextFileInterceptor, logTee, pgStores.Heartbeats, pgStores.ConfigPermissions, pgStores.SystemConfigs, pgStores.Tenants, pgStores.SkillTenantCfgs, audioMgr, usageCapSvc) // Phase 3: Agent hooks RPC methods (hooks.list/create/update/delete/toggle/test/history). if hs, ok := pgStores.Hooks.(hooks.HookStore); ok && hs != nil { @@ -519,6 +523,7 @@ func runGateway() { instanceLoader = channels.NewInstanceLoader(pgStores.ChannelInstances, pgStores.Agents, channelMgr, msgBus, pgStores.Pairing) instanceLoader.SetProviderRegistry(providerRegistry) instanceLoader.SetPendingCompactionConfig(cfg.Channels.PendingCompaction) + instanceLoader.SetUsageCapService(usageCapSvc) instanceLoader.RegisterFactory(channels.TypeTelegram, telegram.FactoryWithStoresAndAudio(pgStores.Agents, pgStores.ConfigPermissions, pgStores.Teams, pgStores.SubagentTasks, pgStores.PendingMessages, audioMgr)) instanceLoader.RegisterFactory(channels.TypeDiscord, discord.FactoryWithStoresAndAudio(pgStores.Agents, pgStores.ConfigPermissions, pgStores.PendingMessages, audioMgr)) instanceLoader.RegisterFactory(channels.TypeFeishu, feishu.FactoryWithPendingStoreAndAudio(pgStores.PendingMessages, audioMgr)) diff --git a/cmd/gateway_consumer.go b/cmd/gateway_consumer.go index cbca58e3..3bd7e78a 100644 --- a/cmd/gateway_consumer.go +++ b/cmd/gateway_consumer.go @@ -19,6 +19,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/scheduler" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/pkg/protocol" ) @@ -26,7 +27,7 @@ import ( // and routes them through the scheduler/agent loop, then publishes the response back. // Also handles subagent announcements: routes them through the parent agent's session // (matching TS subagent-announce.ts pattern) so the agent can reformulate for the user. -func consumeInboundMessages(ctx context.Context, msgBus *bus.MessageBus, agents *agent.Router, cfg *config.Config, sched *scheduler.Scheduler, channelMgr *channels.Manager, teamStore store.TeamStore, quotaChecker *channels.QuotaChecker, sessStore store.SessionStore, agentStore store.AgentStore, contactCollector *store.ContactCollector, postTurn tools.PostTurnProcessor, subagentMgr *tools.SubagentManager) { +func consumeInboundMessages(ctx context.Context, msgBus *bus.MessageBus, agents *agent.Router, cfg *config.Config, sched *scheduler.Scheduler, channelMgr *channels.Manager, teamStore store.TeamStore, quotaChecker *channels.QuotaChecker, sessStore store.SessionStore, agentStore store.AgentStore, contactCollector *store.ContactCollector, postTurn tools.PostTurnProcessor, subagentMgr *tools.SubagentManager, usageCapSvc *usagecaps.Service) { slog.Info("inbound message consumer started") // Inbound message deduplication (matching TS src/infra/dedupe.ts + inbound-dedupe.ts). @@ -58,6 +59,7 @@ func consumeInboundMessages(ctx context.Context, msgBus *bus.MessageBus, agents QuotaChecker: quotaChecker, ContactCollector: contactCollector, SubagentMgr: subagentMgr, + UsageCaps: usageCapSvc, GetAnnounceMu: getAnnounceMu, } diff --git a/cmd/gateway_consumer_deps.go b/cmd/gateway_consumer_deps.go index faf3755c..11174f2a 100644 --- a/cmd/gateway_consumer_deps.go +++ b/cmd/gateway_consumer_deps.go @@ -10,6 +10,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/scheduler" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // ConsumerDeps bundles shared dependencies for consumer message handlers. @@ -28,6 +29,7 @@ type ConsumerDeps struct { ContactCollector *store.ContactCollector TaskRunSessions sync.Map SubagentMgr *tools.SubagentManager + UsageCaps *usagecaps.Service BgWg sync.WaitGroup GetAnnounceMu func(string) *sync.Mutex } diff --git a/cmd/gateway_consumer_normal.go b/cmd/gateway_consumer_normal.go index 0b243cf8..a719c446 100644 --- a/cmd/gateway_consumer_normal.go +++ b/cmd/gateway_consumer_normal.go @@ -333,7 +333,11 @@ func processNormalMessage( if locale == "" { locale = "en" } - intent := agent.ClassifyIntent(ctx, loop.Provider(), loop.Model(), msg.Content) + classifyCtx := ctx + if uid := loop.UUID(); uid != uuid.Nil { + classifyCtx = store.WithAgentID(classifyCtx, uid) + } + intent := agent.ClassifyIntentWithUsageCaps(classifyCtx, deps.UsageCaps, loop.Provider(), loop.Model(), msg.Content) switch intent { case agent.IntentStatusQuery: status := deps.Agents.GetActivity(sessionKey) diff --git a/cmd/gateway_deps.go b/cmd/gateway_deps.go index 0487f358..c30ca4f2 100644 --- a/cmd/gateway_deps.go +++ b/cmd/gateway_deps.go @@ -14,6 +14,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/skills" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/internal/vault" ) @@ -28,13 +29,14 @@ type gatewayDeps struct { channelMgr *channels.Manager agentRouter *agent.Router toolsReg *tools.Registry - skillsLoader *skills.Loader // optional: enables skill creation in evolution approval + skillsLoader *skills.Loader // optional: enables skill creation in evolution approval permCache *cache.PermissionCache // nil if no tenant store; closed on shutdown to stop sweep goroutines - enrichProgress *vault.EnrichProgress // nil if enrichment worker not registered - enrichWorker *vault.EnrichWorker // nil if enrichment worker not registered; for stop/enqueue + enrichProgress *vault.EnrichProgress // nil if enrichment worker not registered + enrichWorker *vault.EnrichWorker // nil if enrichment worker not registered; for stop/enqueue workspace string dataDir string domainBus eventbus.DomainEventBus + usageCapSvc *usagecaps.Service audioMgr *audio.Manager // nil if TTS not configured; used by TTSHandler ttsHandler *httpapi.TTSHandler // nil if TTS not configured; for hot-reload } diff --git a/cmd/gateway_hooks.go b/cmd/gateway_hooks.go index 26b35887..29340bc5 100644 --- a/cmd/gateway_hooks.go +++ b/cmd/gateway_hooks.go @@ -7,12 +7,13 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/config" "github.com/nextlevelbuilder/goclaw/internal/edition" "github.com/nextlevelbuilder/goclaw/internal/hooks" - hookhandlers "github.com/nextlevelbuilder/goclaw/internal/hooks/handlers" "github.com/nextlevelbuilder/goclaw/internal/hooks/budget" + hookhandlers "github.com/nextlevelbuilder/goclaw/internal/hooks/handlers" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/security" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/store/pg" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // sharedHookHandlers is populated by wireExtras so the gateway.go router @@ -27,7 +28,7 @@ var sharedHookHandlers map[hooks.HandlerType]hooks.Handler // Budget wiring (C1 fix): the PromptHandler receives a budget.Store bound // to pg.NewPGHookBudget so token spend is atomically deducted per tenant. // When the DB handle is unavailable, budget falls back to nil (Lite desktop). -func buildHookHandlers(stores *store.Stores, providerReg *providers.Registry, hooksCfg config.HooksConfig) map[hooks.HandlerType]hooks.Handler { +func buildHookHandlers(stores *store.Stores, providerReg *providers.Registry, hooksCfg config.HooksConfig, usageCapSvc *usagecaps.Service) map[hooks.HandlerType]hooks.Handler { encryptKey := os.Getenv("GOCLAW_ENCRYPTION_KEY") var budgetStore *budget.Store @@ -38,6 +39,7 @@ func buildHookHandlers(stores *store.Stores, providerReg *providers.Registry, ho promptHandler := &hookhandlers.PromptHandler{ Resolver: hookhandlers.NewRegistryResolver(providerReg, stores.SystemConfigs), Budget: budgetStore, + UsageCaps: usageCapSvc, DefaultModel: "haiku", } diff --git a/cmd/gateway_http_handlers.go b/cmd/gateway_http_handlers.go index 5ad49409..ce8586ca 100644 --- a/cmd/gateway_http_handlers.go +++ b/cmd/gateway_http_handlers.go @@ -6,10 +6,11 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // wireHTTP creates HTTP handlers (agents + skills + traces + MCP + channel instances + providers + builtin tools + pending messages). -func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir string, msgBus *bus.MessageBus, toolsReg *tools.Registry, providerReg *providers.Registry, modelReg providers.ModelRegistry, isOwner func(string) bool, gatewayAddr string, mcpToolLister httpapi.MCPToolLister) (*httpapi.AgentsHandler, *httpapi.SkillsHandler, *httpapi.TracesHandler, *httpapi.MCPHandler, *httpapi.ChannelInstancesHandler, *httpapi.ProvidersHandler, *httpapi.BuiltinToolsHandler, *httpapi.PendingMessagesHandler, *httpapi.TeamEventsHandler, *httpapi.SecureCLIHandler, *httpapi.SecureCLIGrantHandler, *httpapi.MCPUserCredentialsHandler) { +func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir string, msgBus *bus.MessageBus, toolsReg *tools.Registry, providerReg *providers.Registry, modelReg providers.ModelRegistry, isOwner func(string) bool, gatewayAddr string, mcpToolLister httpapi.MCPToolLister, usageCapSvc *usagecaps.Service) (*httpapi.AgentsHandler, *httpapi.SkillsHandler, *httpapi.TracesHandler, *httpapi.MCPHandler, *httpapi.ChannelInstancesHandler, *httpapi.ProvidersHandler, *httpapi.BuiltinToolsHandler, *httpapi.PendingMessagesHandler, *httpapi.TeamEventsHandler, *httpapi.SecureCLIHandler, *httpapi.SecureCLIGrantHandler, *httpapi.MCPUserCredentialsHandler) { var agentsH *httpapi.AgentsHandler var skillsH *httpapi.SkillsHandler var tracesH *httpapi.TracesHandler @@ -24,7 +25,7 @@ func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir if stores != nil && stores.Agents != nil { var summoner *httpapi.AgentSummoner if providerReg != nil { - summoner = httpapi.NewAgentSummoner(stores.Agents, providerReg, msgBus) + summoner = httpapi.NewAgentSummoner(stores.Agents, providerReg, msgBus, usageCapSvc) } agentsH = httpapi.NewAgentsHandler(stores.Agents, stores.Providers, providerReg, stores.DB, stores.Tracing, defaultWorkspace, msgBus, summoner, isOwner) agentsH.SetImportStores(stores.Memory, stores.KnowledgeGraph) @@ -61,6 +62,7 @@ func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir if stores != nil && stores.Providers != nil { providersH = httpapi.NewProvidersHandler(stores.Providers, stores.ConfigSecrets, providerReg, gatewayAddr) providersH.SetMessageBus(msgBus) + providersH.SetUsageCapService(usageCapSvc) if modelReg != nil { providersH.SetModelRegistry(modelReg) } @@ -90,6 +92,7 @@ func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir if stores != nil && stores.PendingMessages != nil { pendingMessagesH = httpapi.NewPendingMessagesHandler(stores.PendingMessages, stores.Agents, providerReg) + pendingMessagesH.SetUsageCapService(usageCapSvc) } if stores != nil && stores.SecureCLI != nil { diff --git a/cmd/gateway_http_wiring.go b/cmd/gateway_http_wiring.go index 76c782a3..3eac9db3 100644 --- a/cmd/gateway_http_wiring.go +++ b/cmd/gateway_http_wiring.go @@ -234,7 +234,9 @@ func (d *gatewayDeps) wireHTTPHandlersOnServer( // Knowledge graph API if d.pgStores != nil && d.pgStores.KnowledgeGraph != nil { - d.server.SetKnowledgeGraphHandler(httpapi.NewKnowledgeGraphHandler(d.pgStores.KnowledgeGraph, d.providerRegistry)) + kgHandler := httpapi.NewKnowledgeGraphHandler(d.pgStores.KnowledgeGraph, d.providerRegistry) + kgHandler.SetUsageCapService(d.usageCapSvc) + d.server.SetKnowledgeGraphHandler(kgHandler) } // V3: Evolution metrics + suggestions API diff --git a/cmd/gateway_lifecycle.go b/cmd/gateway_lifecycle.go index 3a8ef20a..631bdd4e 100644 --- a/cmd/gateway_lifecycle.go +++ b/cmd/gateway_lifecycle.go @@ -140,7 +140,7 @@ func (d *gatewayDeps) runLifecycle( d.channelMgr.SetContactCollector(contactCollector) } - go consumeInboundMessages(ctx, d.msgBus, d.agentRouter, d.cfg, deps.sched, d.channelMgr, deps.consumerTeamStore, deps.quotaChecker, d.pgStores.Sessions, d.pgStores.Agents, contactCollector, deps.postTurn, deps.subagentMgr) + go consumeInboundMessages(ctx, d.msgBus, d.agentRouter, d.cfg, deps.sched, d.channelMgr, deps.consumerTeamStore, deps.quotaChecker, d.pgStores.Sessions, d.pgStores.Agents, contactCollector, deps.postTurn, deps.subagentMgr, d.usageCapSvc) // Webhook callback worker — delivers async webhook_calls rows to receiver callback_url. // Runs in both editions: Standard (PG, concurrency=4) and Lite (SQLite, concurrency=1). diff --git a/cmd/gateway_managed.go b/cmd/gateway_managed.go index 0db67050..798bff9f 100644 --- a/cmd/gateway_managed.go +++ b/cmd/gateway_managed.go @@ -185,7 +185,7 @@ func wireExtras( "disabled_count", n, "edition", edition.Current().Name) } - handlers := buildHookHandlers(stores, providerReg, appCfg.Hooks) + handlers := buildHookHandlers(stores, providerReg, appCfg.Hooks, usageCapSvc) stdOpts := hooks.StdDispatcherOpts{ Store: hs, Audit: hooks.NewAuditWriter(hs, ""), @@ -304,7 +304,7 @@ func wireExtras( writeMemIntc = tools.NewMemoryInterceptor(stores.Memory, workspace) // Hook KG extraction on memory writes if KG store is available if stores.KnowledgeGraph != nil && stores.BuiltinTools != nil { - writeMemIntc.SetKGExtractFunc(buildKGExtractFunc(stores.KnowledgeGraph, stores.BuiltinTools, providerReg)) + writeMemIntc.SetKGExtractFunc(buildKGExtractFunc(stores.KnowledgeGraph, stores.BuiltinTools, providerReg, usageCapSvc)) } } if readTool, ok := toolsReg.Get("read_file"); ok { @@ -712,7 +712,7 @@ type kgSettings struct { // buildKGExtractFunc returns a callback that extracts entities from memory content. // Settings are read from the builtin_tools table on each invocation (not cached), // so changes take effect immediately without restart. -func buildKGExtractFunc(kgStore store.KnowledgeGraphStore, bts store.BuiltinToolStore, providerReg *providers.Registry) tools.KGExtractFunc { +func buildKGExtractFunc(kgStore store.KnowledgeGraphStore, bts store.BuiltinToolStore, providerReg *providers.Registry, usageCapSvc *usagecaps.Service) tools.KGExtractFunc { return func(ctx context.Context, agentID, userID, content string) { slog.Info("kg extract: triggered", "agent", agentID, "user", userID, "content_len", len(content)) // Read settings from DB on each call so admin changes take effect immediately @@ -736,6 +736,7 @@ func buildKGExtractFunc(kgStore store.KnowledgeGraphStore, bts store.BuiltinTool return } extractor := kg.NewExtractor(p, settings.ExtractionModel, settings.MinConfidence) + extractor.SetUsageCapService(usageCapSvc) result, err := extractor.Extract(ctx, content) if err != nil { slog.Warn("kg extract: extraction failed", "agent", agentID, "error", err) diff --git a/cmd/gateway_methods.go b/cmd/gateway_methods.go index d994217e..7d7f0664 100644 --- a/cmd/gateway_methods.go +++ b/cmd/gateway_methods.go @@ -12,14 +12,16 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/gateway/methods" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) -func registerAllMethods(server *gateway.Server, agents *agent.Router, sessStore store.SessionStore, cronStore store.CronStore, pairingStore store.PairingStore, cfg *config.Config, cfgPath, workspace, dataDir string, msgBus *bus.MessageBus, execApprovalMgr *tools.ExecApprovalManager, agentStore store.AgentStore, skillStore store.SkillStore, configSecretsStore store.ConfigSecretsStore, teamStore store.TeamStore, contextFileInterceptor *tools.ContextFileInterceptor, logTee *gateway.LogTee, heartbeatStore store.HeartbeatStore, configPermStore store.ConfigPermissionStore, sysConfigStore store.SystemConfigStore, tenantStore store.TenantStore, skillTenantCfgStore store.SkillTenantConfigStore, audioMgr *audio.Manager) (*methods.PairingMethods, *methods.HeartbeatMethods, *methods.ChatMethods, *methods.ConfigPermissionsMethods) { +func registerAllMethods(server *gateway.Server, agents *agent.Router, sessStore store.SessionStore, cronStore store.CronStore, pairingStore store.PairingStore, cfg *config.Config, cfgPath, workspace, dataDir string, msgBus *bus.MessageBus, execApprovalMgr *tools.ExecApprovalManager, agentStore store.AgentStore, skillStore store.SkillStore, configSecretsStore store.ConfigSecretsStore, teamStore store.TeamStore, contextFileInterceptor *tools.ContextFileInterceptor, logTee *gateway.LogTee, heartbeatStore store.HeartbeatStore, configPermStore store.ConfigPermissionStore, sysConfigStore store.SystemConfigStore, tenantStore store.TenantStore, skillTenantCfgStore store.SkillTenantConfigStore, audioMgr *audio.Manager, usageCapSvc *usagecaps.Service) (*methods.PairingMethods, *methods.HeartbeatMethods, *methods.ChatMethods, *methods.ConfigPermissionsMethods) { router := server.Router() // Phase 1: Core methods chatMethods := methods.NewChatMethods(agents, sessStore, cfg, server.RateLimiter(), msgBus) chatMethods.SetAudioManager(audioMgr) // Wire TTS auto-apply for WS responses + chatMethods.SetUsageCapService(usageCapSvc) chatMethods.Register(router) methods.NewAgentsMethods(agents, cfg, cfgPath, workspace, agentStore, contextFileInterceptor, msgBus).Register(router) methods.NewSessionsMethods(sessStore, msgBus, cfg).Register(router) diff --git a/internal/agent/intent_classify.go b/internal/agent/intent_classify.go index 2fd4e660..46201708 100644 --- a/internal/agent/intent_classify.go +++ b/internal/agent/intent_classify.go @@ -8,6 +8,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/i18n" "github.com/nextlevelbuilder/goclaw/internal/providers" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // IntentType represents the classified intent of a user message. @@ -108,6 +109,10 @@ func containsWholeWord(s, kw string) bool { // Uses keyword fast-path first, then falls back to LLM classification. // Falls back to IntentNewTask on any error. func ClassifyIntent(ctx context.Context, provider providers.Provider, model, userMessage string) IntentType { + return ClassifyIntentWithUsageCaps(ctx, nil, provider, model, userMessage) +} + +func ClassifyIntentWithUsageCaps(ctx context.Context, usageCaps *usagecaps.Service, provider providers.Provider, model, userMessage string) IntentType { // Fast-path: keyword matching for obvious patterns (no LLM cost). if intent, ok := quickClassify(userMessage); ok { return intent @@ -116,7 +121,7 @@ func ClassifyIntent(ctx context.Context, provider providers.Provider, model, use ctx, cancel := context.WithTimeout(ctx, intentClassifyTimeout) defer cancel() - resp, err := provider.Chat(ctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{ {Role: "system", Content: intentSystemPrompt}, {Role: "user", Content: userMessage}, @@ -126,6 +131,11 @@ func ClassifyIntent(ctx context.Context, provider providers.Provider, model, use providers.OptMaxTokens: 20, providers.OptTemperature: 0.0, }, + } + resp, err := usageCaps.Chat(ctx, provider, req, usagecaps.ChatOptions{ + ModelID: model, + Purpose: "intent-classify", + MaxOutputTokens: 20, }) if err != nil { return IntentNewTask diff --git a/internal/agent/title_generate.go b/internal/agent/title_generate.go index eab223cf..67fcd302 100644 --- a/internal/agent/title_generate.go +++ b/internal/agent/title_generate.go @@ -7,6 +7,7 @@ import ( "time" "github.com/nextlevelbuilder/goclaw/internal/providers" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) const titleGenerateTimeout = 120 * time.Second @@ -16,10 +17,14 @@ const titleSystemPrompt = `Generate a short title (max 15 words) for this conver // GenerateTitle uses a lightweight LLM call to create a short conversation title // from the user's first message. Returns empty string on error. func GenerateTitle(ctx context.Context, provider providers.Provider, model, userMessage string) string { + return GenerateTitleWithUsageCaps(ctx, nil, provider, model, userMessage) +} + +func GenerateTitleWithUsageCaps(ctx context.Context, usageCaps *usagecaps.Service, provider providers.Provider, model, userMessage string) string { ctx, cancel := context.WithTimeout(ctx, titleGenerateTimeout) defer cancel() - resp, err := provider.Chat(ctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{ {Role: "system", Content: titleSystemPrompt}, {Role: "user", Content: userMessage}, @@ -29,13 +34,18 @@ func GenerateTitle(ctx context.Context, provider providers.Provider, model, user // Larger budget: thinking-capable models (Gemini 2.5/3, GPT-5 reasoning) // can consume output tokens on reasoning traces. 256 leaves room for a // 15-word title even when the provider allocates some budget to thinking. - providers.OptMaxTokens: 256, - providers.OptTemperature: 0.3, + providers.OptMaxTokens: 256, + providers.OptTemperature: 0.3, // Disable extended thinking for title generation — it's a trivial task // that doesn't benefit from reasoning and defaults (esp. Gemini's "high") // otherwise eat the entire max_tokens budget, truncating the title to 1 word. providers.OptThinkingLevel: "off", }, + } + resp, err := usageCaps.Chat(ctx, provider, req, usagecaps.ChatOptions{ + ModelID: model, + Purpose: "session-title", + MaxOutputTokens: 256, }) if err != nil { slog.Warn("title generation failed", "error", err) diff --git a/internal/channels/history_compaction.go b/internal/channels/history_compaction.go index 6480a6ad..79b20624 100644 --- a/internal/channels/history_compaction.go +++ b/internal/channels/history_compaction.go @@ -11,6 +11,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // CompactionConfig configures LLM-based history compaction. @@ -20,6 +21,7 @@ type CompactionConfig struct { MaxTokens int // max output tokens for summarization (default 4096) Provider providers.Provider // LLM provider for summarization Model string // model to use for summarization + UsageCaps *usagecaps.Service } // MaybeCompact checks if compaction is needed for a history key and triggers it in background. @@ -55,7 +57,7 @@ func (ph *PendingHistory) MaybeCompact(historyKey string, currentCount int, cfg // CompactGroup performs LLM-based compaction on a pending message group. // Reused by both auto-compact (channel) and HTTP compact endpoint. // Returns the number of entries remaining after compaction. -func CompactGroup(ctx context.Context, s store.PendingMessageStore, channelName, historyKey string, provider providers.Provider, model string, keepRecent, maxTokens int) (int, error) { +func CompactGroup(ctx context.Context, s store.PendingMessageStore, channelName, historyKey string, provider providers.Provider, model string, keepRecent, maxTokens int, usageCaps *usagecaps.Service) (int, error) { if keepRecent <= 0 { keepRecent = 40 } @@ -94,13 +96,18 @@ func CompactGroup(ctx context.Context, s store.PendingMessageStore, channelName, fmt.Fprintf(&sb, "%s%s: %s\n", prefix, ts, e.Body) } - resp, err := provider.Chat(ctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{{ Role: "user", Content: "Summarize these group chat messages concisely, preserving key topics, decisions, names, and important context:\n\n" + sb.String(), }}, Model: model, Options: map[string]any{"max_tokens": maxTokens, "temperature": 0.3}, + } + resp, err := usageCaps.Chat(ctx, provider, req, usagecaps.ChatOptions{ + ModelID: model, + Purpose: "pending-history-compaction", + MaxOutputTokens: maxTokens, }) if err != nil { return 0, fmt.Errorf("llm summarize: %w", err) @@ -185,7 +192,7 @@ func (ph *PendingHistory) runCompaction(historyKey string, cfg *CompactionConfig return } - _, err = CompactGroup(ctx, ph.store, ph.channelName, historyKey, cfg.Provider, cfg.Model, cfg.KeepRecent, cfg.MaxTokens) + _, err = CompactGroup(ctx, ph.store, ph.channelName, historyKey, cfg.Provider, cfg.Model, cfg.KeepRecent, cfg.MaxTokens, cfg.UsageCaps) if err != nil { slog.Warn("compaction.failed", "channel", ph.channelName, "key", historyKey, "error", err) return diff --git a/internal/channels/instance_loader.go b/internal/channels/instance_loader.go index df6d677f..7f845d80 100644 --- a/internal/channels/instance_loader.go +++ b/internal/channels/instance_loader.go @@ -16,6 +16,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/providerresolve" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // reloadStartTimeout bounds how long Reload() will wait for a single channel's @@ -39,6 +40,7 @@ type InstanceLoader struct { agentStore store.AgentStore providerReg *providers.Registry pendingCompactCfg *config.PendingCompactionConfig + usageCaps *usagecaps.Service factories map[string]ChannelFactory manager *Manager msgBus *bus.MessageBus @@ -78,6 +80,10 @@ func (l *InstanceLoader) SetPendingCompactionConfig(cfg *config.PendingCompactio l.pendingCompactCfg = cfg } +func (l *InstanceLoader) SetUsageCapService(s *usagecaps.Service) { + l.usageCaps = s +} + // RegisterFactory registers a factory for a channel type (e.g., "telegram", "discord"). func (l *InstanceLoader) RegisterFactory(channelType string, factory ChannelFactory) { l.factories[channelType] = factory @@ -306,8 +312,9 @@ func (l *InstanceLoader) loadInstance(ctx context.Context, inst store.ChannelIns if p != nil && model != "" { cc := &CompactionConfig{ - Provider: p, - Model: model, + Provider: p, + Model: model, + UsageCaps: l.usageCaps, } if l.pendingCompactCfg != nil { cc.Threshold = l.pendingCompactCfg.Threshold diff --git a/internal/consolidation/dreaming_worker.go b/internal/consolidation/dreaming_worker.go index 4a97bf76..4a981793 100644 --- a/internal/consolidation/dreaming_worker.go +++ b/internal/consolidation/dreaming_worker.go @@ -11,9 +11,10 @@ import ( "github.com/google/uuid" "github.com/nextlevelbuilder/goclaw/internal/bgalert" "github.com/nextlevelbuilder/goclaw/internal/eventbus" - "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providerresolve" + "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) const ( @@ -31,6 +32,7 @@ type dreamingWorker struct { systemConfigs store.SystemConfigStore // per-tenant provider config registry *providers.Registry // provider resolution alertDeps bgalert.AlertDeps + usageCaps *usagecaps.Service // threshold/debounce are the global defaults. Per-agent overrides come // from resolveConfig which reads the agent's MemoryConfig.Dreaming JSONB. @@ -95,6 +97,11 @@ func (w *dreamingWorker) Handle(ctx context.Context, event eventbus.DomainEvent) ctx = store.WithTenantID(ctx, tid) } } + if event.AgentID != "" { + if aid, err := uuid.Parse(event.AgentID); err == nil { + ctx = store.WithAgentID(ctx, aid) + } + } agentID := event.AgentID userID := event.UserID @@ -205,7 +212,7 @@ func (w *dreamingWorker) synthesize(ctx context.Context, provider providers.Prov } body := strings.Join(summaries, "\n---\n") - resp, err := provider.Chat(ctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{ {Role: "system", Content: dreamingSystemPrompt}, {Role: "user", Content: "Session summaries:\n---\n" + body + "\n---"}, @@ -214,6 +221,11 @@ func (w *dreamingWorker) synthesize(ctx context.Context, provider providers.Prov Options: map[string]any{ providers.OptMaxTokens: dreamingMaxTokens, }, + } + resp, err := w.usageCaps.Chat(ctx, provider, req, usagecaps.ChatOptions{ + ModelID: model, + Purpose: "dreaming-synthesis", + MaxOutputTokens: dreamingMaxTokens, }) if err != nil { return "", fmt.Errorf("dreaming chat: %w", err) diff --git a/internal/consolidation/episodic_worker.go b/internal/consolidation/episodic_worker.go index 8a5f8cf5..b93a5383 100644 --- a/internal/consolidation/episodic_worker.go +++ b/internal/consolidation/episodic_worker.go @@ -10,19 +10,21 @@ import ( "github.com/google/uuid" "github.com/nextlevelbuilder/goclaw/internal/bgalert" "github.com/nextlevelbuilder/goclaw/internal/eventbus" - "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providerresolve" + "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // episodicWorker handles session.completed events → creates episodic summaries. type episodicWorker struct { store store.EpisodicStore - sessions store.SessionCoreStore // for reading session messages during summarization - systemConfigs store.SystemConfigStore // per-tenant provider config - registry *providers.Registry // provider resolution + sessions store.SessionCoreStore // for reading session messages during summarization + systemConfigs store.SystemConfigStore // per-tenant provider config + registry *providers.Registry // provider resolution eventBus eventbus.DomainEventBus alertDeps bgalert.AlertDeps + usageCaps *usagecaps.Service } // resolveProvider delegates to shared background provider resolution. @@ -53,6 +55,7 @@ func (w *episodicWorker) Handle(ctx context.Context, event eventbus.DomainEvent) if err != nil { return fmt.Errorf("episodic: invalid agent_id %q: %w", event.AgentID, err) } + ctx = store.WithAgentID(ctx, agentUUID) // Build source_id for idempotency sourceID := fmt.Sprintf("%s:%d", payload.SessionKey, payload.CompactionCount) @@ -168,13 +171,18 @@ func (w *episodicWorker) summarizeFromMessages(ctx context.Context, provider pro sctx, cancel := context.WithTimeout(ctx, 30*time.Second) defer cancel() - resp, err := provider.Chat(sctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{ {Role: "system", Content: summarizationPrompt}, {Role: "user", Content: sb.String()}, }, Model: model, Options: map[string]any{"max_tokens": 1024, "temperature": 0.3}, + } + resp, err := w.usageCaps.Chat(sctx, provider, req, usagecaps.ChatOptions{ + ModelID: model, + Purpose: "episodic-summary", + MaxOutputTokens: 1024, }) if err != nil { return "", err diff --git a/internal/consolidation/workers.go b/internal/consolidation/workers.go index 1c5fb82a..2c3be24e 100644 --- a/internal/consolidation/workers.go +++ b/internal/consolidation/workers.go @@ -13,6 +13,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/eventbus" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // ConsolidationDeps bundles all dependencies for the consolidation pipeline. @@ -26,6 +27,7 @@ type ConsolidationDeps struct { Registry *providers.Registry // provider resolution Extractor EntityExtractor AlertDeps bgalert.AlertDeps // for reporting non-retryable LLM errors + UsageCaps *usagecaps.Service // AgentStore is optional: when present, the dreaming worker reads // per-agent overrides from MemoryConfig.Dreaming. If nil, the worker // uses its built-in defaults for every agent. @@ -42,6 +44,7 @@ func Register(deps ConsolidationDeps) func() { registry: deps.Registry, eventBus: deps.EventBus, alertDeps: deps.AlertDeps, + usageCaps: deps.UsageCaps, } semantic := &semanticWorker{ kgStore: deps.KGStore, @@ -59,6 +62,7 @@ func Register(deps ConsolidationDeps) func() { systemConfigs: deps.SystemConfigs, registry: deps.Registry, alertDeps: deps.AlertDeps, + usageCaps: deps.UsageCaps, threshold: dreamingDefaultThreshold, debounce: dreamingDefaultDebounce, resolveConfig: newAgentStoreResolver(deps.AgentStore), diff --git a/internal/gateway/methods/chat.go b/internal/gateway/methods/chat.go index 1281a9c5..aa0e98cc 100644 --- a/internal/gateway/methods/chat.go +++ b/internal/gateway/methods/chat.go @@ -20,6 +20,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/sessions" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/pkg/protocol" ) @@ -32,6 +33,7 @@ type ChatMethods struct { eventBus bus.EventPublisher postTurn tools.PostTurnProcessor audioMgr *audio.Manager // for TTS auto-apply on WS responses (nil = disabled) + usageCaps *usagecaps.Service debouncer *chatDebouncer } @@ -46,6 +48,10 @@ func (m *ChatMethods) SetAudioManager(mgr *audio.Manager) { m.audioMgr = mgr } +func (m *ChatMethods) SetUsageCapService(s *usagecaps.Service) { + m.usageCaps = s +} + // SetPostTurnProcessor sets the post-turn processor for team task dispatch. func (m *ChatMethods) SetPostTurnProcessor(pt tools.PostTurnProcessor) { m.postTurn = pt @@ -352,7 +358,10 @@ func (m *ChatMethods) dispatchChatSends(requests []chatSendRequest) { // Use runCtxBase (WithoutCancel + tenant-aware) so title save uses correct tenant. titleCtx := runCtxBase go func() { - title := agent.GenerateTitle(titleCtx, agentProvider, agentModel, userMsg) + if uid := loop.UUID(); uid != uuid.Nil { + titleCtx = store.WithAgentID(titleCtx, uid) + } + title := agent.GenerateTitleWithUsageCaps(titleCtx, m.usageCaps, agentProvider, agentModel, userMsg) if title == "" { return } diff --git a/internal/hooks/handlers/prompt.go b/internal/hooks/handlers/prompt.go index 2ec9fb1f..062c3c1f 100644 --- a/internal/hooks/handlers/prompt.go +++ b/internal/hooks/handlers/prompt.go @@ -17,6 +17,8 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/hooks" "github.com/nextlevelbuilder/goclaw/internal/hooks/budget" "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // ── Public surface ────────────────────────────────────────────────────────── @@ -46,6 +48,9 @@ type PromptHandler struct { // budget checks are skipped (Lite edition behavior). Budget *budget.Store + // UsageCaps enforces provider/model/tenant cost caps before the hook LLM call. + UsageCaps *usagecaps.Service + // DefaultModel is used when a hook config does not specify one. // Recommended: "haiku" for cheap evaluation. DefaultModel string @@ -155,7 +160,17 @@ func (h *PromptHandler) Execute(ctx context.Context, cfg hooks.HookConfig, ev ho req := h.buildChatRequest(cfg, ev, resolvedModel) // 6. Call provider. - resp, err := provider.Chat(ctx, req) + callCtx := ctx + if ev.TenantID != uuid.Nil { + callCtx = store.WithTenantID(callCtx, ev.TenantID) + } + resp, err := h.UsageCaps.Chat(callCtx, provider, req, usagecaps.ChatOptions{ + TenantID: ev.TenantID, + ProviderName: provider.Name(), + ModelID: resolvedModel, + Purpose: "hook-prompt", + MaxOutputTokens: promptMaxTokens(req), + }) if err != nil { // Fail-closed on transport/provider error for blocking events. if ev.HookEvent.IsBlocking() { @@ -300,6 +315,23 @@ func (h *PromptHandler) buildChatRequest(cfg hooks.HookConfig, ev hooks.Event, m } } +func promptMaxTokens(req providers.ChatRequest) int { + if req.Options == nil { + return 512 + } + if v, ok := req.Options[providers.OptMaxTokens]; ok { + switch n := v.(type) { + case int: + return n + case int64: + return int(n) + case float64: + return int(n) + } + } + return 512 +} + // sanitizeToolInput canonicalizes a tool_input map into stable JSON with // sorted keys, stripping only structural noise. Actual injection-attack // detection is delegated to the evaluator LLM (which has the anti-injection diff --git a/internal/http/knowledge_graph.go b/internal/http/knowledge_graph.go index 7333f851..fcc61b90 100644 --- a/internal/http/knowledge_graph.go +++ b/internal/http/knowledge_graph.go @@ -7,12 +7,14 @@ import ( kg "github.com/nextlevelbuilder/goclaw/internal/knowledgegraph" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // KnowledgeGraphHandler handles KG entity/relation management endpoints. type KnowledgeGraphHandler struct { store store.KnowledgeGraphStore providerReg *providers.Registry + usageCaps *usagecaps.Service } // NewKnowledgeGraphHandler creates a handler for KG management endpoints. @@ -20,6 +22,10 @@ func NewKnowledgeGraphHandler(s store.KnowledgeGraphStore, providerReg *provider return &KnowledgeGraphHandler{store: s, providerReg: providerReg} } +func (h *KnowledgeGraphHandler) SetUsageCapService(s *usagecaps.Service) { + h.usageCaps = s +} + // NewExtractor creates an Extractor from the given provider name and model. func (h *KnowledgeGraphHandler) NewExtractor(ctx context.Context, providerName, model string, minConfidence float64) *kg.Extractor { if h.providerReg == nil || providerName == "" || model == "" { @@ -29,7 +35,9 @@ func (h *KnowledgeGraphHandler) NewExtractor(ctx context.Context, providerName, if err != nil { return nil } - return kg.NewExtractor(p, model, minConfidence) + extractor := kg.NewExtractor(p, model, minConfidence) + extractor.SetUsageCapService(h.usageCaps) + return extractor } // RegisterRoutes registers all KG routes on the given mux. diff --git a/internal/http/knowledge_graph_handlers.go b/internal/http/knowledge_graph_handlers.go index 80cb420f..8e31b319 100644 --- a/internal/http/knowledge_graph_handlers.go +++ b/internal/http/knowledge_graph_handlers.go @@ -5,6 +5,7 @@ import ( "net/http" "strconv" + "github.com/google/uuid" "github.com/nextlevelbuilder/goclaw/internal/i18n" "github.com/nextlevelbuilder/goclaw/internal/store" ) @@ -156,6 +157,10 @@ func (h *KnowledgeGraphHandler) handleTraverse(w http.ResponseWriter, r *http.Re func (h *KnowledgeGraphHandler) handleExtract(w http.ResponseWriter, r *http.Request) { locale := extractLocale(r) agentID := r.PathValue("agentID") + callCtx := r.Context() + if parsedAgentID, err := uuid.Parse(agentID); err == nil { + callCtx = store.WithAgentID(callCtx, parsedAgentID) + } var body struct { Text string `json:"text"` @@ -176,13 +181,13 @@ func (h *KnowledgeGraphHandler) handleExtract(w http.ResponseWriter, r *http.Req return } - extractor := h.NewExtractor(r.Context(), body.Provider, body.Model, body.MinConf) + extractor := h.NewExtractor(callCtx, body.Provider, body.Model, body.MinConf) if extractor == nil { writeJSON(w, http.StatusBadRequest, map[string]string{"error": i18n.T(locale, i18n.MsgInvalidProviderOrModel)}) return } - result, err := extractor.Extract(r.Context(), body.Text) + result, err := extractor.Extract(callCtx, body.Text) if err != nil { slog.Warn("kg.extract failed", "error", err) writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) @@ -213,10 +218,10 @@ func (h *KnowledgeGraphHandler) handleExtract(w http.ResponseWriter, r *http.Req } writeJSON(w, http.StatusOK, map[string]any{ - "entities": len(result.Entities), - "relations": len(result.Relations), - "dedup_merged": dedupMerged, - "dedup_flagged": dedupFlagged, + "entities": len(result.Entities), + "relations": len(result.Relations), + "dedup_merged": dedupMerged, + "dedup_flagged": dedupFlagged, }) } diff --git a/internal/http/pending_messages.go b/internal/http/pending_messages.go index 6bc7d0a6..3f144fa7 100644 --- a/internal/http/pending_messages.go +++ b/internal/http/pending_messages.go @@ -11,6 +11,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/providerresolve" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // PendingMessagesHandler handles pending message HTTP endpoints. @@ -22,6 +23,7 @@ type PendingMessagesHandler struct { maxTokens int // max output tokens for LLM summarization (0 = use default) cfgProvider string // config-level provider override (empty = resolve from agent) cfgModel string // config-level model override (empty = resolve from agent) + usageCaps *usagecaps.Service } func NewPendingMessagesHandler(s store.PendingMessageStore, agentStore store.AgentStore, providerReg *providers.Registry) *PendingMessagesHandler { @@ -40,6 +42,10 @@ func (h *PendingMessagesHandler) SetProviderModel(provider, model string) { h.cfgModel = model } +func (h *PendingMessagesHandler) SetUsageCapService(s *usagecaps.Service) { + h.usageCaps = s +} + func (h *PendingMessagesHandler) RegisterRoutes(mux *http.ServeMux) { mux.HandleFunc("GET /v1/pending-messages", h.authMiddleware(h.handleListGroups)) mux.HandleFunc("GET /v1/pending-messages/messages", h.authMiddleware(h.handleListMessages)) @@ -148,7 +154,7 @@ func (h *PendingMessagesHandler) handleCompact(w http.ResponseWriter, r *http.Re go func() { ctx, cancel := context.WithTimeout(store.WithTenantID(context.Background(), tenantID), 180*time.Second) defer cancel() - remaining, err := channels.CompactGroup(ctx, h.store, req.ChannelName, req.HistoryKey, provider, model, keepRecent, h.maxTokens) + remaining, err := channels.CompactGroup(ctx, h.store, req.ChannelName, req.HistoryKey, provider, model, keepRecent, h.maxTokens, h.usageCaps) if err != nil { slog.Warn("compact.failed", "channel", req.ChannelName, "key", req.HistoryKey, "error", err) } else { diff --git a/internal/http/provider_verify.go b/internal/http/provider_verify.go index 48e4e591..8b3e402e 100644 --- a/internal/http/provider_verify.go +++ b/internal/http/provider_verify.go @@ -16,6 +16,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/i18n" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // HandleVerifyProviderForTest invokes the verify handler directly without auth @@ -126,8 +127,9 @@ func (h *ProvidersHandler) handleVerifyProvider(w http.ResponseWriter, r *http.R ctx, cancel := context.WithTimeout(r.Context(), 30*time.Second) defer cancel() + ctx = store.WithTenantID(ctx, p.TenantID) - _, err = provider.Chat(ctx, providers.ChatRequest{ + reqChat := providers.ChatRequest{ Messages: []providers.Message{ {Role: "user", Content: "hi"}, }, @@ -136,6 +138,13 @@ func (h *ProvidersHandler) handleVerifyProvider(w http.ResponseWriter, r *http.R // Use a small but safe value — reasoning models need headroom beyond 1 token. "max_tokens": 50, }, + } + _, err = h.usageCaps.Chat(ctx, provider, reqChat, usagecaps.ChatOptions{ + TenantID: p.TenantID, + ProviderName: p.Name, + ModelID: req.Model, + Purpose: "provider-verify", + MaxOutputTokens: 50, }) if err != nil { writeJSON(w, http.StatusOK, map[string]any{"valid": false, "error": friendlyVerifyError(err)}) diff --git a/internal/http/providers.go b/internal/http/providers.go index bee754f0..9127982f 100644 --- a/internal/http/providers.go +++ b/internal/http/providers.go @@ -22,6 +22,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/permissions" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/pkg/protocol" ) @@ -39,6 +40,7 @@ type ProvidersHandler struct { tracingStore store.TracingStore // optional: for provider-scoped pool activity agents store.AgentCRUDStore // optional: for provider pool activity agent lookup modelReg providers.ModelRegistry // optional: forward-compat model resolver for Anthropic + usageCaps *usagecaps.Service } // NewProvidersHandler creates a handler for provider management endpoints. @@ -85,6 +87,10 @@ func (h *ProvidersHandler) SetModelRegistry(r providers.ModelRegistry) { h.modelReg = r } +func (h *ProvidersHandler) SetUsageCapService(s *usagecaps.Service) { + h.usageCaps = s +} + // resolveAPIBase returns the provider's api_base, falling back to config/env if empty. // For Ollama/OllamaCloud providers, applies a safety-net normalization: if the stored // value is missing the /v1 suffix (pre-existing record before write-time normalization), diff --git a/internal/http/summoner.go b/internal/http/summoner.go index 540257c1..64f2effb 100644 --- a/internal/http/summoner.go +++ b/internal/http/summoner.go @@ -12,6 +12,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/bus" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // Summoning event type constants. @@ -48,14 +49,16 @@ type AgentSummoner struct { agents store.AgentStore providerReg *providers.Registry msgBus *bus.MessageBus + usageCaps *usagecaps.Service } // NewAgentSummoner creates a summoner backed by the given stores and provider registry. -func NewAgentSummoner(agents store.AgentStore, providerReg *providers.Registry, msgBus *bus.MessageBus) *AgentSummoner { +func NewAgentSummoner(agents store.AgentStore, providerReg *providers.Registry, msgBus *bus.MessageBus, usageCaps *usagecaps.Service) *AgentSummoner { return &AgentSummoner{ agents: agents, providerReg: providerReg, msgBus: msgBus, + usageCaps: usageCaps, } } @@ -73,6 +76,7 @@ const singleCallTimeout = 300 * time.Second func (s *AgentSummoner) SummonAgent(agentID uuid.UUID, tenantID uuid.UUID, providerName, model, description string) { ctx, cancel := context.WithTimeout(store.WithTenantID(context.Background(), tenantID), 600*time.Second) defer cancel() + ctx = store.WithAgentID(ctx, agentID) s.ensureBackfillFiles(ctx, agentID) s.emitEvent(agentID, tenantID, SummonEventStarted, "", "") diff --git a/internal/http/summoner_regenerate.go b/internal/http/summoner_regenerate.go index e3237719..126f87db 100644 --- a/internal/http/summoner_regenerate.go +++ b/internal/http/summoner_regenerate.go @@ -14,6 +14,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/bus" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/pkg/protocol" ) @@ -23,6 +24,7 @@ import ( func (s *AgentSummoner) RegenerateAgent(agentID uuid.UUID, tenantID uuid.UUID, providerName, model, editPrompt string) { ctx, cancel := context.WithTimeout(store.WithTenantID(context.Background(), tenantID), 300*time.Second) defer cancel() + ctx = store.WithAgentID(ctx, agentID) s.ensureBackfillFiles(ctx, agentID) @@ -121,7 +123,7 @@ func (s *AgentSummoner) generateFiles(ctx context.Context, providerName, model, slog.Info("summoning: calling LLM", "provider", providerName, "model", model, "prompt_len", len(prompt)) - resp, err := provider.Chat(ctx, providers.ChatRequest{ + req := providers.ChatRequest{ Messages: []providers.Message{ {Role: "system", Content: "You are a file generator. Output ONLY the requested XML-tagged files. No extra commentary."}, {Role: "user", Content: prompt}, @@ -133,6 +135,12 @@ func (s *AgentSummoner) generateFiles(ctx context.Context, providerName, model, providers.OptSessionKey: summonSessionKey, providers.OptDisableTools: true, }, + } + resp, err := s.usageCaps.Chat(ctx, provider, req, usagecaps.ChatOptions{ + ProviderName: providerName, + ModelID: model, + Purpose: "agent-summoner", + MaxOutputTokens: 8192, }) if err != nil { return nil, fmt.Errorf("%s: %w", providerName, err) diff --git a/internal/http/usage_caps.go b/internal/http/usage_caps.go index cf1940ca..bf2c3718 100644 --- a/internal/http/usage_caps.go +++ b/internal/http/usage_caps.go @@ -13,6 +13,7 @@ import ( "time" "github.com/google/uuid" + "github.com/nextlevelbuilder/goclaw/internal/i18n" "github.com/nextlevelbuilder/goclaw/internal/permissions" "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/usage/pricing" @@ -64,11 +65,15 @@ func (h *UsageCapsHandler) masterAuth(next http.HandlerFunc) http.HandlerFunc { }) } +func writeUsageCapError(w http.ResponseWriter, r *http.Request, status int, key string, args ...any) { + writeJSON(w, status, map[string]string{"error": i18n.T(store.LocaleFromContext(r.Context()), key, args...)}) +} + func (h *UsageCapsHandler) handleListPolicies(w http.ResponseWriter, r *http.Request) { scope := store.UsageCapScope{TenantID: tenantIDOrMaster(r)} policies, err := h.store.ListUsageCapPolicies(r.Context(), scope, true) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "list policies failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsageCapsListPoliciesFailed) return } writeJSON(w, http.StatusOK, map[string]any{"policies": policies}) @@ -77,17 +82,17 @@ func (h *UsageCapsHandler) handleListPolicies(w http.ResponseWriter, r *http.Req func (h *UsageCapsHandler) handleCreatePolicy(w http.ResponseWriter, r *http.Request) { var body policyBody if err := json.NewDecoder(r.Body).Decode(&body); err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid json"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidJSON) return } p, err := body.toPolicy(tenantIDOrMaster(r)) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": err.Error()}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidRequest, err.Error()) return } if err := h.store.CreateUsageCapPolicy(r.Context(), &p); err != nil { slog.Warn("usage_caps.create_policy_failed", "error", err) - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "usage cap policy validation failed"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgUsageCapPolicyValidationFailed) return } writeJSON(w, http.StatusCreated, p) @@ -96,27 +101,27 @@ func (h *UsageCapsHandler) handleCreatePolicy(w http.ResponseWriter, r *http.Req func (h *UsageCapsHandler) handleUpdatePolicy(w http.ResponseWriter, r *http.Request) { id, err := uuid.Parse(r.PathValue("id")) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid policy id"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidID, "policy") return } bodyBytes, err := io.ReadAll(r.Body) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid json"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidJSON) return } patch, err := policyPatchFromBody(bodyBytes) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": err.Error()}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidRequest, err.Error()) return } p, err := h.store.UpdateUsageCapPolicy(r.Context(), tenantIDOrMaster(r), id, patch) if err != nil { if errors.Is(err, store.ErrUsageCapPolicyManaged) { - writeJSON(w, http.StatusConflict, map[string]string{"error": err.Error()}) + writeUsageCapError(w, r, http.StatusConflict, i18n.MsgUsageCapPolicyManaged) return } slog.Warn("usage_caps.update_policy_failed", "error", err) - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "usage cap policy validation failed"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgUsageCapPolicyValidationFailed) return } writeJSON(w, http.StatusOK, p) @@ -125,15 +130,15 @@ func (h *UsageCapsHandler) handleUpdatePolicy(w http.ResponseWriter, r *http.Req func (h *UsageCapsHandler) handleDeletePolicy(w http.ResponseWriter, r *http.Request) { id, err := uuid.Parse(r.PathValue("id")) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid policy id"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidID, "policy") return } if err := h.store.DeleteUsageCapPolicy(r.Context(), tenantIDOrMaster(r), id); err != nil { if errors.Is(err, store.ErrUsageCapPolicyManaged) { - writeJSON(w, http.StatusConflict, map[string]string{"error": err.Error()}) + writeUsageCapError(w, r, http.StatusConflict, i18n.MsgUsageCapPolicyManaged) return } - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "delete failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsageCapsDeletePolicyFailed) return } w.WriteHeader(http.StatusNoContent) @@ -142,7 +147,7 @@ func (h *UsageCapsHandler) handleDeletePolicy(w http.ResponseWriter, r *http.Req func (h *UsageCapsHandler) handleUtilization(w http.ResponseWriter, r *http.Request) { rows, err := h.store.ListUsageCapUtilization(r.Context(), tenantIDOrMaster(r)) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "utilization failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsageCapsUtilizationFailed) return } writeJSON(w, http.StatusOK, map[string]any{"rows": rows}) @@ -152,7 +157,7 @@ func (h *UsageCapsHandler) handleEvents(w http.ResponseWriter, r *http.Request) limit, _ := strconv.Atoi(r.URL.Query().Get("limit")) events, err := h.store.ListUsageCapEvents(r.Context(), tenantIDOrMaster(r), limit) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "events failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsageCapsEventsFailed) return } writeJSON(w, http.StatusOK, map[string]any{"events": events}) @@ -164,12 +169,12 @@ func (h *UsageCapsHandler) handleSyncOpenRouter(w http.ResponseWriter, r *http.R entries, err := pricing.FetchOpenRouterCatalog(ctx, h.client) if err != nil { slog.Warn("usage_pricing.openrouter_sync", "error", err) - writeJSON(w, http.StatusBadGateway, map[string]string{"error": err.Error()}) + writeUsageCapError(w, r, http.StatusBadGateway, i18n.MsgUsagePricingSyncOpenRouterFailed, err.Error()) return } count, err := h.store.UpsertPricingCatalog(r.Context(), entries) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "store catalog failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsagePricingStoreCatalogFailed) return } writeJSON(w, http.StatusOK, map[string]any{"count": count}) @@ -181,7 +186,7 @@ func (h *UsageCapsHandler) handleListPricing(w http.ResponseWriter, r *http.Requ Limit: queryInt(r, "limit", 100), }) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "list pricing failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsagePricingListFailed) return } writeJSON(w, http.StatusOK, map[string]any{"models": rows}) @@ -190,12 +195,12 @@ func (h *UsageCapsHandler) handleListPricing(w http.ResponseWriter, r *http.Requ func (h *UsageCapsHandler) handlePutOverride(w http.ResponseWriter, r *http.Request) { var body overrideBody if err := json.NewDecoder(r.Body).Decode(&body); err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid json"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidJSON) return } providerID, err := uuid.Parse(body.ProviderID) if err != nil || body.ModelID == "" { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "provider_id and model_id are required"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgUsagePricingProviderModelRequired) return } o := &store.UsagePricingOverride{ @@ -205,7 +210,7 @@ func (h *UsageCapsHandler) handlePutOverride(w http.ResponseWriter, r *http.Requ } if err := h.store.PutPricingOverride(r.Context(), o); err != nil { slog.Warn("usage_pricing.put_override_failed", "error", err) - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "pricing override validation failed"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgUsagePricingOverrideValidationFailed) return } writeJSON(w, http.StatusOK, o) @@ -217,13 +222,13 @@ func (h *UsageCapsHandler) handleListOverrides(w http.ResponseWriter, r *http.Re var err error providerID, err = uuid.Parse(raw) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid provider_id"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidID, "provider") return } } rows, err := h.store.ListPricingOverrides(r.Context(), store.UsagePricingQuery{TenantID: tenantIDOrMaster(r), ProviderID: providerID}) if err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "list overrides failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsagePricingListOverridesFailed) return } writeJSON(w, http.StatusOK, map[string]any{"overrides": rows}) @@ -232,11 +237,11 @@ func (h *UsageCapsHandler) handleListOverrides(w http.ResponseWriter, r *http.Re func (h *UsageCapsHandler) handleDeleteOverride(w http.ResponseWriter, r *http.Request) { id, err := uuid.Parse(r.PathValue("id")) if err != nil { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid override id"}) + writeUsageCapError(w, r, http.StatusBadRequest, i18n.MsgInvalidID, "override") return } if err := h.store.DeletePricingOverride(r.Context(), tenantIDOrMaster(r), id); err != nil { - writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "delete failed"}) + writeUsageCapError(w, r, http.StatusInternalServerError, i18n.MsgUsagePricingDeleteOverrideFailed) return } w.WriteHeader(http.StatusNoContent) diff --git a/internal/i18n/catalog_en.go b/internal/i18n/catalog_en.go index 4013913d..c87a7a03 100644 --- a/internal/i18n/catalog_en.go +++ b/internal/i18n/catalog_en.go @@ -87,6 +87,21 @@ func init() { // Provider MsgProviderReqFailed: "%s: request failed: %s", + // Usage caps / pricing + MsgUsageCapsListPoliciesFailed: "failed to list usage cap policies", + MsgUsageCapPolicyValidationFailed: "usage cap policy validation failed", + MsgUsageCapPolicyManaged: "managed usage cap policies cannot be modified", + MsgUsageCapsDeletePolicyFailed: "failed to delete usage cap policy", + MsgUsageCapsUtilizationFailed: "failed to load usage cap utilization", + MsgUsageCapsEventsFailed: "failed to load usage cap events", + MsgUsagePricingSyncOpenRouterFailed: "failed to sync OpenRouter pricing: %s", + MsgUsagePricingStoreCatalogFailed: "failed to store pricing catalog", + MsgUsagePricingListFailed: "failed to list model pricing", + MsgUsagePricingProviderModelRequired: "provider_id and model_id are required", + MsgUsagePricingOverrideValidationFailed: "pricing override validation failed", + MsgUsagePricingListOverridesFailed: "failed to list pricing overrides", + MsgUsagePricingDeleteOverrideFailed: "failed to delete pricing override", + // Unknown method MsgUnknownMethod: "unknown method: %s", @@ -200,10 +215,10 @@ func init() { MsgTenantScopeRequired: "tenant scope is required for this operation", // TTS / Voices - MsgTtsUnknownModel: "unknown tts model: %s", - MsgVoicesListFailed: "failed to list voices: %s", - MsgTtsGeminiInvalidVoice: "invalid Gemini voice: %s", - MsgTtsGeminiSpeakerLimit: "Gemini TTS supports at most 2 speakers", + MsgTtsUnknownModel: "unknown tts model: %s", + MsgVoicesListFailed: "failed to list voices: %s", + MsgTtsGeminiInvalidVoice: "invalid Gemini voice: %s", + MsgTtsGeminiSpeakerLimit: "Gemini TTS supports at most 2 speakers", MsgTtsGeminiInvalidModel: "invalid Gemini TTS model: %s", MsgTtsGeminiTextOnly: "Gemini refused to generate audio. Try simpler text without translation or commentary.", MsgTtsParamOutOfRange: "TTS param %q value %v is out of range [%v, %v]", diff --git a/internal/i18n/catalog_vi.go b/internal/i18n/catalog_vi.go index fe5c1073..db80a909 100644 --- a/internal/i18n/catalog_vi.go +++ b/internal/i18n/catalog_vi.go @@ -87,6 +87,21 @@ func init() { // Provider MsgProviderReqFailed: "%s: yêu cầu thất bại: %s", + // Usage caps / pricing + MsgUsageCapsListPoliciesFailed: "không thể liệt kê chính sách usage cap", + MsgUsageCapPolicyValidationFailed: "xác thực chính sách usage cap thất bại", + MsgUsageCapPolicyManaged: "không thể chỉnh sửa chính sách usage cap do hệ thống quản lý", + MsgUsageCapsDeletePolicyFailed: "không thể xóa chính sách usage cap", + MsgUsageCapsUtilizationFailed: "không thể tải mức sử dụng usage cap", + MsgUsageCapsEventsFailed: "không thể tải sự kiện usage cap", + MsgUsagePricingSyncOpenRouterFailed: "không thể đồng bộ giá OpenRouter: %s", + MsgUsagePricingStoreCatalogFailed: "không thể lưu catalog giá", + MsgUsagePricingListFailed: "không thể liệt kê giá model", + MsgUsagePricingProviderModelRequired: "provider_id và model_id là bắt buộc", + MsgUsagePricingOverrideValidationFailed: "xác thực override giá thất bại", + MsgUsagePricingListOverridesFailed: "không thể liệt kê override giá", + MsgUsagePricingDeleteOverrideFailed: "không thể xóa override giá", + // Unknown method MsgUnknownMethod: "phương thức không xác định: %s", @@ -200,10 +215,10 @@ func init() { MsgTenantScopeRequired: "cần xác định tenant để thực hiện thao tác này", // TTS / Giọng đọc - MsgTtsUnknownModel: "model tts không hỗ trợ: %s", - MsgVoicesListFailed: "không tải được danh sách giọng đọc: %s", - MsgTtsGeminiInvalidVoice: "giọng đọc Gemini không hợp lệ: %s", - MsgTtsGeminiSpeakerLimit: "Gemini TTS hỗ trợ tối đa 2 người nói", + MsgTtsUnknownModel: "model tts không hỗ trợ: %s", + MsgVoicesListFailed: "không tải được danh sách giọng đọc: %s", + MsgTtsGeminiInvalidVoice: "giọng đọc Gemini không hợp lệ: %s", + MsgTtsGeminiSpeakerLimit: "Gemini TTS hỗ trợ tối đa 2 người nói", MsgTtsGeminiInvalidModel: "mô hình Gemini TTS không hợp lệ: %s", MsgTtsGeminiTextOnly: "Gemini từ chối tạo âm thanh. Vui lòng thử văn bản đơn giản hơn, không dịch hay bình luận.", MsgTtsParamOutOfRange: "tham số TTS %q có giá trị %v nằm ngoài phạm vi [%v, %v]", diff --git a/internal/i18n/catalog_zh.go b/internal/i18n/catalog_zh.go index 0fac3cbb..698ff34c 100644 --- a/internal/i18n/catalog_zh.go +++ b/internal/i18n/catalog_zh.go @@ -87,6 +87,21 @@ func init() { // Provider MsgProviderReqFailed: "%s:请求失败:%s", + // Usage caps / pricing + MsgUsageCapsListPoliciesFailed: "无法列出 usage cap 策略", + MsgUsageCapPolicyValidationFailed: "usage cap 策略验证失败", + MsgUsageCapPolicyManaged: "无法修改系统托管的 usage cap 策略", + MsgUsageCapsDeletePolicyFailed: "无法删除 usage cap 策略", + MsgUsageCapsUtilizationFailed: "无法加载 usage cap 使用量", + MsgUsageCapsEventsFailed: "无法加载 usage cap 事件", + MsgUsagePricingSyncOpenRouterFailed: "无法同步 OpenRouter 价格:%s", + MsgUsagePricingStoreCatalogFailed: "无法保存价格目录", + MsgUsagePricingListFailed: "无法列出模型价格", + MsgUsagePricingProviderModelRequired: "provider_id 和 model_id 是必填项", + MsgUsagePricingOverrideValidationFailed: "价格覆盖验证失败", + MsgUsagePricingListOverridesFailed: "无法列出价格覆盖", + MsgUsagePricingDeleteOverrideFailed: "无法删除价格覆盖", + // Unknown method MsgUnknownMethod: "未知方法:%s", @@ -200,10 +215,10 @@ func init() { MsgTenantScopeRequired: "此操作需要指定租户范围", // TTS / 声音 - MsgTtsUnknownModel: "未知的 tts 模型:%s", - MsgVoicesListFailed: "获取声音列表失败:%s", - MsgTtsGeminiInvalidVoice: "无效的 Gemini 声音:%s", - MsgTtsGeminiSpeakerLimit: "Gemini TTS 最多支持 2 位发言人", + MsgTtsUnknownModel: "未知的 tts 模型:%s", + MsgVoicesListFailed: "获取声音列表失败:%s", + MsgTtsGeminiInvalidVoice: "无效的 Gemini 声音:%s", + MsgTtsGeminiSpeakerLimit: "Gemini TTS 最多支持 2 位发言人", MsgTtsGeminiInvalidModel: "无效的 Gemini TTS 模型:%s", MsgTtsGeminiTextOnly: "Gemini 拒绝生成音频。请尝试更简单的文本,不要翻译或添加评论。", MsgTtsParamOutOfRange: "TTS 参数 %q 的值 %v 超出范围 [%v, %v]", diff --git a/internal/i18n/keys.go b/internal/i18n/keys.go index f6644b51..b53a7648 100644 --- a/internal/i18n/keys.go +++ b/internal/i18n/keys.go @@ -88,6 +88,21 @@ const ( // --- Provider --- MsgProviderReqFailed = "error.provider_request_failed" // "%s: request failed: %s" + // --- Usage caps / pricing --- + MsgUsageCapsListPoliciesFailed = "usage_caps.list_policies_failed" + MsgUsageCapPolicyValidationFailed = "usage_caps.policy_validation_failed" + MsgUsageCapPolicyManaged = "usage_caps.policy_managed" + MsgUsageCapsDeletePolicyFailed = "usage_caps.delete_policy_failed" + MsgUsageCapsUtilizationFailed = "usage_caps.utilization_failed" + MsgUsageCapsEventsFailed = "usage_caps.events_failed" + MsgUsagePricingSyncOpenRouterFailed = "usage_pricing.sync_openrouter_failed" + MsgUsagePricingStoreCatalogFailed = "usage_pricing.store_catalog_failed" + MsgUsagePricingListFailed = "usage_pricing.list_failed" + MsgUsagePricingProviderModelRequired = "usage_pricing.provider_model_required" + MsgUsagePricingOverrideValidationFailed = "usage_pricing.override_validation_failed" + MsgUsagePricingListOverridesFailed = "usage_pricing.list_overrides_failed" + MsgUsagePricingDeleteOverrideFailed = "usage_pricing.delete_override_failed" + // --- Unknown method --- MsgUnknownMethod = "error.unknown_method" // "unknown method: %s" @@ -117,14 +132,14 @@ const ( MsgInvalidVisibility = "error.invalid_visibility" // "invalid visibility %q: must be one of private, public" // --- Package updates (Phase 4+5) --- - MsgPackageNotInstalled = "packages.update.not_installed" // "Package {name} is not installed" - MsgPackageUpdateLocked = "packages.update.locked" // "Package {name} is being updated by another request" + MsgPackageNotInstalled = "packages.update.not_installed" // "Package {name} is not installed" + MsgPackageUpdateLocked = "packages.update.locked" // "Package {name} is being updated by another request" MsgReleaseNotFound = "packages.update.release_not_found" // "Release {tag} not found for {repo}" - MsgAssetNotFound = "packages.update.asset_not_found" // "No compatible asset for {os}/{arch}" + MsgAssetNotFound = "packages.update.asset_not_found" // "No compatible asset for {os}/{arch}" MsgChecksumMismatch = "packages.update.checksum_mismatch" // "Checksum mismatch for {name}" - MsgUpdateSwapFailed = "packages.update.swap_failed" // "Failed to install {name}; previous version restored" - MsgUpdateManifestDesync = "packages.update.manifest_desync" // "Binary updated but manifest save failed — manual recovery required for {name}" - MsgUpdateCacheStale = "packages.update.cache_stale" // "Updates cache stale; run refresh before applying an update" + MsgUpdateSwapFailed = "packages.update.swap_failed" // "Failed to install {name}; previous version restored" + MsgUpdateManifestDesync = "packages.update.manifest_desync" // "Binary updated but manifest save failed — manual recovery required for {name}" + MsgUpdateCacheStale = "packages.update.cache_stale" // "Updates cache stale; run refresh before applying an update" // Package update source labels MsgPackagesUpdatesSourceGithub = "packages.updates.source.github" // "GitHub" @@ -234,15 +249,15 @@ const ( MsgInvalidRole = "error.invalid_role" // "invalid role: allowed values are owner, admin, operator, member, viewer" // --- TTS / Voices --- - MsgTtsUnknownModel = "error.tts_unknown_model" // "unknown tts model: %s" - MsgVoicesListFailed = "error.voices_list_failed" // "failed to list voices: %s" - MsgTtsGeminiInvalidVoice = "error.tts_gemini_invalid_voice" // "invalid Gemini voice: %s" - MsgTtsGeminiSpeakerLimit = "error.tts_gemini_speaker_limit" // "Gemini TTS supports at most 2 speakers" - MsgTtsGeminiInvalidModel = "error.tts_gemini_invalid_model" // "invalid Gemini TTS model: %s" - MsgTtsGeminiTextOnly = "error.tts_gemini_text_only" // "Gemini refused to generate audio; try simpler text without translation or commentary" - MsgTtsParamOutOfRange = "error.tts_param_out_of_range" // "TTS param %q value %v is out of range [%v, %v]" - MsgTtsParamUnknownKey = "error.tts_param_unknown_key" // "TTS param %q is not supported by this provider" - MsgTtsMiniMaxVoicesFailed = "error.tts_minimax_voices_failed" // "failed to fetch MiniMax voices: %s" + MsgTtsUnknownModel = "error.tts_unknown_model" // "unknown tts model: %s" + MsgVoicesListFailed = "error.voices_list_failed" // "failed to list voices: %s" + MsgTtsGeminiInvalidVoice = "error.tts_gemini_invalid_voice" // "invalid Gemini voice: %s" + MsgTtsGeminiSpeakerLimit = "error.tts_gemini_speaker_limit" // "Gemini TTS supports at most 2 speakers" + MsgTtsGeminiInvalidModel = "error.tts_gemini_invalid_model" // "invalid Gemini TTS model: %s" + MsgTtsGeminiTextOnly = "error.tts_gemini_text_only" // "Gemini refused to generate audio; try simpler text without translation or commentary" + MsgTtsParamOutOfRange = "error.tts_param_out_of_range" // "TTS param %q value %v is out of range [%v, %v]" + MsgTtsParamUnknownKey = "error.tts_param_unknown_key" // "TTS param %q is not supported by this provider" + MsgTtsMiniMaxVoicesFailed = "error.tts_minimax_voices_failed" // "failed to fetch MiniMax voices: %s" // --- STT --- MsgSTTAllProvidersFailed = "error.stt_all_providers_failed" // "All STT providers failed" @@ -258,50 +273,50 @@ const ( MsgTenantScopeRequired = "error.tenant_scope_required" // "tenant scope is required for this operation" // --- Webhooks --- - MsgWebhookAuthFailed = "webhook.auth_failed" // "webhook authentication failed" - MsgWebhookHMACInvalid = "webhook.hmac_invalid" // "HMAC signature is invalid" - MsgWebhookHMACTimestampSkew = "webhook.hmac_timestamp_skew" // "request timestamp outside acceptable window" - MsgWebhookBearerRequiredHMAC = "webhook.bearer_required_hmac" // "this webhook requires HMAC authentication" - MsgWebhookRevoked = "webhook.revoked" // "webhook has been revoked" - MsgWebhookKindMismatch = "webhook.kind_mismatch" // "request kind does not match webhook configuration" - MsgWebhookRateLimited = "webhook.rate_limited" // "webhook rate limit exceeded" - MsgWebhookBodyTooLarge = "webhook.body_too_large" // "request body exceeds size limit" - MsgWebhookIdempotencyConflict = "webhook.idempotency_conflict" // "idempotency key conflict: request body mismatch" - MsgWebhookTenantMismatch = "webhook.tenant_mismatch" // "webhook tenant mismatch" - MsgWebhookAgentNotFound = "webhook.agent_not_found" // "webhook agent not found" - MsgWebhookChannelNotFound = "webhook.channel_not_found" // "webhook channel not found" - MsgWebhookMediaSSRFBlocked = "webhook.media_ssrf_blocked" // "media URL blocked by SSRF policy" - MsgWebhookMediaTooLarge = "webhook.media_too_large" // "media file exceeds size limit" - MsgWebhookMediaMIMEDenied = "webhook.media_mime_denied" // "media MIME type is not allowed" - MsgWebhookCallbackURLInvalid = "webhook.callback_url_invalid" // "callback URL is invalid or blocked" - MsgWebhookLLMTimeout = "webhook.llm_timeout" // "LLM processing timed out" - MsgWebhookLaneSaturated = "webhook.lane_saturated" // "webhook processing lane is at capacity" + MsgWebhookAuthFailed = "webhook.auth_failed" // "webhook authentication failed" + MsgWebhookHMACInvalid = "webhook.hmac_invalid" // "HMAC signature is invalid" + MsgWebhookHMACTimestampSkew = "webhook.hmac_timestamp_skew" // "request timestamp outside acceptable window" + MsgWebhookBearerRequiredHMAC = "webhook.bearer_required_hmac" // "this webhook requires HMAC authentication" + MsgWebhookRevoked = "webhook.revoked" // "webhook has been revoked" + MsgWebhookKindMismatch = "webhook.kind_mismatch" // "request kind does not match webhook configuration" + MsgWebhookRateLimited = "webhook.rate_limited" // "webhook rate limit exceeded" + MsgWebhookBodyTooLarge = "webhook.body_too_large" // "request body exceeds size limit" + MsgWebhookIdempotencyConflict = "webhook.idempotency_conflict" // "idempotency key conflict: request body mismatch" + MsgWebhookTenantMismatch = "webhook.tenant_mismatch" // "webhook tenant mismatch" + MsgWebhookAgentNotFound = "webhook.agent_not_found" // "webhook agent not found" + MsgWebhookChannelNotFound = "webhook.channel_not_found" // "webhook channel not found" + MsgWebhookMediaSSRFBlocked = "webhook.media_ssrf_blocked" // "media URL blocked by SSRF policy" + MsgWebhookMediaTooLarge = "webhook.media_too_large" // "media file exceeds size limit" + MsgWebhookMediaMIMEDenied = "webhook.media_mime_denied" // "media MIME type is not allowed" + MsgWebhookCallbackURLInvalid = "webhook.callback_url_invalid" // "callback URL is invalid or blocked" + MsgWebhookLLMTimeout = "webhook.llm_timeout" // "LLM processing timed out" + MsgWebhookLaneSaturated = "webhook.lane_saturated" // "webhook processing lane is at capacity" MsgWebhookLocalhostOnlyViolation = "webhook.localhost_only_violation" // "this webhook is restricted to localhost callers" MsgWebhookMediaChannelUnsupported = "webhook.media_channel_unsupported" // "channel does not support media attachments" MsgWebhookIPDenied = "webhook.ip_denied" // "request origin is not in the IP allowlist" MsgWebhookEncryptionUnavailable = "webhook.encryption_unavailable" // "webhook encryption key not configured; set GOCLAW_ENCRYPTION_KEY to enable webhooks" // --- Workstation permissions --- - MsgWorkstationCmdDenied = "error.workstation_cmd_denied" // "command denied by workstation policy: %s" - MsgWorkstationEnvDenied = "error.workstation_env_denied" // "env var denied by policy: %s" - MsgWorkstationInputInvalid = "error.workstation_input_invalid" // "command contains invalid characters: %s" - MsgWorkstationRateLimit = "error.workstation_rate_limit" // "workstation rate limit exceeded" - MsgWorkstationPermNotFound = "error.workstation_perm_not_found" // "permission entry not found: %s" + MsgWorkstationCmdDenied = "error.workstation_cmd_denied" // "command denied by workstation policy: %s" + MsgWorkstationEnvDenied = "error.workstation_env_denied" // "env var denied by policy: %s" + MsgWorkstationInputInvalid = "error.workstation_input_invalid" // "command contains invalid characters: %s" + MsgWorkstationRateLimit = "error.workstation_rate_limit" // "workstation rate limit exceeded" + MsgWorkstationPermNotFound = "error.workstation_perm_not_found" // "permission entry not found: %s" // --- Workstation activity (Phase 7) --- - MsgWorkstationActivityTitle = "ui.workstations.activity.title" // "Recent Activity" - MsgWorkstationActionExec = "ui.workstations.activity.action_exec" // "Exec" - MsgWorkstationActionDeny = "ui.workstations.activity.action_deny" // "Denied" + MsgWorkstationActivityTitle = "ui.workstations.activity.title" // "Recent Activity" + MsgWorkstationActionExec = "ui.workstations.activity.action_exec" // "Exec" + MsgWorkstationActionDeny = "ui.workstations.activity.action_deny" // "Denied" // --- Workstation --- - MsgWorkstationNotFound = "error.workstation_not_found" // "workstation not found: %s" - MsgWorkstationKeyExists = "error.workstation_key_exists" // "workstation key already in use: %s" - MsgInvalidBackend = "error.invalid_backend" // "invalid backend type: %s (must be ssh|docker)" - MsgWorkstationInactive = "error.workstation_inactive" // "workstation is inactive: %s" - MsgInvalidMetadataShape = "error.invalid_metadata_shape" // "invalid metadata for %s backend: %s" - MsgWorkstationRequired = "error.workstation_required" // "no workstation bound to agent; pass workstation_id" + MsgWorkstationNotFound = "error.workstation_not_found" // "workstation not found: %s" + MsgWorkstationKeyExists = "error.workstation_key_exists" // "workstation key already in use: %s" + MsgInvalidBackend = "error.invalid_backend" // "invalid backend type: %s (must be ssh|docker)" + MsgWorkstationInactive = "error.workstation_inactive" // "workstation is inactive: %s" + MsgInvalidMetadataShape = "error.invalid_metadata_shape" // "invalid metadata for %s backend: %s" + MsgWorkstationRequired = "error.workstation_required" // "no workstation bound to agent; pass workstation_id" MsgWorkstationAccessDenied = "error.workstation_access_denied" // "agent %s not authorized for workstation %s" - MsgBackendNotReady = "error.backend_not_ready" // "workstation backend not ready: %s" + MsgBackendNotReady = "error.backend_not_ready" // "workstation backend not ready: %s" // --- Hooks --- MsgHookInvalidMatcher = "hook.invalid_matcher" // "invalid matcher regex: %s" diff --git a/internal/knowledgegraph/extractor.go b/internal/knowledgegraph/extractor.go index 491c6c2c..addb80df 100644 --- a/internal/knowledgegraph/extractor.go +++ b/internal/knowledgegraph/extractor.go @@ -9,6 +9,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" ) // ExtractionResult holds entities and relations extracted from text. @@ -22,6 +23,7 @@ type Extractor struct { provider providers.Provider model string minConfidence float64 + usageCaps *usagecaps.Service } // NewExtractor creates a new Extractor with the given provider, model, and confidence threshold. @@ -32,6 +34,11 @@ func NewExtractor(provider providers.Provider, model string, minConfidence float return &Extractor{provider: provider, model: model, minConfidence: minConfidence} } +// SetUsageCapService enables cost enforcement for LLM extraction calls. +func (e *Extractor) SetUsageCapService(s *usagecaps.Service) { + e.usageCaps = s +} + const maxChunkChars = 12000 // Extract calls the LLM to extract entities and relations from text. @@ -72,7 +79,11 @@ func (e *Extractor) extractChunk(ctx context.Context, text string) (*ExtractionR }, } - resp, err := e.provider.Chat(ctx, req) + resp, err := e.usageCaps.Chat(ctx, e.provider, req, usagecaps.ChatOptions{ + ModelID: e.model, + Purpose: "knowledge-graph-extract", + MaxOutputTokens: 8192, + }) if err != nil { return nil, fmt.Errorf("kg extraction LLM call: %w", err) } @@ -85,7 +96,11 @@ func (e *Extractor) extractChunk(ctx context.Context, text string) (*ExtractionR text = text[:retryMaxChars] + "\n\n[...truncated]" } req.Messages[1].Content = text - resp, err = e.provider.Chat(ctx, req) + resp, err = e.usageCaps.Chat(ctx, e.provider, req, usagecaps.ChatOptions{ + ModelID: e.model, + Purpose: "knowledge-graph-extract-retry", + MaxOutputTokens: 8192, + }) if err != nil { return nil, fmt.Errorf("kg extraction LLM retry: %w", err) } @@ -286,4 +301,3 @@ func stripCodeBlock(s string) string { } return strings.TrimSpace(s) } - diff --git a/internal/store/pg/usage_caps.go b/internal/store/pg/usage_caps.go index 06bdeaa4..f693cb63 100644 --- a/internal/store/pg/usage_caps.go +++ b/internal/store/pg/usage_caps.go @@ -267,9 +267,16 @@ func (s *PGUsageCapStore) ListUsageCapUtilization(ctx context.Context, tenantID for _, p := range policies { start, end := usageWindow(time.Now().UTC(), p.Window) u := store.UsageCapUtilization{Policy: p, WindowStart: start, WindowEnd: end} - _ = s.db.QueryRowContext(ctx, ` + err := s.db.QueryRowContext(ctx, ` SELECT used_tokens, reserved_tokens, used_cost_micros, reserved_cost_micros FROM usage_cap_counters WHERE policy_id=$1 AND window_start=$2`, p.ID, start).Scan(&u.UsedTokens, &u.ReservedTokens, &u.UsedCostMicros, &u.ReservedCostMicros) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + out = append(out, u) + continue + } + return nil, err + } out = append(out, u) } return out, nil diff --git a/internal/usage/caps/chat_call.go b/internal/usage/caps/chat_call.go new file mode 100644 index 00000000..5883d263 --- /dev/null +++ b/internal/usage/caps/chat_call.go @@ -0,0 +1,150 @@ +package caps + +import ( + "context" + "errors" + "fmt" + + "github.com/google/uuid" + "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store" +) + +// ChatOptions identifies a billable non-agent LLM call for usage-cap enforcement. +type ChatOptions struct { + TenantID uuid.UUID + AgentID uuid.UUID + ProviderName string + ModelID string + ReservationKey string + Purpose string + MaxOutputTokens int +} + +// Chat wraps Provider.Chat with the same usage-cap preflight and reconciliation +// used by agent loops. A nil service intentionally falls back to direct calls +// for Lite/subscription-only runtimes. +func (s *Service) Chat(ctx context.Context, provider providers.Provider, req providers.ChatRequest, opts ChatOptions) (*providers.ChatResponse, error) { + if provider == nil { + return nil, errors.New("usage cap chat: provider is nil") + } + if s == nil || s.store == nil { + return provider.Chat(ctx, req) + } + if fallback, ok := provider.(*providers.ModelFallbackProvider); ok { + return fallback.ChatWithHook(ctx, req, func(callCtx context.Context, entry providers.FallbackCandidate, actualReq providers.ChatRequest) (providers.FallbackAfterCall, error) { + callOpts := opts + callOpts.ProviderName = entry.ProviderName + if callOpts.ProviderName == "" && entry.Provider != nil { + callOpts.ProviderName = entry.Provider.Name() + } + callOpts.ModelID = actualReq.Model + callOpts.ReservationKey = "" + usageReq := s.chatRequest(callCtx, entry.Provider, actualReq, callOpts) + scopedCtx := scopedRequestContext(callCtx, usageReq) + reservation, err := s.Preflight(scopedCtx, usageReq) + if err != nil { + return nil, err + } + return func(resp *providers.ChatResponse, callErr error, _ providers.FallbackCallInfo) { + if reservation != nil { + reservation.Reconcile(scopedCtx, resp, callErr) + } + }, nil + }) + } + + usageReq := s.chatRequest(ctx, provider, req, opts) + scopedCtx := scopedRequestContext(ctx, usageReq) + reservation, err := s.Preflight(scopedCtx, usageReq) + if err != nil { + return nil, err + } + resp, err := provider.Chat(scopedCtx, req) + if reservation != nil { + reservation.Reconcile(scopedCtx, resp, err) + } + return resp, err +} + +func (s *Service) chatRequest(ctx context.Context, provider providers.Provider, req providers.ChatRequest, opts ChatOptions) Request { + tenantID := opts.TenantID + if tenantID == uuid.Nil { + tenantID = store.TenantIDFromContext(ctx) + } + if tenantID == uuid.Nil { + tenantID = store.MasterTenantID + } + agentID := opts.AgentID + if agentID == uuid.Nil { + agentID = store.AgentIDFromContext(ctx) + } + providerName := opts.ProviderName + if providerName == "" && provider != nil { + providerName = provider.Name() + } + modelID := opts.ModelID + if modelID == "" { + modelID = req.Model + } + if modelID == "" && provider != nil { + modelID = provider.DefaultModel() + } + return Request{ + TenantID: tenantID, + AgentID: agentID, + ProviderName: providerName, + ModelID: modelID, + ReservationKey: reservationKey(opts), + Messages: req.Messages, + MaxOutputTokens: maxOutputTokens(req, opts.MaxOutputTokens), + } +} + +func scopedRequestContext(ctx context.Context, req Request) context.Context { + if req.TenantID != uuid.Nil && store.TenantIDFromContext(ctx) != req.TenantID { + ctx = store.WithTenantID(ctx, req.TenantID) + } + if req.AgentID != uuid.Nil && store.AgentIDFromContext(ctx) != req.AgentID { + ctx = store.WithAgentID(ctx, req.AgentID) + } + return ctx +} + +func reservationKey(opts ChatOptions) string { + if opts.ReservationKey != "" { + return opts.ReservationKey + } + purpose := opts.Purpose + if purpose == "" { + purpose = "llm" + } + return fmt.Sprintf("%s:%s", purpose, uuid.NewString()) +} + +func maxOutputTokens(req providers.ChatRequest, fallback int) int { + if fallback <= 0 { + fallback = 1024 + } + if req.Options == nil { + return fallback + } + v, ok := req.Options[providers.OptMaxTokens] + if !ok { + return fallback + } + switch n := v.(type) { + case int: + return n + case int64: + return int(n) + case int32: + return int(n) + case float64: + return int(n) + case float32: + return int(n) + default: + return fallback + } +} diff --git a/internal/usage/caps/service.go b/internal/usage/caps/service.go index d3fbbaf3..30061bec 100644 --- a/internal/usage/caps/service.go +++ b/internal/usage/caps/service.go @@ -66,6 +66,7 @@ func (s *Service) Preflight(ctx context.Context, req Request) (*Reservation, err if s == nil || s.store == nil { return skippedReservation(req, "service_disabled"), nil } + ctx = scopedRequestContext(ctx, req) providerData, err := s.resolveProvider(ctx, req.TenantID, req.ProviderName) if err != nil { return skippedReservation(req, "provider_metadata_missing"), nil @@ -107,6 +108,12 @@ func (s *Service) Preflight(ctx context.Context, req Request) (*Reservation, err resolved, err := s.store.ResolvePricing(ctx, req.TenantID, providerData.ID, providerData.Name, providerData.ProviderType, req.ModelID) if err != nil { if errors.Is(err, sql.ErrNoRows) { + _ = s.store.InsertUsageCapEvent(ctx, &store.UsageCapEvent{ + TenantID: req.TenantID, + ReservationKey: key, Decision: store.UsageCapEventBlock, Reason: "pricing_unknown", + EstimatedTokens: usage.TotalTokens(), EstimatedCostMicros: 0, + Metadata: mustJSON(map[string]any{"model_id": req.ModelID, "provider": req.ProviderName}), + }) return blockedReservation(req, scope, key, usage, 0, uuid.Nil, "pricing_unknown"), fmt.Errorf("%w: %s", ErrPricingUnknown, req.ModelID) } return nil, err diff --git a/internal/usage/caps/service_test.go b/internal/usage/caps/service_test.go index 7a14a084..6a597404 100644 --- a/internal/usage/caps/service_test.go +++ b/internal/usage/caps/service_test.go @@ -89,7 +89,7 @@ func TestPreflightIncludesRequestPricingWhenConfigured(t *testing.T) { Name: "openrouter", ProviderType: store.ProviderOpenRouter, APIKey: "sk-test", - }} + }, requireTenant: policy.TenantID} svc := NewService(usageStore, providerStore) _, err := svc.Preflight(context.Background(), Request{ @@ -279,6 +279,79 @@ func TestPreflightTraceMetadataForCapExceeded(t *testing.T) { } } +func TestPreflightRecordsPricingUnknownBlockEvent(t *testing.T) { + policy := store.UsageCapPolicy{ID: uuid.New(), TenantID: uuid.New(), MaxCostMicros: int64Ptr(1000), Enabled: true} + usageStore := &fakeUsageCapStore{ + policies: []store.UsageCapPolicy{policy}, + resolveErr: sql.ErrNoRows, + } + providerStore := &fakeProviderStore{provider: &store.LLMProviderData{ + BaseModel: store.BaseModel{ID: uuid.New()}, + Name: "openrouter", + ProviderType: store.ProviderOpenRouter, + APIKey: "sk-test", + }} + svc := NewService(usageStore, providerStore) + + reservation, err := svc.Preflight(context.Background(), Request{ + TenantID: policy.TenantID, ProviderName: "openrouter", ModelID: "missing/model", + ReservationKey: "pricing-missing", Messages: []providers.Message{{Role: "user", Content: "hello"}}, + MaxOutputTokens: 10, + }) + if !errors.Is(err, ErrPricingUnknown) { + t.Fatalf("Preflight error = %v, want ErrPricingUnknown", err) + } + if reservation == nil { + t.Fatal("Preflight returned nil reservation") + } + metadata := reservation.TraceMetadata() + if metadata.Reason != "pricing_unknown" { + t.Fatalf("reservation metadata = %+v, want pricing_unknown", metadata) + } + if len(usageStore.events) != 1 { + t.Fatalf("events = %d, want 1", len(usageStore.events)) + } + event := usageStore.events[0] + if event.Decision != store.UsageCapEventBlock || event.Reason != "pricing_unknown" { + t.Fatalf("event decision/reason = %q/%q, want block/pricing_unknown", event.Decision, event.Reason) + } + if event.ReservationKey != "pricing-missing" { + t.Fatalf("event reservation_key = %q, want pricing-missing", event.ReservationKey) + } +} + +func TestServiceChatBlocksBeforeProviderCall(t *testing.T) { + policy := store.UsageCapPolicy{ID: uuid.New(), TenantID: uuid.New(), MaxTokens: int64Ptr(10), Enabled: true} + usageStore := &fakeUsageCapStore{ + policies: []store.UsageCapPolicy{policy}, + reserveErr: &store.UsageCapExceededError{PolicyID: policy.ID, Reason: "token_cap_exceeded"}, + } + providerStore := &fakeProviderStore{provider: &store.LLMProviderData{ + BaseModel: store.BaseModel{ID: uuid.New()}, + Name: "openrouter", + ProviderType: store.ProviderOpenRouter, + APIKey: "sk-test", + }, requireTenant: policy.TenantID} + svc := NewService(usageStore, providerStore) + provider := &fakeChatProvider{name: "openrouter", model: "token/model"} + + _, err := svc.Chat(context.Background(), provider, providers.ChatRequest{ + Messages: []providers.Message{{Role: "user", Content: "hello"}}, + Model: "token/model", + Options: map[string]any{providers.OptMaxTokens: 20}, + }, ChatOptions{ + TenantID: policy.TenantID, + ProviderName: "openrouter", + Purpose: "test-block", + }) + if !errors.Is(err, ErrCapExceeded) { + t.Fatalf("Chat error = %v, want ErrCapExceeded", err) + } + if provider.calls != 0 { + t.Fatalf("provider calls = %d, want 0", provider.calls) + } +} + func TestMergeTraceMetadataPreservesExistingSections(t *testing.T) { existing := json.RawMessage(`{"thinking":{"effort":"high"}}`) merged := MergeTraceMetadata(existing, []TraceMetadata{{ @@ -319,6 +392,32 @@ func TestCountImagesOnlyCountsImageMIMEs(t *testing.T) { } } +type fakeChatProvider struct { + name string + model string + calls int + resp *providers.ChatResponse + err error +} + +func (p *fakeChatProvider) Chat(context.Context, providers.ChatRequest) (*providers.ChatResponse, error) { + p.calls++ + if p.resp != nil || p.err != nil { + return p.resp, p.err + } + return &providers.ChatResponse{ + Content: "ok", + Usage: &providers.Usage{PromptTokens: 1, CompletionTokens: 1, TotalTokens: 2}, + }, nil +} + +func (p *fakeChatProvider) ChatStream(ctx context.Context, req providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) { + return p.Chat(ctx, req) +} + +func (p *fakeChatProvider) DefaultModel() string { return p.model } +func (p *fakeChatProvider) Name() string { return p.name } + type fakeUsageCapStore struct { policies []store.UsageCapPolicy resolved *store.ResolvedUsagePricing @@ -329,6 +428,7 @@ type fakeUsageCapStore struct { reconciled store.UsageReconcileRequest reconcileCalls int reconcileCtxCanceled bool + events []store.UsageCapEvent } func (s *fakeUsageCapStore) UpsertPricingCatalog(context.Context, []store.UsagePricingCatalogEntry) (int, error) { @@ -384,13 +484,17 @@ func (s *fakeUsageCapStore) ListUsageCapUtilization(context.Context, uuid.UUID) func (s *fakeUsageCapStore) ListUsageCapEvents(context.Context, uuid.UUID, int) ([]store.UsageCapEvent, error) { return nil, nil } -func (s *fakeUsageCapStore) InsertUsageCapEvent(context.Context, *store.UsageCapEvent) error { +func (s *fakeUsageCapStore) InsertUsageCapEvent(_ context.Context, event *store.UsageCapEvent) error { + if event != nil { + s.events = append(s.events, *event) + } return nil } type fakeProviderStore struct { provider *store.LLMProviderData masterProvider *store.LLMProviderData + requireTenant uuid.UUID } func (s *fakeProviderStore) CreateProvider(context.Context, *store.LLMProviderData) error { return nil } @@ -398,6 +502,9 @@ func (s *fakeProviderStore) GetProvider(context.Context, uuid.UUID) (*store.LLMP return s.provider, nil } func (s *fakeProviderStore) GetProviderByName(ctx context.Context, _ string) (*store.LLMProviderData, error) { + if s.requireTenant != uuid.Nil && store.TenantIDFromContext(ctx) != s.requireTenant { + return nil, sql.ErrNoRows + } if store.TenantIDFromContext(ctx) == store.MasterTenantID && s.masterProvider != nil { return s.masterProvider, nil } diff --git a/internal/vault/enrich_classify.go b/internal/vault/enrich_classify.go index b8aaf857..227a27cc 100644 --- a/internal/vault/enrich_classify.go +++ b/internal/vault/enrich_classify.go @@ -7,6 +7,7 @@ import ( "maps" "slices" + "github.com/google/uuid" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" ) @@ -40,6 +41,14 @@ func (w *EnrichWorker) classifyLinks(ctx context.Context, provider providers.Pro if provider == nil { return } + if tid, err := uuid.Parse(tenantID); err == nil { + ctx = store.WithTenantID(ctx, tid) + } + if agentID != "" { + if aid, err := uuid.Parse(agentID); err == nil { + ctx = store.WithAgentID(ctx, aid) + } + } capped := results if len(capped) > classifyMaxSourceDocs { diff --git a/internal/vault/enrich_worker.go b/internal/vault/enrich_worker.go index 18d9f49e..2fd4cb1c 100644 --- a/internal/vault/enrich_worker.go +++ b/internal/vault/enrich_worker.go @@ -16,20 +16,21 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/bgalert" "github.com/nextlevelbuilder/goclaw/internal/bus" "github.com/nextlevelbuilder/goclaw/internal/eventbus" - "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/providerresolve" + "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "golang.org/x/sync/semaphore" ) const ( - enrichMaxDedupEntries = 10000 - enrichSimilarityLimit = 10 - enrichSimilarityMin = 0.7 - enrichMaxConcurrent = 3 // max concurrent batch summarize calls across chunks - enrichBatchSize = 5 // docs per enrichment chunk (1 LLM call per chunk) - enrichBatchItemMaxRunes = 3000 // per-file content limit in batch summarize - enrichMaxRetries = 3 // shared retry count for LLM calls (summarize + classify) + enrichMaxDedupEntries = 10000 + enrichSimilarityLimit = 10 + enrichSimilarityMin = 0.7 + enrichMaxConcurrent = 3 // max concurrent batch summarize calls across chunks + enrichBatchSize = 5 // docs per enrichment chunk (1 LLM call per chunk) + enrichBatchItemMaxRunes = 3000 // per-file content limit in batch summarize + enrichMaxRetries = 3 // shared retry count for LLM calls (summarize + classify) ) // Shared retry config for all enrichment LLM calls. @@ -41,12 +42,13 @@ var ( // EnrichWorkerDeps bundles dependencies for the vault enrichment worker. type EnrichWorkerDeps struct { VaultStore store.VaultStore - SystemConfigs store.SystemConfigStore // per-tenant provider config - Registry *providers.Registry // provider resolution + SystemConfigs store.SystemConfigStore // per-tenant provider config + Registry *providers.Registry // provider resolution EventBus eventbus.DomainEventBus - MsgBus bus.EventPublisher // for WS event broadcast - TeamStore store.TaskCommentStore // for Phase 2.5 task-based auto-linking (nil-safe) - AlertDeps bgalert.AlertDeps // for reporting non-retryable LLM errors + MsgBus bus.EventPublisher // for WS event broadcast + TeamStore store.TaskCommentStore // for Phase 2.5 task-based auto-linking (nil-safe) + AlertDeps bgalert.AlertDeps // for reporting non-retryable LLM errors + UsageCaps *usagecaps.Service } // RegisterEnrichWorker subscribes the enrichment worker to vault doc events. @@ -60,6 +62,7 @@ func RegisterEnrichWorker(deps EnrichWorkerDeps) (func(), *EnrichProgress, *Enri registry: deps.Registry, msgBus: deps.MsgBus, alertDeps: deps.AlertDeps, + usageCaps: deps.UsageCaps, dedup: make(map[string]string), sem: semaphore.NewWeighted(enrichMaxConcurrent), progress: progress, @@ -74,11 +77,12 @@ func RegisterEnrichWorker(deps EnrichWorkerDeps) (func(), *EnrichProgress, *Enri // Exported so HTTP handlers can call Stop/EnqueueUnenriched. type EnrichWorker struct { vault store.VaultStore - teamStore store.TaskCommentStore // nil-tolerant — Phase 2.5 disabled when nil - systemConfigs store.SystemConfigStore // per-tenant provider config - registry *providers.Registry // provider resolution - msgBus bus.EventPublisher // for error event broadcast - alertDeps bgalert.AlertDeps // for reporting non-retryable LLM errors + teamStore store.TaskCommentStore // nil-tolerant — Phase 2.5 disabled when nil + systemConfigs store.SystemConfigStore // per-tenant provider config + registry *providers.Registry // provider resolution + msgBus bus.EventPublisher // for error event broadcast + alertDeps bgalert.AlertDeps // for reporting non-retryable LLM errors + usageCaps *usagecaps.Service queue enrichBatchQueue progress *EnrichProgress @@ -296,6 +300,14 @@ func (w *EnrichWorker) processChunk(ctx context.Context, items []eventbus.VaultD // Batch-fetch all existing docs in a single query. tenantID := pending[0].TenantID + if tid, err := uuid.Parse(tenantID); err == nil { + ctx = store.WithTenantID(ctx, tid) + } + if pending[0].AgentID != "" { + if aid, err := uuid.Parse(pending[0].AgentID); err == nil { + ctx = store.WithAgentID(ctx, aid) + } + } // Resolve provider once per chunk (all items share tenantID) provider, model := w.resolveProviderForTenant(ctx, tenantID) @@ -515,7 +527,11 @@ func (w *EnrichWorker) chatWithRetry(ctx context.Context, provider providers.Pro } } cctx, cancel := context.WithTimeout(ctx, enrichRetryTimeouts[attempt]) - resp, err := provider.Chat(cctx, req) + resp, err := w.usageCaps.Chat(cctx, provider, req, usagecaps.ChatOptions{ + ModelID: req.Model, + Purpose: logPrefix, + MaxOutputTokens: 4096, + }) cancel() if err != nil { lastErr = err @@ -563,7 +579,6 @@ func (w *EnrichWorker) syncWikilinks(ctx context.Context, p eventbus.VaultDocUps } } - // recordDedup stores a processed hash and evicts ~25% entries if over capacity. func (w *EnrichWorker) recordDedup(docID, hash string) { w.dedupMu.Lock()