diff --git a/apps/server/AGENTS.md b/apps/server/AGENTS.md index 16dd68e8..23af31ca 100644 --- a/apps/server/AGENTS.md +++ b/apps/server/AGENTS.md @@ -11,13 +11,14 @@ All sources under `apps/server/src/`. ## Local Contracts ### HTTP REST routes -- `GET/POST /sessions` — session browse and start +- `GET/POST /sessions` — session browse and start (`POST` accepts an optional `intent` brief, parity with the WS `StartSession`) - `POST /sessions/{id}/resume` — resume a session after restart - `GET/POST /tasks` — task listing and management - `GET/POST /providers` — provider configuration - `GET/POST /workflows` — workflow management - `GET /health` — health report (probes: event-store, llama-server, disk watermark) - `GET /stats` — metrics report (MetricsProjection) +- `GET /metrics/tool-reliability` — per-model tool-call validity across the event log (`ToolReliabilityInspectionService`); groundwork for capability-aware routing ### WebSocket protocol (`/ws`) - **ServerMessage** (server → client): sealed hierarchy — `SessionMessage` (event-derived, carries `sequence` + `sessionSequence`) and `NonEventMessage` (control/infra). Variants include session lifecycle, approval requests, clarification requests, narration, proposed workflows, health/metrics pushes. diff --git a/apps/server/build.gradle b/apps/server/build.gradle index 95399e63..57fd4701 100644 --- a/apps/server/build.gradle +++ b/apps/server/build.gradle @@ -36,6 +36,7 @@ dependencies { implementation project(':infrastructure:persistence') implementation project(':infrastructure:inference') implementation project(':infrastructure:inference:llama_cpp') + implementation project(':infrastructure:inference:openai_compat') implementation project(':infrastructure:inference:commons') implementation project(':core:router') implementation project(':core:tools') diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/Application.kt b/apps/server/src/main/kotlin/com/correx/apps/server/Application.kt index 2bdee935..4355f037 100644 --- a/apps/server/src/main/kotlin/com/correx/apps/server/Application.kt +++ b/apps/server/src/main/kotlin/com/correx/apps/server/Application.kt @@ -1,6 +1,7 @@ package com.correx.apps.server import com.correx.apps.server.health.HealthInspectionService +import com.correx.apps.server.metrics.ToolReliabilityInspectionService import com.correx.apps.server.routes.providerRoutes import com.correx.apps.server.routes.sessionRoutes import com.correx.apps.server.routes.taskRoutes @@ -38,6 +39,7 @@ fun Application.configureServer(module: ServerModule) { val globalStreamHandler = GlobalStreamHandler(module) val healthInspection = HealthInspectionService(module.eventStore) + val toolReliability = ToolReliabilityInspectionService(module.eventStore) routing { get("/health") { @@ -51,6 +53,12 @@ fun Application.configureServer(module: ServerModule) { call.respond(healthInspection.inspect()) } + // Per-model tool-call reliability across the whole event log (valid vs. failed tool calls + // per provider). Measurement groundwork for capability-aware routing of tool-heavy stages. + get("/metrics/tool-reliability") { + call.respond(toolReliability.inspect()) + } + webSocket("/stream") { globalStreamHandler.handle(this) } diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/Main.kt b/apps/server/src/main/kotlin/com/correx/apps/server/Main.kt index 1b96d0b8..9b6c7738 100644 --- a/apps/server/src/main/kotlin/com/correx/apps/server/Main.kt +++ b/apps/server/src/main/kotlin/com/correx/apps/server/Main.kt @@ -76,7 +76,7 @@ import com.correx.core.events.types.SessionId import com.correx.infrastructure.InfrastructureModule import com.correx.infrastructure.inference.DefaultProviderRegistry import com.correx.infrastructure.workflow.ExecutionPlanCompiler -import com.correx.infrastructure.inference.FirstAvailableRoutingStrategy +import com.correx.infrastructure.inference.CapabilityAwareRoutingStrategy import com.correx.infrastructure.inference.commons.ManagedInferenceRouter import com.correx.infrastructure.inference.commons.AmdResourceProbe import com.correx.infrastructure.inference.commons.NvidiaResourceProbe @@ -84,7 +84,6 @@ import com.correx.infrastructure.inference.commons.ResourceProbe import com.correx.infrastructure.inference.commons.SystemResourceProbe import com.correx.infrastructure.inference.commons.UnavailableProbe import com.correx.core.inference.InferenceProvider -import com.correx.infrastructure.inference.llama.cpp.LlamaCppInferenceProvider import com.correx.infrastructure.tools.DispatchingToolExecutor import com.correx.infrastructure.tools.task.TaskTools import com.correx.infrastructure.tools.FileEditConfig @@ -148,12 +147,9 @@ fun main() { firstProvider = managedProvider infraRegistry = InfrastructureModule.createProviderRegistry(listOf(managedProvider)) // Also register any static providers from [[providers]] (non-managed remotes) - correxConfig.providers.mapNotNull { providerConfig -> - when (providerConfig.type.lowercase()) { - "llamacpp" -> buildProviderFromConfig(providerConfig) - else -> { log.error("Unknown provider type: {}", providerConfig.type); null } - } - }.forEach { infraRegistry.register(it) } + correxConfig.providers + .mapNotNull { providerConfig -> buildProviderFromConfig(providerConfig) } + .forEach { infraRegistry.register(it) } // Managed router owns per-stage model selection: it ensures the resolved model is resident // before routing (stage.modelId > capability match > default), evicting the prior model. val managedRouter = ManagedInferenceRouter(modelManager, descriptors, targetModelConfig.id) @@ -174,7 +170,7 @@ fun main() { require(staticProviders.isNotEmpty()) { "At least one provider must be configured" } firstProvider = staticProviders.first() infraRegistry = InfrastructureModule.createProviderRegistry(staticProviders) - inferenceRouter = DefaultInferenceRouter(infraRegistry, FirstAvailableRoutingStrategy()) + inferenceRouter = DefaultInferenceRouter(infraRegistry, CapabilityAwareRoutingStrategy()) llamaBaseUrl = correxConfig.providers.firstOrNull()?.url ?: System.getenv("CORREX_LLAMA_URL") ?: "http://127.0.0.1:10000" @@ -566,7 +562,9 @@ fun main() { module.start() log.info("==============================") - embeddedServer(Netty, port = 8080) { configureServer(module) }.start(wait = true) + embeddedServer(Netty, port = correxConfig.server.port, host = correxConfig.server.host) { + configureServer(module) + }.start(wait = true) } fun renderProjectProfileText(profile: com.correx.core.sessions.BoundProjectProfile?): String? { @@ -651,18 +649,10 @@ private fun resolveConfigRelativePath(raw: String, configDir: Path?): Path { } } -private fun buildProviders(config: CorrexConfig): List { +private fun buildProviders(config: CorrexConfig): List { return if (config.providers.isNotEmpty()) { log.info("Loading {} provider(s) from config", config.providers.size) - config.providers.mapNotNull { providerConfig -> - when (providerConfig.type.lowercase()) { - "llamacpp" -> buildProviderFromConfig(providerConfig) - else -> { - log.error("Unknown provider type: {}", providerConfig.type) - null - } - } - } + config.providers.mapNotNull { providerConfig -> buildProviderFromConfig(providerConfig) } } else { log.info("No providers in config; using env var fallback") listOf( @@ -675,14 +665,48 @@ private fun buildProviders(config: CorrexConfig): List InfrastructureModule.createLlamaCppProvider( + modelId = config.modelId, + modelPath = config.modelPath, + baseUrl = config.url, + capabilities = capabilities, + ) + + "nim", "openai" -> { + val apiKey = resolveApiKey(config) + if (apiKey.isBlank()) { + log.error( + "Provider '{}' (type {}) has no API key (set api_key or api_key_env); skipping", + config.id, config.type, + ) + null + } else { + InfrastructureModule.createOpenAiCompatProvider( + modelId = config.modelId, + baseUrl = config.url, + apiKey = apiKey, + idPrefix = config.type.lowercase(), + capabilities = capabilities, + ) + } + } + + else -> { + log.error("Unknown provider type: {}", config.type) + null + } + } +} + +// Prefer the literal api_key; otherwise read the named env var. Returns "" when neither resolves. +private fun resolveApiKey(config: ProviderConfig): String = when { + config.apiKey.isNotBlank() -> config.apiKey + config.apiKeyEnv.isNotBlank() -> System.getenv(config.apiKeyEnv).orEmpty() + else -> "" } private fun parseCapabilitiesFromConfig(capsMap: Map): Set { @@ -740,6 +764,7 @@ private fun buildToolConfig( fileRead = FileReadConfig( enabled = toolsConfig.fileReadEnabled, allowedPaths = allowed, + workingDir = workingDir, ), fileWrite = FileWriteConfig( allowedPaths = allowed, diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/metrics/ToolReliabilityInspectionService.kt b/apps/server/src/main/kotlin/com/correx/apps/server/metrics/ToolReliabilityInspectionService.kt new file mode 100644 index 00000000..9448c37b --- /dev/null +++ b/apps/server/src/main/kotlin/com/correx/apps/server/metrics/ToolReliabilityInspectionService.kt @@ -0,0 +1,107 @@ +package com.correx.apps.server.metrics + +import com.correx.core.events.events.InferenceCompletedEvent +import com.correx.core.events.events.ToolExecutionCompletedEvent +import com.correx.core.events.events.ToolExecutionFailedEvent +import com.correx.core.events.stores.EventStore +import kotlinx.serialization.Serializable + +private const val PERCENT = 100.0 +private const val REASON_BUCKET_MAX = 80 +private const val TOP_REASONS = 5 + +@Serializable +data class ReasonCount(val reason: String, val count: Long) + +@Serializable +data class ToolReliabilityRow( + val toolName: String, + val calls: Long, + val failures: Long, + val validityPct: Double, +) + +@Serializable +data class ProviderToolReliability( + val provider: String, + val calls: Long, + val failures: Long, + val validityPct: Double, + val byTool: List, + val topFailureReasons: List, +) + +@Serializable +data class ToolReliabilityReport(val providers: List) + +/** + * Per-model tool-call reliability, derived by replaying the whole event log: for each provider, how + * often did the tool calls it produced actually execute vs. fail (a failed `ToolExecutionFailedEvent` + * is the deterministic signal of a malformed/invalid call — wrong args, missing required field, …). + * + * Tool events don't carry a providerId, so each tool outcome is attributed to the provider of the + * most recent `InferenceCompletedEvent` in the same session (one model drives a session's inference + * under the default routing). This is the measurement groundwork for capability-aware routing — + * steering tool-heavy stages to models with a high observed validity rate. + */ +class ToolReliabilityInspectionService(private val eventStore: EventStore) { + + private class Agg { + var calls: Long = 0 + var failures: Long = 0 + val reasons: MutableMap = linkedMapOf() + } + + fun inspect(): ToolReliabilityReport { + val sessionProvider = mutableMapOf() + val byProvider = linkedMapOf() + val byProviderTool = linkedMapOf, Agg>() + + fun record(provider: String, tool: String, failureReason: String?) { + val p = byProvider.getOrPut(provider) { Agg() } + val pt = byProviderTool.getOrPut(provider to tool) { Agg() } + p.calls++; pt.calls++ + if (failureReason != null) { + p.failures++; pt.failures++ + val bucket = failureReason.lineSequence().firstOrNull().orEmpty().take(REASON_BUCKET_MAX) + p.reasons[bucket] = (p.reasons[bucket] ?: 0) + 1 + } + } + + eventStore.allEvents().forEach { stored -> + when (val payload = stored.payload) { + is InferenceCompletedEvent -> sessionProvider[payload.sessionId.value] = payload.providerId.value + is ToolExecutionCompletedEvent -> + record(sessionProvider[payload.sessionId.value] ?: "unknown", payload.toolName, null) + is ToolExecutionFailedEvent -> + record(sessionProvider[payload.sessionId.value] ?: "unknown", payload.toolName, payload.reason) + else -> Unit + } + } + + val providers = byProvider.entries + .sortedByDescending { it.value.calls } + .map { (provider, agg) -> + ProviderToolReliability( + provider = provider, + calls = agg.calls, + failures = agg.failures, + validityPct = validity(agg.calls, agg.failures), + byTool = byProviderTool.entries + .filter { it.key.first == provider } + .sortedByDescending { it.value.calls } + .map { (key, a) -> + ToolReliabilityRow(key.second, a.calls, a.failures, validity(a.calls, a.failures)) + }, + topFailureReasons = agg.reasons.entries + .sortedByDescending { it.value } + .take(TOP_REASONS) + .map { ReasonCount(it.key, it.value) }, + ) + } + return ToolReliabilityReport(providers) + } + + private fun validity(calls: Long, failures: Long): Double = + if (calls == 0L) PERCENT else (calls - failures).toDouble() / calls * PERCENT +} diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/routes/SessionRoutes.kt b/apps/server/src/main/kotlin/com/correx/apps/server/routes/SessionRoutes.kt index 08dbdf4d..18277694 100644 --- a/apps/server/src/main/kotlin/com/correx/apps/server/routes/SessionRoutes.kt +++ b/apps/server/src/main/kotlin/com/correx/apps/server/routes/SessionRoutes.kt @@ -9,6 +9,7 @@ import com.correx.apps.server.replay.ReplayInspectionService import com.correx.apps.server.serialization.payloadDiscriminator import com.correx.apps.server.ws.SessionStreamHandler import com.correx.core.events.events.EventPayload +import com.correx.core.events.events.InitialIntentEvent import com.correx.core.events.serialization.eventJson import com.correx.core.events.types.SessionId import com.correx.core.utils.TypeId @@ -39,7 +40,13 @@ data class EventRow( ) @Serializable -data class StartSessionRequest(val workflowId: String, val config: SessionConfigDto? = null) +data class StartSessionRequest( + val workflowId: String, + val config: SessionConfigDto? = null, + // Optional freeform brief for intent-driven workflows (REST parity with the WS StartSession). + // Seeded as an InitialIntentEvent before the run so it lands in the decision journal. + val intent: String? = null, +) @Serializable data class SessionSummaryResponse( @@ -108,6 +115,9 @@ private fun Route.startSessionRoute(module: ServerModule) { val graph = module.workflowRegistry.find(body.workflowId) ?: return@post call.respond(HttpStatusCode.BadRequest, "Unknown workflowId: ${body.workflowId}") val sessionId: SessionId = TypeId(UUID.randomUUID().toString()) + body.intent?.takeIf { it.isNotBlank() }?.let { intent -> + EventDispatcher(module.eventStore).emit(InitialIntentEvent(sessionId, intent), sessionId) + } module.launchSessionRun(sessionId, graph) call.respond(HttpStatusCode.Accepted, StartSessionResponse(sessionId.value)) } diff --git a/core/config/src/main/kotlin/com/correx/core/config/ConfigLoader.kt b/core/config/src/main/kotlin/com/correx/core/config/ConfigLoader.kt index 1dd5451c..69e4fc0a 100644 --- a/core/config/src/main/kotlin/com/correx/core/config/ConfigLoader.kt +++ b/core/config/src/main/kotlin/com/correx/core/config/ConfigLoader.kt @@ -537,6 +537,8 @@ object ConfigLoader { modelPath = modelPath, url = url, capabilities = capabilities, + apiKey = asString(providerMap["api_key"], ""), + apiKeyEnv = asString(providerMap["api_key_env"], ""), ) } diff --git a/core/config/src/main/kotlin/com/correx/core/config/CorrexConfig.kt b/core/config/src/main/kotlin/com/correx/core/config/CorrexConfig.kt index 30b8206e..2c944050 100644 --- a/core/config/src/main/kotlin/com/correx/core/config/CorrexConfig.kt +++ b/core/config/src/main/kotlin/com/correx/core/config/CorrexConfig.kt @@ -173,6 +173,11 @@ data class ProviderConfig( val modelPath: String = "", val url: String = "http://127.0.0.1:10000", val capabilities: Map = emptyMap(), + // Bearer credentials for remote providers (e.g. NVIDIA NIM). `apiKey` is the literal token; + // `apiKeyEnv` names an env var to read it from at startup (preferred — keeps the secret out + // of config files). Both are blank/unused for local providers like llamacpp. + val apiKey: String = "", + val apiKeyEnv: String = "", ) @Serializable diff --git a/core/config/src/main/kotlin/com/correx/core/config/CorrexConfigWriter.kt b/core/config/src/main/kotlin/com/correx/core/config/CorrexConfigWriter.kt index 69b3da68..8744aeab 100644 --- a/core/config/src/main/kotlin/com/correx/core/config/CorrexConfigWriter.kt +++ b/core/config/src/main/kotlin/com/correx/core/config/CorrexConfigWriter.kt @@ -113,6 +113,8 @@ object CorrexConfigWriter { b.kv("model_path", str(p.modelPath)) b.kv("url", str(p.url)) if (p.capabilities.isNotEmpty()) b.kv("capabilities", caps(p.capabilities)) + if (p.apiKey.isNotEmpty()) b.kv("api_key", str(p.apiKey)) + if (p.apiKeyEnv.isNotEmpty()) b.kv("api_key_env", str(p.apiKeyEnv)) } cfg.models.forEach { m -> diff --git a/core/kernel/AGENTS.md b/core/kernel/AGENTS.md index e77d07eb..7bba11ce 100644 --- a/core/kernel/AGENTS.md +++ b/core/kernel/AGENTS.md @@ -13,7 +13,9 @@ CORREX kernel team. This is the integration point for all other `core/` modules. - `SessionOrchestrator` / `DefaultSessionOrchestrator` — primary entry point for launching and advancing sessions through workflow stages. - `OrchestrationState` / `OrchestrationReducer` (`DefaultOrchestrationReducer`) / `OrchestrationProjector` / `OrchestrationRepository` — standard event-sourcing stack for orchestration state. - `RetryCoordinator` / `DefaultRetryCoordinator` — manages retry logic per `RetryPolicy`. -- `ApprovalGateway` — kernel-side approval bridge; calls `core:approvals` engine before executing gated operations. +- On a recoverable tool failure `dispatchToolCalls` feeds the failing tool's argument schema back into context alongside the error (`toolArgsHint`), so the model self-corrects a malformed call instead of repeating it — the contract stays strict; the feedback is what loosens. +- Stages that grant tools (`allowedTools` non-empty) request `ModelCapability.ToolCalling` on top of their declared capabilities when routing, so the capability-aware strategy steers them to the best tool-calling model. +- `ApprovalGateway` — kernel-side approval bridge; calls `core:approvals` engine before executing gated operations. Per-tool gating in `dispatchToolCalls` builds the `ApprovalContext` mode from the session's bound operator profile (`boundProfile.approvalMode`, mapped by `approvalModeFor`): unset/`prompt` keeps a human in the loop (default), `auto` auto-approves up to T2, `yolo` all tiers, `deny` blocks above T0. The engine is always consulted (Invariant #4 holds); policy/plane-2 BLOCK stays terminal regardless of mode. - `ReplayOrchestrator` / `ReplayInferenceProvider` / `ReplayStrategy` — deterministic replay of a session from its event log. `ReplayInferenceProvider` returns recorded responses — no live LLM (Hard Invariant #8). - `SubagentRunner` / `InSessionSubagentRunner` — runs sub-agent invocations within an active session. - `StaticAnalysisRunner` / `ProcessStaticAnalysisRunner` — runs static analysis tools and records results as events. diff --git a/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt b/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt index dd49f549..43edf139 100644 --- a/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt +++ b/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt @@ -106,6 +106,7 @@ import com.correx.core.inference.ResponseFormat import com.correx.core.inference.Tokenizer import com.correx.core.inference.ToolCallRequest import com.correx.core.inference.ToolDefinition +import com.correx.core.inference.ModelCapability import com.correx.core.sessions.ApprovalMode import com.correx.core.inference.ToolFunction import com.correx.core.kernel.execution.WorkflowResult @@ -561,7 +562,8 @@ abstract class SessionOrchestrator( continue } val toolEntries = dispatchToolCalls(sessionId, stageId, - inferenceResult.response.toolCalls, stageConfig, effectives) + inferenceResult.response.toolCalls, stageConfig, effectives, + approvalModeFor(session.state.boundProfile?.approvalMode)) val fatalEntry = toolEntries.firstOrNull { it.content.startsWith("FATAL:") } if (fatalEntry != null) { emitProcessResultEvents(sessionId, stageId, stageConfig) @@ -654,12 +656,33 @@ abstract class SessionOrchestrator( } @Suppress("CyclomaticComplexMethod") + // Maps the operator profile's free-text approval_mode onto the engine's [ApprovalMode]. + // Unset/unknown falls back to PROMPT so an absent profile keeps the human-in-the-loop + // default; only an explicit auto/yolo opts a session into unattended approval. + private fun approvalModeFor(profileMode: String?): ApprovalMode = + when (profileMode?.trim()?.lowercase()) { + "deny" -> ApprovalMode.DENY + "auto" -> ApprovalMode.AUTO + "yolo" -> ApprovalMode.YOLO + else -> ApprovalMode.PROMPT + } + + // On a recoverable tool failure, append the tool's argument schema to the error fed back to the + // model so it can self-correct its next attempt (it already sees its own malformed call in the + // assistant entry above) rather than repeating the same mistake. Kept compact to bound context. + private fun toolArgsHint(tool: Tool?): String = + tool?.let { + "\nThe '${it.name}' tool requires arguments matching this JSON schema — re-issue the " + + "call with corrected arguments: ${it.parametersSchema}" + }.orEmpty() + private suspend fun dispatchToolCalls( sessionId: SessionId, stageId: StageId, toolCalls: List, stageConfig: StageConfig, effectives: RunEffectives, + approvalMode: ApprovalMode, ): List { val executor = effectives.executor ?: return emptyList() val processResultSlots = stageConfig.produces.filter { it.kind.id == "process_result" } @@ -793,7 +816,7 @@ abstract class SessionOrchestrator( val projectId = effectives.policy?.workspaceRoot?.let { ProjectIdentity.of(it.toString()) } val approvalCtx = ApprovalContext( identity = ApprovalScopeIdentity(sessionId, stageId, projectId = projectId), - mode = ApprovalMode.PROMPT, + mode = approvalMode, ) val requestId = ApprovalRequestId(UUID.randomUUID().toString()) val toolPreview = computeToolPreview(toolCall.function.name, parameters) @@ -968,7 +991,7 @@ abstract class SessionOrchestrator( ?: result.output is ToolResult.Failure -> { if (!result.recoverable) "FATAL: ${result.reason}" - else "ERROR: ${result.reason}" + else "ERROR: ${result.reason}${toolArgsHint(tool)}" } } val resultEntry = ContextEntry( @@ -1772,7 +1795,12 @@ abstract class SessionOrchestrator( responseFormat: ResponseFormat = ResponseFormat.Text, effectives: RunEffectives = RunEffectives(toolRegistry, toolExecutor, workspacePolicy), ): InferenceResult { - val provider = inferenceRouter.route(stageId, stageConfig.requiredCapabilities, stageConfig.modelId) + // A stage that grants tools needs a tool-calling model, so request ToolCalling on top of any + // declared capabilities — the capability-aware strategy then ranks eligible providers by their + // ToolCalling score and routes the stage to the best tool-caller. + val requiredCapabilities = stageConfig.requiredCapabilities + + if (stageConfig.allowedTools.isNotEmpty()) setOf(ModelCapability.ToolCalling) else emptySet() + val provider = inferenceRouter.route(stageId, requiredCapabilities, stageConfig.modelId) log.debug( "[Orchestrator] inference session={} stage={} provider={} timeoutMs={}", sessionId.value, stageId.value, provider.id.value, timeoutMs, diff --git a/infrastructure/build.gradle b/infrastructure/build.gradle index 77e47c5a..d8558551 100644 --- a/infrastructure/build.gradle +++ b/infrastructure/build.gradle @@ -15,6 +15,7 @@ dependencies { implementation project(":infrastructure:inference") implementation project(":infrastructure:inference:commons") implementation project(":infrastructure:inference:llama_cpp") + implementation project(":infrastructure:inference:openai_compat") implementation project(":infrastructure:persistence") implementation project(":infrastructure:router:turbovec") implementation project(":infrastructure:tools") diff --git a/infrastructure/inference/AGENTS.md b/infrastructure/inference/AGENTS.md index 7c1229e5..c1317836 100644 --- a/infrastructure/inference/AGENTS.md +++ b/infrastructure/inference/AGENTS.md @@ -2,7 +2,7 @@ ## Purpose -Inference adapter layer. The root module provides `DefaultProviderRegistry` (registers `InferenceProvider` instances) and `FirstAvailableRoutingStrategy`. Submodules cover shared HTTP client infrastructure (`commons/`) and the llama.cpp server adapter (`llama_cpp/`). +Inference adapter layer. The root module provides `DefaultProviderRegistry` (registers `InferenceProvider` instances) and `FirstAvailableRoutingStrategy`. Submodules cover shared HTTP client infrastructure (`commons/`), the llama.cpp server adapter (`llama_cpp/`), and the remote OpenAI-compatible adapter (`openai_compat/`, e.g. NVIDIA NIM). ## Ownership @@ -11,10 +11,11 @@ Adapter for LLM inference backends. Implements `core:inference` interfaces. No d ## Local Contracts - `DefaultProviderRegistry` implements `ProviderRegistry` from `core:inference`. -- `FirstAvailableRoutingStrategy` is the default routing policy; extend only in `core:inference`, not here. +- Two `RoutingStrategy` implementations (both hard-filter to providers declaring every required capability): `FirstAvailableRoutingStrategy` picks the first match; `CapabilityAwareRoutingStrategy` ranks the matches by summed required-capability score (ties keep list order, so it's a superset of first-available). The server wires `CapabilityAwareRoutingStrategy` so tool-heavy stages (which request `ModelCapability.ToolCalling`) route to the best tool-caller. Extend the `RoutingStrategy` interface only in `core:inference`, not here. - All LLM responses are proposals — they must be validated by the core before affecting state (invariant #7). Adapters return raw responses; they do not validate. -- Network calls to the llama.cpp server are environment observations; results must be recorded as events by callers if replay must reproduce them (invariant #9). +- Network calls to inference backends are environment observations; results must be recorded as events by callers if replay must reproduce them (invariant #9). - Model lifecycle (spawn/own the llama-server process) lives in `llama_cpp/`; the autonomous scheduler is intentionally unbuilt (see root CLAUDE.md). +- `openai_compat/` is fully remote (no local process/GPU). It speaks `POST {baseUrl}/chat/completions` with `Authorization: Bearer`; `baseUrl` must include the version segment (e.g. `/v1`). It has no `/tokenize`, so it uses a heuristic tokenizer, and no GBNF — JSON artifacts rely on the core's validate-after-retry. Server dispatch keys `[[providers]] type = "nim" | "openai"`; the key comes from `api_key` or `api_key_env`. ## Work Guidance @@ -26,9 +27,11 @@ Standard adapter rules apply (see parent `AGENTS.md`). HTTP client code uses Kto ./gradlew :infrastructure:inference:test --rerun-tasks ./gradlew :infrastructure:inference:commons:test --rerun-tasks ./gradlew :infrastructure:inference:llama_cpp:test --rerun-tasks +./gradlew :infrastructure:inference:openai_compat:test --rerun-tasks ``` ## Child DOX Index - `commons/` — shared `ManagedInferenceProvider`, `ModelManager`, `ResourceProbe` (Nvidia/AMD), `ResidencyMode`; no separate AGENTS.md (sub-leaf, covered by this doc) - `llama_cpp/` — `LlamaCppInferenceProvider`, `LlamaProcess` (spawns/owns llama-server), `LlamaCppEmbedder`, `LlamaCppTokenizer`, `GbnfGrammarConverter`; no separate AGENTS.md (sub-leaf, covered by this doc) +- `openai_compat/` — `OpenAiCompatInferenceProvider` (remote Bearer-auth chat completions for NVIDIA NIM/OpenAI), `HeuristicTokenizer`, `OpenAiApiModels`; no separate AGENTS.md (sub-leaf, covered by this doc) diff --git a/infrastructure/inference/openai_compat/build.gradle b/infrastructure/inference/openai_compat/build.gradle new file mode 100644 index 00000000..5ee48a72 --- /dev/null +++ b/infrastructure/inference/openai_compat/build.gradle @@ -0,0 +1,29 @@ +plugins { + id 'java-library' + id 'org.jetbrains.kotlin.jvm' + id 'org.jetbrains.kotlin.plugin.serialization' +} + +ext { + ktor_version = '3.0.3' + ext.ktor_version = '3.0.3' +} + +dependencies { + implementation project(':core:inference') + implementation project(':core:events') + implementation project(':core:artifacts') + implementation project(':core:context') + implementation project(':infrastructure:inference:commons') + + implementation "io.ktor:ktor-client-core:$ktor_version" + implementation "io.ktor:ktor-client-cio:$ktor_version" + implementation "io.ktor:ktor-client-content-negotiation:$ktor_version" + implementation "io.ktor:ktor-serialization-kotlinx-json:$ktor_version" + implementation "org.slf4j:slf4j-api:2.0.16" + testImplementation "org.junit.jupiter:junit-jupiter" + testImplementation "io.ktor:ktor-client-mock:$ktor_version" + testImplementation "org.jetbrains.kotlin:kotlin-test" +} + +tasks.named("koverVerify").configure { enabled = false } diff --git a/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/HeuristicTokenizer.kt b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/HeuristicTokenizer.kt new file mode 100644 index 00000000..dabcee91 --- /dev/null +++ b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/HeuristicTokenizer.kt @@ -0,0 +1,21 @@ +package com.correx.infrastructure.inference.openai + +import com.correx.core.inference.Token +import com.correx.core.inference.Tokenizer + +/** + * Remote OpenAI-compatible backends (NIM, OpenAI) expose no `/tokenize` endpoint, so we can't + * get exact token ids without shipping a tokenizer. This is a length-based approximation + * (~4 chars/token) used only for context-budget estimates; the synthesized [Token] ids are + * placeholders and must not be treated as real model tokens. + */ +class HeuristicTokenizer(private val charsPerToken: Int = 4) : Tokenizer { + + override suspend fun tokenize(text: String): List { + val count = countTokens(text) + return List(count) { Token(it) } + } + + override suspend fun countTokens(text: String): Int = + if (text.isEmpty()) 0 else (text.length + charsPerToken - 1) / charsPerToken +} diff --git a/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiApiModels.kt b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiApiModels.kt new file mode 100644 index 00000000..b89e83e4 --- /dev/null +++ b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiApiModels.kt @@ -0,0 +1,53 @@ +package com.correx.infrastructure.inference.openai + +import com.correx.core.inference.ToolCallRequest +import com.correx.core.inference.ToolDefinition +import kotlinx.serialization.SerialName +import kotlinx.serialization.Serializable + +/** + * Wire models for the OpenAI-compatible `/chat/completions` endpoint (NVIDIA NIM, OpenAI, + * vLLM, etc.). Unlike the llama.cpp models there is no `grammar` field — structured output + * on these backends is `response_format`/tools, not GBNF. JSON artifacts rely on the + * orchestrator's validate-after-retry path (invariant #7). + */ +@Serializable +data class OpenAiChatCompletionRequest( + val model: String, + val messages: List, + val temperature: Double, + @SerialName("top_p") val topP: Double, + @SerialName("max_tokens") val maxTokens: Int, + @SerialName("stop") val stopSequences: List? = null, + val seed: Long? = null, + val stream: Boolean = false, + val tools: List? = null, +) + +@Serializable +data class OpenAiChatMessage( + val role: String, + val content: String? = null, + @SerialName("tool_calls") val toolCalls: List = emptyList(), + @SerialName("tool_call_id") val toolCallId: String? = null, +) + +@Serializable +data class OpenAiChatCompletionResponse( + val id: String? = null, + val choices: List = emptyList(), + val usage: OpenAiUsage? = null, +) + +@Serializable +data class OpenAiChoice( + val message: OpenAiChatMessage, + @SerialName("finish_reason") val finishReason: String? = null, +) + +@Serializable +data class OpenAiUsage( + @SerialName("prompt_tokens") val promptTokens: Int = 0, + @SerialName("completion_tokens") val completionTokens: Int = 0, + @SerialName("total_tokens") val totalTokens: Int = 0, +) diff --git a/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProvider.kt b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProvider.kt new file mode 100644 index 00000000..07d311f7 --- /dev/null +++ b/infrastructure/inference/openai_compat/src/main/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProvider.kt @@ -0,0 +1,169 @@ +package com.correx.infrastructure.inference.openai + +import com.correx.core.events.types.ProviderId +import com.correx.core.inference.CapabilityScore +import com.correx.core.inference.FinishReason +import com.correx.core.inference.InferenceProvider +import com.correx.core.inference.InferenceRequest +import com.correx.core.inference.InferenceResponse +import com.correx.core.inference.PromptRenderer +import com.correx.core.inference.ProviderHealth +import com.correx.core.inference.TokenUsage +import com.correx.core.inference.Tokenizer +import com.correx.core.inference.ToolCallRequest +import io.ktor.client.HttpClient +import io.ktor.client.call.body +import io.ktor.client.engine.cio.CIO +import io.ktor.client.plugins.HttpTimeout +import io.ktor.client.plugins.contentnegotiation.ContentNegotiation +import io.ktor.client.request.accept +import io.ktor.client.request.get +import io.ktor.client.request.header +import io.ktor.client.request.post +import io.ktor.client.request.setBody +import io.ktor.client.statement.bodyAsText +import io.ktor.http.ContentType +import io.ktor.http.HttpHeaders +import io.ktor.http.contentType +import io.ktor.serialization.kotlinx.json.json +import kotlinx.serialization.encodeToString +import kotlinx.serialization.json.Json +import org.slf4j.LoggerFactory + +private const val DEFAULT_REQUEST_TIMEOUT_MS = 600_000L + +private val json = Json { + ignoreUnknownKeys = true + explicitNulls = false + encodeDefaults = true +} + +private fun defaultHttpClient(): HttpClient = HttpClient(CIO) { + install(ContentNegotiation) { json(json) } + install(HttpTimeout) { requestTimeoutMillis = DEFAULT_REQUEST_TIMEOUT_MS } +} + +private val log = LoggerFactory.getLogger(OpenAiCompatInferenceProvider::class.java) + +// Some models emit a tool call as a raw JSON blob in `content` instead of the native +// tool_calls array; salvage it so the orchestrator gets a real call rather than treating +// the blob as a (failing) artifact. Mirrors the llama.cpp provider's recovery path. +internal fun salvageToolCalls(content: String?): List { + val s = content?.trim().orEmpty() + if (s.isEmpty() || (s.first() != '{' && s.first() != '[')) return emptyList() + return runCatching { json.decodeFromString>(s) } + .recoverCatching { listOf(json.decodeFromString(s)) } + .getOrDefault(emptyList()) + .filter { it.function.name.isNotBlank() } +} + +/** + * [InferenceProvider] for OpenAI-compatible chat APIs reached over HTTPS with a Bearer key — + * NVIDIA NIM (`https://integrate.api.nvidia.com/v1`), OpenAI, vLLM, etc. Inference is fully + * remote: no local process, no GPU, no GGUF. [baseUrl] must already include the API version + * segment (e.g. `.../v1`); this class appends `/chat/completions` and `/models`. + * + * @param idPrefix label for the [ProviderId] (e.g. `nim`, `openai`). + * @param apiKey Bearer token; sent as `Authorization: Bearer ` when non-blank. + */ +@Suppress("TooGenericExceptionCaught", "MagicNumber") +class OpenAiCompatInferenceProvider( + private val modelId: String, + private val baseUrl: String, + private val apiKey: String, + private val capabilities: Set, + private val idPrefix: String = "openai", + private val httpClient: HttpClient = defaultHttpClient(), +) : InferenceProvider { + + override val id: ProviderId = ProviderId("$idPrefix:$modelId") + override val name: String = "OpenAI-compatible ($idPrefix:$modelId)" + override val tokenizer: Tokenizer = HeuristicTokenizer() + + override suspend fun infer(request: InferenceRequest): InferenceResponse { + val startTime = System.currentTimeMillis() + + val messages = PromptRenderer.render(request.contextPack).map { msg -> + // Strict OpenAI-compatible servers reject a `tool` role that isn't bound to a prior + // tool_calls turn (we don't carry tool_call_ids through the context pack). Fold any + // non-standard role into a user turn so the model still sees the content. + if (msg.role in STANDARD_ROLES) { + OpenAiChatMessage(role = msg.role, content = msg.content) + } else { + OpenAiChatMessage(role = "user", content = "[${msg.role}] ${msg.content}") + } + } + + val tools = request.tools.takeIf { it.isNotEmpty() } + + val body = OpenAiChatCompletionRequest( + model = modelId, + messages = messages, + temperature = request.generationConfig.temperature, + topP = request.generationConfig.topP, + maxTokens = request.generationConfig.maxTokens, + stopSequences = request.generationConfig.stopSequences.ifEmpty { null }, + seed = request.generationConfig.seed, + stream = false, + tools = tools, + ) + + val encoded = json.encodeToString(body) + log.debug("sending request to {}: {}", id.value, encoded) + + val httpResponse = httpClient.post("$baseUrl/chat/completions") { + contentType(ContentType.Application.Json) + accept(ContentType.Application.Json) + if (apiKey.isNotBlank()) header(HttpHeaders.Authorization, "Bearer $apiKey") + setBody(encoded) + } + if (httpResponse.status.value !in 200..299) { + val errorBody = httpResponse.bodyAsText() + error("$idPrefix returned ${httpResponse.status.value} ${httpResponse.status.description}: $errorBody") + } + val response = httpResponse.body() + log.debug("got response from {}: {}", id.value, response) + + val choice = response.choices.firstOrNull() + ?: error("$idPrefix returned no choices") + val message = choice.message + val salvaged = if (message.toolCalls.isEmpty()) salvageToolCalls(message.content) else emptyList() + val toolCalls = message.toolCalls.ifEmpty { salvaged } + val finishReason = when { + toolCalls.isNotEmpty() -> FinishReason.ToolCall + choice.finishReason?.lowercase() == "length" -> FinishReason.Length + else -> FinishReason.Stop + } + + return InferenceResponse( + requestId = request.requestId, + text = if (salvaged.isNotEmpty()) "" else message.content ?: "", + finishReason = finishReason, + tokensUsed = TokenUsage( + promptTokens = response.usage?.promptTokens ?: 0, + completionTokens = response.usage?.completionTokens ?: 0, + ), + latencyMs = System.currentTimeMillis() - startTime, + toolCalls = toolCalls, + ) + } + + override suspend fun healthCheck(): ProviderHealth = try { + val response = httpClient.get("$baseUrl/models") { + if (apiKey.isNotBlank()) header(HttpHeaders.Authorization, "Bearer $apiKey") + } + if (response.status.value in 200..299) { + ProviderHealth.Healthy + } else { + ProviderHealth.Unavailable("Health check returned non-2xx status: ${response.status}") + } + } catch (e: Exception) { + ProviderHealth.Unavailable("Health check failed: ${e.message}") + } + + override fun capabilities(): Set = capabilities + + private companion object { + private val STANDARD_ROLES = setOf("system", "user", "assistant") + } +} diff --git a/infrastructure/inference/openai_compat/src/test/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProviderTest.kt b/infrastructure/inference/openai_compat/src/test/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProviderTest.kt new file mode 100644 index 00000000..e3768071 --- /dev/null +++ b/infrastructure/inference/openai_compat/src/test/kotlin/com/correx/infrastructure/inference/openai/OpenAiCompatInferenceProviderTest.kt @@ -0,0 +1,130 @@ +package com.correx.infrastructure.inference.openai + +import com.correx.core.context.model.ContextPack +import com.correx.core.events.types.ContextPackId +import com.correx.core.events.types.InferenceRequestId +import com.correx.core.events.types.SessionId +import com.correx.core.events.types.StageId +import com.correx.core.inference.CapabilityScore +import com.correx.core.inference.FinishReason +import com.correx.core.inference.GenerationConfig +import com.correx.core.inference.InferenceRequest +import com.correx.core.inference.ModelCapability +import io.ktor.client.HttpClient +import io.ktor.client.engine.mock.MockEngine +import io.ktor.client.engine.mock.respond +import io.ktor.client.plugins.contentnegotiation.ContentNegotiation +import io.ktor.client.request.HttpRequestData +import io.ktor.http.ContentType +import io.ktor.http.HttpHeaders +import io.ktor.http.headersOf +import io.ktor.serialization.kotlinx.json.json +import kotlinx.coroutines.runBlocking +import kotlinx.serialization.json.Json +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +class OpenAiCompatInferenceProviderTest { + + private fun request() = InferenceRequest( + requestId = InferenceRequestId("req1"), + sessionId = SessionId("s1"), + stageId = StageId("stage1"), + contextPack = ContextPack( + id = ContextPackId("cp1"), + sessionId = SessionId("s1"), + stageId = StageId("stage1"), + layers = emptyMap(), + budgetUsed = 0, + budgetLimit = 4096, + ), + generationConfig = GenerationConfig(temperature = 0.2, topP = 0.9, maxTokens = 256), + ) + + private fun clientReturning(body: String, capture: (HttpRequestData) -> Unit = {}): HttpClient { + val engine = MockEngine { req -> + capture(req) + respond( + content = body, + headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()), + ) + } + return HttpClient(engine) { + install(ContentNegotiation) { json(Json { ignoreUnknownKeys = true }) } + } + } + + private fun provider(client: HttpClient) = OpenAiCompatInferenceProvider( + modelId = "deepseek-ai/deepseek-v4-flash", + baseUrl = "https://integrate.api.nvidia.com/v1", + apiKey = "nvapi-test", + capabilities = setOf(CapabilityScore(ModelCapability.General, 1.0)), + idPrefix = "nim", + httpClient = client, + ) + + @Test + fun `infer sends bearer auth to chat completions and returns text`(): Unit = runBlocking { + var seenAuth: String? = null + var seenUrl: String? = null + val client = clientReturning( + """ + {"id":"x","choices":[{"message":{"role":"assistant","content":"hello there"}, + "finish_reason":"stop"}],"usage":{"prompt_tokens":11,"completion_tokens":3,"total_tokens":14}} + """.trimIndent(), + ) { req -> + seenAuth = req.headers[HttpHeaders.Authorization] + seenUrl = req.url.toString() + } + + val resp = provider(client).infer(request()) + + assertEquals("Bearer nvapi-test", seenAuth) + assertEquals("https://integrate.api.nvidia.com/v1/chat/completions", seenUrl) + assertEquals("hello there", resp.text) + assertEquals(FinishReason.Stop, resp.finishReason) + assertEquals(11, resp.tokensUsed.promptTokens) + assertEquals(3, resp.tokensUsed.completionTokens) + } + + @Test + fun `infer surfaces native tool calls as ToolCall finish`(): Unit = runBlocking { + val client = clientReturning( + """ + {"id":"x","choices":[{"message":{"role":"assistant","content":null, + "tool_calls":[{"id":"c1","type":"function","function":{"name":"file_read","arguments":"{}"}}]}, + "finish_reason":"tool_calls"}],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}} + """.trimIndent(), + ) + + val resp = provider(client).infer(request()) + + assertEquals(FinishReason.ToolCall, resp.finishReason) + assertEquals(1, resp.toolCalls.size) + assertEquals("file_read", resp.toolCalls.first().function.name) + } + + @Test + fun `infer salvages a tool call emitted as a JSON blob in content`(): Unit = runBlocking { + val client = clientReturning( + """ + {"id":"x","choices":[{"message":{"role":"assistant", + "content":"{\"function\":{\"name\":\"file_read\",\"arguments\":\"{}\"}}"}, + "finish_reason":"stop"}],"usage":{"prompt_tokens":4,"completion_tokens":9,"total_tokens":13}} + """.trimIndent(), + ) + + val resp = provider(client).infer(request()) + + assertEquals(FinishReason.ToolCall, resp.finishReason) + assertEquals("file_read", resp.toolCalls.first().function.name) + assertEquals("", resp.text) + } + + @Test + fun `salvageToolCalls leaves plain artifact JSON alone`() { + assertTrue(salvageToolCalls("""{"title":"a plan","steps":[]}""").isEmpty()) + assertTrue(salvageToolCalls("not json").isEmpty()) + } +} diff --git a/infrastructure/inference/src/main/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategy.kt b/infrastructure/inference/src/main/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategy.kt new file mode 100644 index 00000000..86b20ff6 --- /dev/null +++ b/infrastructure/inference/src/main/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategy.kt @@ -0,0 +1,38 @@ +package com.correx.infrastructure.inference + +import com.correx.core.events.types.StageId +import com.correx.core.inference.InferenceProvider +import com.correx.core.inference.ModelCapability +import com.correx.core.inference.NoEligibleProviderException +import com.correx.core.inference.RoutingStrategy + +/** + * Like [FirstAvailableRoutingStrategy] it hard-filters to providers that declare every + * [requiredCapabilities] — but among those it picks the one with the highest summed capability + * **score** for the required set (falling back to total declared score when nothing is required). + * Ties keep list order, so behaviour is identical to first-available when scores are equal. + * + * This is the seam the orchestrator relies on to steer tool-heavy stages (which request + * `ModelCapability.ToolCalling`) to the model that calls tools best, rather than whichever healthy + * provider happens to come first. The scores come from each provider's declared + * `capabilities()` (config-supplied); observed per-model reliability (the `/metrics/tool-reliability` + * groundwork) can later feed in as an additional weight. + */ +class CapabilityAwareRoutingStrategy : RoutingStrategy { + override fun select( + candidates: List, + requiredCapabilities: Set, + ): InferenceProvider { + val eligible = candidates.filter { provider -> + provider.capabilities().map { it.capability }.toSet().containsAll(requiredCapabilities) + } + // maxByOrNull keeps the first element on ties, preserving first-available order. + return eligible.maxByOrNull { score(it, requiredCapabilities) } + ?: throw NoEligibleProviderException(StageId("routing"), requiredCapabilities) + } + + private fun score(provider: InferenceProvider, required: Set): Double { + val scores = provider.capabilities().associate { it.capability to it.score } + return if (required.isEmpty()) scores.values.sum() else required.sumOf { scores[it] ?: 0.0 } + } +} diff --git a/infrastructure/inference/src/test/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategyTest.kt b/infrastructure/inference/src/test/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategyTest.kt new file mode 100644 index 00000000..b19d9d16 --- /dev/null +++ b/infrastructure/inference/src/test/kotlin/com/correx/infrastructure/inference/CapabilityAwareRoutingStrategyTest.kt @@ -0,0 +1,73 @@ +package com.correx.infrastructure.inference + +import com.correx.core.events.types.ProviderId +import com.correx.core.inference.CapabilityScore +import com.correx.core.inference.InferenceProvider +import com.correx.core.inference.InferenceRequest +import com.correx.core.inference.InferenceResponse +import com.correx.core.inference.ModelCapability +import com.correx.core.inference.NoEligibleProviderException +import com.correx.core.inference.ProviderHealth +import com.correx.core.inference.Token +import com.correx.core.inference.Tokenizer +import com.correx.core.utils.TypeId +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.assertThrows + +class CapabilityAwareRoutingStrategyTest { + + private val strategy = CapabilityAwareRoutingStrategy() + + private fun fakeProvider(id: String, vararg caps: Pair): InferenceProvider = + object : InferenceProvider { + override val id: ProviderId = TypeId(id) + override val name: String = id + override val tokenizer: Tokenizer = object : Tokenizer { + override suspend fun tokenize(text: String): List = emptyList() + override suspend fun countTokens(text: String): Int = 0 + } + override suspend fun infer(request: InferenceRequest): InferenceResponse = error("unused") + override suspend fun healthCheck(): ProviderHealth = ProviderHealth.Healthy + override fun capabilities(): Set = + caps.map { CapabilityScore(it.first, it.second) }.toSet() + } + + @Test + fun `picks the highest-scoring provider for the required capability`() { + val weak = fakeProvider("weak", ModelCapability.ToolCalling to 0.4, ModelCapability.General to 1.0) + val strong = fakeProvider("strong", ModelCapability.ToolCalling to 0.95, ModelCapability.General to 1.0) + val selected = strategy.select(listOf(weak, strong), setOf(ModelCapability.ToolCalling)) + assertEquals("strong", selected.id.value) + } + + @Test + fun `still hard-filters to providers declaring the required capability`() { + val noTools = fakeProvider("noTools", ModelCapability.General to 1.0) + val tools = fakeProvider("tools", ModelCapability.ToolCalling to 0.3) + val selected = strategy.select(listOf(noTools, tools), setOf(ModelCapability.ToolCalling)) + assertEquals("tools", selected.id.value) + } + + @Test + fun `throws when no candidate declares the required capability`() { + val general = fakeProvider("general", ModelCapability.General to 1.0) + assertThrows { + strategy.select(listOf(general), setOf(ModelCapability.ToolCalling)) + } + } + + @Test + fun `keeps list order on a score tie (first-available parity)`() { + val a = fakeProvider("a", ModelCapability.ToolCalling to 0.5) + val b = fakeProvider("b", ModelCapability.ToolCalling to 0.5) + assertEquals("a", strategy.select(listOf(a, b), setOf(ModelCapability.ToolCalling)).id.value) + } + + @Test + fun `with no required capabilities ranks by total declared score`() { + val lean = fakeProvider("lean", ModelCapability.General to 0.5) + val rich = fakeProvider("rich", ModelCapability.General to 1.0, ModelCapability.Coding to 0.9) + assertEquals("rich", strategy.select(listOf(lean, rich), emptySet()).id.value) + } +} diff --git a/infrastructure/src/main/kotlin/com/correx/infrastructure/InfrastructureModule.kt b/infrastructure/src/main/kotlin/com/correx/infrastructure/InfrastructureModule.kt index 8b65bc9e..51725ed5 100644 --- a/infrastructure/src/main/kotlin/com/correx/infrastructure/InfrastructureModule.kt +++ b/infrastructure/src/main/kotlin/com/correx/infrastructure/InfrastructureModule.kt @@ -47,6 +47,7 @@ import com.correx.infrastructure.inference.commons.ResidencyMode import com.correx.infrastructure.inference.llama.cpp.DefaultModelManager import com.correx.infrastructure.inference.llama.cpp.LlamaCppEmbedder import com.correx.infrastructure.inference.llama.cpp.LlamaCppInferenceProvider +import com.correx.infrastructure.inference.openai.OpenAiCompatInferenceProvider import io.ktor.client.HttpClient import io.ktor.client.engine.cio.CIO import com.correx.infrastructure.router.turbovec.TurboVecL3MemoryStore @@ -127,6 +128,20 @@ object InfrastructureModule { baseUrl = baseUrl, ) + fun createOpenAiCompatProvider( + modelId: String, + baseUrl: String, + apiKey: String, + idPrefix: String = "openai", + capabilities: Set = DEFAULT_LLAMA_CAPABILITIES, + ): OpenAiCompatInferenceProvider = OpenAiCompatInferenceProvider( + modelId = modelId, + baseUrl = baseUrl, + apiKey = apiKey, + idPrefix = idPrefix, + capabilities = capabilities, + ) + fun createModelManager( settings: ModelsSettings, eventStore: EventStore, diff --git a/infrastructure/tools/AGENTS.md b/infrastructure/tools/AGENTS.md index debbd3e6..6d51f4c3 100644 --- a/infrastructure/tools/AGENTS.md +++ b/infrastructure/tools/AGENTS.md @@ -16,6 +16,7 @@ Adapter for `core:tools`. Depends on `core:tools`, `core:events`, `core:approval - Web search and web fetch results are environment observations; they must be recorded as events by callers to preserve replay determinism (invariant #9). - `ToolConfig` is the only configuration surface; pass via `InfrastructureModule.createToolExecutor()`. - `buildTools()` extension on `ToolConfig` assembles the full tool list; add new tools there, not in the registry directly. +- Filesystem mutation is split by intent: `file_write` only writes (`{path, content}`), `file_edit` edits, and `file_delete` only deletes (`{path}`) — deletion is a separately-named capability so a model can never delete by getting a write-mode parameter wrong. `file_delete` shares `file_write`'s path jail and `fileWrite.enabled` toggle and carries `ToolCapability.FILE_WRITE`. ## Work Guidance @@ -30,4 +31,4 @@ Standard adapter rules apply (see parent `AGENTS.md`). Network calls (web tools) ## Child DOX Index -- `filesystem/` — filesystem read/write/list tools implementing `core:tools` contracts; no separate AGENTS.md (sub-leaf, covered by this doc) +- `filesystem/` — filesystem tools implementing `core:tools` contracts: `FileReadTool`, `FileWriteTool` (write-only), `FileDeleteTool`, `FileEditTool`, list; no separate AGENTS.md (sub-leaf, covered by this doc) diff --git a/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteTool.kt b/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteTool.kt new file mode 100644 index 00000000..972fe80c --- /dev/null +++ b/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteTool.kt @@ -0,0 +1,134 @@ +package com.correx.infrastructure.tools.filesystem + +import com.correx.core.approvals.Tier +import com.correx.core.events.events.ToolRequest +import com.correx.core.events.types.ToolInvocationId +import com.correx.core.tools.contract.FileAffectingTool +import com.correx.core.tools.contract.ParamRole +import com.correx.core.tools.contract.Tool +import com.correx.core.tools.contract.ToolCapability +import com.correx.core.tools.contract.ToolExecutor +import com.correx.core.tools.contract.ToolResult +import com.correx.core.tools.contract.ValidationResult +import kotlinx.coroutines.CancellationException +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.withContext +import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.JsonPrimitive +import kotlinx.serialization.json.buildJsonArray +import kotlinx.serialization.json.buildJsonObject +import kotlinx.serialization.json.put +import kotlinx.serialization.json.putJsonObject +import java.io.IOException +import java.nio.file.Files +import java.nio.file.InvalidPathException +import java.nio.file.Path +import java.nio.file.Paths + +/** + * Deletes a file. Split out of [FileWriteTool] so deletion is an explicitly-named capability a model + * must choose deliberately — it can never happen by getting a write-mode parameter wrong. Carries the + * [ToolCapability.FILE_WRITE] capability (file mutation) and the same path jail as the writer. + */ +class FileDeleteTool( + allowedPaths: Set = emptySet(), + private val workingDir: Path? = null, +) : Tool, FileAffectingTool, ToolExecutor { + + private val normalizedAllowedPaths: Set = allowedPaths.map { it.normalize().toAbsolutePath() }.toSet() + + override val name: String = "file_delete" + override val description: String = "Delete the file at the specified relative path" + override val parametersSchema: JsonObject = buildJsonObject { + put("type", "object") + putJsonObject("properties") { + putJsonObject("path") { + put("type", "string") + put("description", "Relative path of the file to delete") + } + } + put("required", buildJsonArray { add(JsonPrimitive("path")) }) + } + override val tier: Tier = Tier.T2 + override val requiredCapabilities: Set = setOf(ToolCapability.FILE_WRITE) + override val paramRoles: Map = mapOf("path" to ParamRole.PATH) + + override fun affectedPaths(request: ToolRequest): Set { + val pathString = request.parameters["path"] as? String ?: return emptySet() + return setOf(resolvePath(pathString)) + } + + private fun resolvePath(pathString: String): Path { + val raw = Paths.get(pathString) + return when { + raw.isAbsolute -> raw.normalize() + workingDir != null -> workingDir.resolve(raw).normalize() + else -> raw.toAbsolutePath().normalize() + } + } + + override fun validateRequest(request: ToolRequest): ValidationResult { + val pathString = request.parameters["path"] as? String + ?: return ValidationResult.Invalid( + """Missing 'path' parameter (string). Call file_delete with {"path": ""}.""", + ) + return runCatching { + val path = resolvePath(pathString) + when { + normalizedAllowedPaths.isEmpty() -> ValidationResult.Invalid("No paths are allowed.") + !PathJail.isContained(path, normalizedAllowedPaths) -> + ValidationResult.Invalid("Path '$pathString' is not in the allowed list.") + + else -> ValidationResult.Valid + } + }.getOrElse { e -> mapExceptionToValidationResult(e) } + } + + private fun mapExceptionToValidationResult(e: Throwable): ValidationResult = + when (e) { + is InvalidPathException -> ValidationResult.Invalid("Invalid path format: ${e.message}") + is IOException -> ValidationResult.Invalid("IO error: ${e.message}") + is SecurityException -> ValidationResult.Invalid("Security error: ${e.message}") + else -> ValidationResult.Invalid(e.message ?: "Unknown error occurred") + } + + override suspend fun execute(request: ToolRequest): ToolResult = withContext(Dispatchers.IO) { + val validation = validateRequest(request) + if (validation is ValidationResult.Invalid) { + return@withContext ToolResult.Failure( + invocationId = request.invocationId, + reason = validation.reason, + recoverable = false, + ) + } + val pathString = request.parameters["path"] as String + val path = resolvePath(pathString) + runCatching { + if (Files.deleteIfExists(path)) { + ToolResult.Success( + invocationId = request.invocationId, + output = "File deleted successfully: $pathString", + ) + } else { + ToolResult.Failure( + invocationId = request.invocationId, + reason = "File not found: $pathString", + recoverable = false, + ) + } + }.getOrElse { e -> handleExecutionException(e, request.invocationId, pathString) } + } + + private fun handleExecutionException( + e: Throwable, + invocationId: ToolInvocationId, + pathString: String, + ): ToolResult = when (e) { + is CancellationException -> throw e + is IOException -> ToolResult.Failure(invocationId, "IO error: ${e.message}", recoverable = false) + is SecurityException -> + ToolResult.Failure(invocationId, "Access denied: $pathString, ${e.message}", recoverable = false) + + else -> ToolResult.Failure(invocationId, e.message ?: "Unknown error occurred", recoverable = false) + } +} diff --git a/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteTool.kt b/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteTool.kt index ccc9932b..6079f8a2 100644 --- a/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteTool.kt +++ b/infrastructure/tools/filesystem/src/main/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteTool.kt @@ -25,6 +25,12 @@ import java.nio.file.InvalidPathException import java.nio.file.Path import java.nio.file.Paths +/** + * Writes content to a file. Write-only by design: deleting is the separate, explicitly-named + * [FileDeleteTool] so a model can never delete a file by getting an `operation` mode wrong — a + * destructive action must name itself. (Previously this tool carried an `operation: write|delete` + * mode; splitting it removes the most-forgotten parameter and makes delete a distinct capability.) + */ class FileWriteTool( allowedPaths: Set = emptySet(), private val workingDir: Path? = null, @@ -33,7 +39,7 @@ class FileWriteTool( private val normalizedAllowedPaths: Set = allowedPaths.map { it.normalize().toAbsolutePath() }.toSet() override val name: String = "file_write" - override val description: String = "Write content to a file at the specified path or delete the file entirely" + override val description: String = "Write content to a file at the specified relative path (creates or overwrites)" override val parametersSchema: JsonObject = buildJsonObject { put("type", "object") putJsonObject("properties") { @@ -43,11 +49,7 @@ class FileWriteTool( } putJsonObject("content") { put("type", "string") - put("description", "File content") - } - putJsonObject("operation") { - put("type", "string") - put("description", "Either 'write' or 'delete'") + put("description", "The full file content to write") } } put( @@ -55,7 +57,6 @@ class FileWriteTool( buildJsonArray { add(JsonPrimitive("path")) add(JsonPrimitive("content")) - add(JsonPrimitive("operation")) }, ) } @@ -78,45 +79,37 @@ class FileWriteTool( } override fun validateRequest(request: ToolRequest): ValidationResult { - val operation = request.parameters["operation"] as? String val pathString = request.parameters["path"] as? String + val hasContent = request.parameters.containsKey("content") return when { - operation == null -> - ValidationResult.Invalid("Missing 'operation' parameter. Expected 'write' or 'delete'.") - pathString == null -> - ValidationResult.Invalid("Missing 'path' parameter. Expected String.") + ValidationResult.Invalid( + "Missing 'path' parameter (string). Call file_write with " + + """{"path": "", "content": ""}.""", + ) + + !hasContent -> + ValidationResult.Invalid( + "Missing 'content' parameter (string). Call file_write with " + + """{"path": "$pathString", "content": ""}.""", + ) else -> runCatching { val path = resolvePath(pathString) - checkPathAllowed(path, pathString, operation, request) + checkPathAllowed(path, pathString) }.getOrElse { e -> mapExceptionToValidationResult(e) } } } - private fun checkPathAllowed( - path: Path, - pathString: String, - operation: String, - request: ToolRequest, - ): ValidationResult { - return when { + private fun checkPathAllowed(path: Path, pathString: String): ValidationResult = + when { normalizedAllowedPaths.isEmpty() -> ValidationResult.Invalid("No paths are allowed.") - !isPathAllowed(path) -> - ValidationResult.Invalid("Path '$pathString' is not in the allowed list.") - - operation == "write" && !request.parameters.containsKey("content") -> - ValidationResult.Invalid("Missing 'content' parameter for 'write' operation.") - - operation != "write" && operation != "delete" -> - ValidationResult.Invalid("Unknown operation: $operation") - + !isPathAllowed(path) -> ValidationResult.Invalid("Path '$pathString' is not in the allowed list.") else -> ValidationResult.Valid } - } private fun isPathAllowed(path: Path): Boolean = PathJail.isContained(path, normalizedAllowedPaths) @@ -139,13 +132,17 @@ class FileWriteTool( ) } - val operation = request.parameters["operation"] as String val pathString = request.parameters["path"] as String + val content = request.parameters["content"] as String val path = resolvePath(pathString) val originalContent = runCatching { Files.readString(path) }.getOrDefault("") val result = runCatching { - performOperation(operation, path, pathString, request) + AtomicFileWriter.write(path, content.toByteArray(Charsets.UTF_8)) + ToolResult.Success( + invocationId = request.invocationId, + output = "File written successfully to $pathString", + ) }.getOrElse { e -> handleExecutionException(e, request.invocationId, pathString) } @@ -165,48 +162,6 @@ class FileWriteTool( } } - private suspend fun performOperation( - operation: String, - path: Path, - pathString: String, - request: ToolRequest, - ): ToolResult = when (operation) { - "write" -> { - val content = request.parameters["content"] as String - withContext(Dispatchers.IO) { - AtomicFileWriter.write(path, content.toByteArray(Charsets.UTF_8)) - } - ToolResult.Success( - invocationId = request.invocationId, - output = "File written successfully to $pathString", - ) - } - - "delete" -> { - if (Files.exists(path)) { - withContext(Dispatchers.IO) { - Files.deleteIfExists(path) - } - ToolResult.Success( - invocationId = request.invocationId, - output = "File deleted successfully: $pathString", - ) - } else { - ToolResult.Failure( - invocationId = request.invocationId, - reason = "File not found: $pathString", - recoverable = false, - ) - } - } - - else -> ToolResult.Failure( - invocationId = request.invocationId, - reason = "Unknown operation: $operation", - recoverable = false, - ) - } - private fun handleExecutionException( e: Throwable, invocationId: ToolInvocationId, diff --git a/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteToolTest.kt b/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteToolTest.kt new file mode 100644 index 00000000..49195b27 --- /dev/null +++ b/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileDeleteToolTest.kt @@ -0,0 +1,84 @@ +package com.correx.infrastructure.tools.filesystem + +import com.correx.core.events.events.ToolRequest +import com.correx.core.events.types.SessionId +import com.correx.core.events.types.StageId +import com.correx.core.events.types.ToolInvocationId +import com.correx.core.tools.contract.ToolResult +import com.correx.core.tools.contract.ValidationResult +import kotlinx.coroutines.runBlocking +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import java.nio.file.Files +import java.util.* + +class FileDeleteToolTest { + + private val invocationId = ToolInvocationId(UUID.randomUUID().toString()) + + private fun createRequest(parameters: Map): ToolRequest = ToolRequest( + invocationId = invocationId, + sessionId = SessionId(UUID.randomUUID().toString()), + stageId = StageId(UUID.randomUUID().toString()), + toolName = "file_delete", + parameters = parameters, + ) + + @Test + fun `validateRequest returns Valid for existing file in allowed directory`(): Unit = runBlocking { + val tempDir = Files.createTempDirectory("file_delete_test") + val filePath = tempDir.resolve("existing.txt") + Files.writeString(filePath, "content") + val tool = FileDeleteTool(allowedPaths = setOf(tempDir)) + assertEquals(ValidationResult.Valid, tool.validateRequest(createRequest(mapOf("path" to filePath.toString())))) + } + + @Test + fun `validateRequest returns Invalid for missing path with actionable message`(): Unit = runBlocking { + val tempDir = Files.createTempDirectory("file_delete_test") + val tool = FileDeleteTool(allowedPaths = setOf(tempDir)) + val result = tool.validateRequest(createRequest(emptyMap())) + assertTrue(result is ValidationResult.Invalid) + assertTrue((result as ValidationResult.Invalid).reason.contains("Missing 'path'")) + } + + @Test + fun `validateRequest returns Invalid for disallowed path`(): Unit = runBlocking { + val tempDir = Files.createTempDirectory("file_delete_test") + val otherDir = Files.createTempDirectory("other_dir") + val tool = FileDeleteTool(allowedPaths = setOf(tempDir)) + val result = tool.validateRequest(createRequest(mapOf("path" to otherDir.resolve("x.txt").toString()))) + assertTrue(result is ValidationResult.Invalid) + assertTrue((result as ValidationResult.Invalid).reason.contains("is not in the allowed list")) + } + + @Test + fun `execute deletes an existing file`(): Unit = runBlocking { + val tempDir = Files.createTempDirectory("file_delete_test") + val filePath = tempDir.resolve("to_delete.txt") + Files.writeString(filePath, "content to delete") + val tool = FileDeleteTool(allowedPaths = setOf(tempDir)) + + val result = tool.execute(createRequest(mapOf("path" to filePath.toString()))) + + assertTrue(result is ToolResult.Success) + assertTrue((result as ToolResult.Success).output.startsWith("File deleted successfully: $filePath")) + assertFalse(Files.exists(filePath)) + } + + @Test + fun `execute returns Failure for deleting non-existent file`(): Unit = runBlocking { + val tempDir = Files.createTempDirectory("file_delete_test") + val filePath = tempDir.resolve("non_existent.txt") + val tool = FileDeleteTool(allowedPaths = setOf(tempDir)) + + val result = tool.execute(createRequest(mapOf("path" to filePath.toString()))) + + assertTrue(result is ToolResult.Failure) + val failure = result as ToolResult.Failure + assertEquals("File not found: $filePath", failure.reason) + assertFalse(failure.recoverable) + } +} diff --git a/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteToolTest.kt b/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteToolTest.kt index ebfb8fc7..2a99fa34 100644 --- a/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteToolTest.kt +++ b/infrastructure/tools/filesystem/src/test/kotlin/com/correx/infrastructure/tools/filesystem/FileWriteToolTest.kt @@ -8,7 +8,6 @@ import com.correx.core.tools.contract.ToolResult import com.correx.core.tools.contract.ValidationResult import kotlinx.coroutines.runBlocking import org.junit.jupiter.api.Assertions.assertEquals -import org.junit.jupiter.api.Assertions.assertFalse import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.Test import java.nio.file.Files @@ -20,29 +19,20 @@ class FileWriteToolTest { private val stageId = StageId(UUID.randomUUID().toString()) private val invocationId = ToolInvocationId(UUID.randomUUID().toString()) - private fun createRequest( - parameters: Map, - toolName: String = "file_write", - ): ToolRequest { - return ToolRequest( - invocationId = invocationId, - sessionId = sessionId, - stageId = stageId, - toolName = toolName, - parameters = parameters, - ) - } + private fun createRequest(parameters: Map): ToolRequest = ToolRequest( + invocationId = invocationId, + sessionId = sessionId, + stageId = stageId, + toolName = "file_write", + parameters = parameters, + ) @Test - fun `validateRequest returns Valid for allowed path and valid operation`(): Unit = runBlocking { + fun `validateRequest returns Valid for allowed path with content`(): Unit = runBlocking { val tempDir = Files.createTempDirectory("file_write_test") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) val request = createRequest( - mapOf( - "operation" to "write", - "path" to tempDir.resolve("file.txt").toString(), - "content" to "hello", - ), + mapOf("path" to tempDir.resolve("file.txt").toString(), "content" to "hello"), ) assertEquals(ValidationResult.Valid, tool.validateRequest(request)) } @@ -53,11 +43,7 @@ class FileWriteToolTest { val otherDir = Files.createTempDirectory("other_dir") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) val request = createRequest( - mapOf( - "operation" to "write", - "path" to otherDir.resolve("file.txt").toString(), - "content" to "hello", - ), + mapOf("path" to otherDir.resolve("file.txt").toString(), "content" to "hello"), ) val result = tool.validateRequest(request) assertTrue(result is ValidationResult.Invalid) @@ -65,171 +51,58 @@ class FileWriteToolTest { } @Test - fun `validateRequest returns Invalid for missing operation`(): Unit = runBlocking { + fun `validateRequest returns Invalid for missing path with actionable message`(): Unit = runBlocking { val tempDir = Files.createTempDirectory("file_write_test") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf( - "path" to tempDir.resolve("file.txt").toString(), - "content" to "hello", - ), - ) - val result = tool.validateRequest(request) + val result = tool.validateRequest(createRequest(mapOf("content" to "hello"))) assertTrue(result is ValidationResult.Invalid) - assertEquals( - "Missing 'operation' parameter. Expected 'write' or 'delete'.", - (result as ValidationResult.Invalid).reason, - ) + val reason = (result as ValidationResult.Invalid).reason + assertTrue(reason.contains("Missing 'path'")) + // actionable: shows the required call shape + assertTrue(reason.contains("\"content\"")) } @Test - fun `validateRequest returns Invalid for missing path`(): Unit = runBlocking { + fun `validateRequest returns Invalid for missing content with actionable message`(): Unit = runBlocking { val tempDir = Files.createTempDirectory("file_write_test") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf( - "operation" to "write", - "content" to "hello", - ), - ) + val request = createRequest(mapOf("path" to tempDir.resolve("file.txt").toString())) val result = tool.validateRequest(request) assertTrue(result is ValidationResult.Invalid) - assertEquals("Missing 'path' parameter. Expected String.", (result as ValidationResult.Invalid).reason) + val reason = (result as ValidationResult.Invalid).reason + assertTrue(reason.contains("Missing 'content'")) + assertTrue(reason.contains("file_write")) } @Test - fun `validateRequest returns Invalid for missing content on write`(): Unit = runBlocking { - val tempDir = Files.createTempDirectory("file_write_test") - val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf( - "operation" to "write", - "path" to tempDir.resolve("file.txt").toString(), - ), - ) - val result = tool.validateRequest(request) - assertTrue(result is ValidationResult.Invalid) - assertEquals("Missing 'content' parameter for 'write' operation.", (result as ValidationResult.Invalid).reason) - } - - @Test - fun `validateRequest returns Invalid for unknown operation`(): Unit = runBlocking { - val tempDir = Files.createTempDirectory("file_write_test") - val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf( - "operation" to "unknown", - "path" to tempDir.resolve("file.txt").toString(), - ), - ) - val result = tool.validateRequest(request) - assertTrue(result is ValidationResult.Invalid) - assertEquals("Unknown operation: unknown", (result as ValidationResult.Invalid).reason) - } - - @Test - fun `execute returns Success for write operation`(): Unit = runBlocking { + fun `execute writes content and attaches a diff`(): Unit = runBlocking { val tempDir = Files.createTempDirectory("file_write_test") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) val filePath = tempDir.resolve("test.txt") val content = "Hello, FileWriteTool!" - val request = createRequest( - mapOf( - "operation" to "write", - "path" to filePath.toString(), - "content" to content, - ), - ) + val request = createRequest(mapOf("path" to filePath.toString(), "content" to content)) val result = tool.execute(request) assertTrue(result is ToolResult.Success) val success = result as ToolResult.Success - // Output starts with the status line; a unified diff of the change is appended (F-019). assertTrue(success.output.startsWith("File written successfully to $filePath")) assertTrue(success.metadata.containsKey("diff")) assertEquals(invocationId, success.invocationId) assertEquals(content, Files.readString(filePath)) } - @Test - fun `execute returns Success for delete operation`(): Unit = runBlocking { - val tempDir = Files.createTempDirectory("file_write_test") - val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val filePath = tempDir.resolve("to_delete.txt") - Files.writeString(filePath, "content to delete") - - val request = createRequest( - mapOf( - "operation" to "delete", - "path" to filePath.toString(), - ), - ) - - val result = tool.execute(request) - assertTrue(result is ToolResult.Success) - val success = result as ToolResult.Success - assertTrue(success.output.startsWith("File deleted successfully: $filePath")) - assertEquals(invocationId, success.invocationId) - assertFalse(Files.exists(filePath)) - } - - @Test - fun `execute returns Failure for deleting non-existent file`(): Unit = runBlocking { - val tempDir = Files.createTempDirectory("file_write_test") - val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val filePath = tempDir.resolve("non_existent.txt") - - val request = createRequest( - mapOf( - "operation" to "delete", - "path" to filePath.toString(), - ), - ) - - val result = tool.execute(request) - - assertTrue(result is ToolResult.Failure) - val failure = result as ToolResult.Failure - assertEquals("File not found: $filePath", failure.reason) - assertFalse(failure.recoverable) - } - @Test fun `execute returns Failure for disallowed path`(): Unit = runBlocking { val tempDir = Files.createTempDirectory("file_write_test") val otherDir = Files.createTempDirectory("other_dir") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val filePath = otherDir.resolve("illegal.txt") - val request = createRequest( - mapOf( - "operation" to "write", - "path" to filePath.toString(), - "content" to "illegal", - ), + mapOf("path" to otherDir.resolve("illegal.txt").toString(), "content" to "illegal"), ) - val result = tool.execute(request) - assertTrue(result is ToolResult.Failure) - val failure = result as ToolResult.Failure - assertTrue(failure.reason.contains("is not in the allowed list")) - } - - @Test - fun `validateRequest returns Valid for existing file in allowed directory`(): Unit = runBlocking { - val tempDir = Files.createTempDirectory("file_write_test_existing") - val filePath = tempDir.resolve("existing.txt") - Files.writeString(filePath, "content") - val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf( - "operation" to "delete", - "path" to filePath.toString(), - ), - ) - assertEquals(ValidationResult.Valid, tool.validateRequest(request)) + assertTrue((result as ToolResult.Failure).reason.contains("is not in the allowed list")) } @Test @@ -237,9 +110,7 @@ class FileWriteToolTest { val tempDir = Files.createTempDirectory("file_write_atomic") val target = tempDir.resolve("nested/deep/file.txt") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf("operation" to "write", "path" to target.toString(), "content" to "hello"), - ) + val request = createRequest(mapOf("path" to target.toString(), "content" to "hello")) val result = tool.execute(request) @@ -254,15 +125,12 @@ class FileWriteToolTest { val target = tempDir.resolve("file.txt") Files.writeString(target, "old") val tool = FileWriteTool(allowedPaths = setOf(tempDir)) - val request = createRequest( - mapOf("operation" to "write", "path" to target.toString(), "content" to "new"), - ) + val request = createRequest(mapOf("path" to target.toString(), "content" to "new")) val result = tool.execute(request) assertTrue(result is ToolResult.Success) assertEquals("new", Files.readString(target)) - // the staged temp file must not survive the swap val leftovers = Files.list(tempDir).use { stream -> stream.filter { it.fileName.toString().endsWith(".tmp") }.count() } diff --git a/infrastructure/tools/src/main/kotlin/com/correx/infrastructure/tools/ToolConfig.kt b/infrastructure/tools/src/main/kotlin/com/correx/infrastructure/tools/ToolConfig.kt index cab7c3f4..54288f27 100644 --- a/infrastructure/tools/src/main/kotlin/com/correx/infrastructure/tools/ToolConfig.kt +++ b/infrastructure/tools/src/main/kotlin/com/correx/infrastructure/tools/ToolConfig.kt @@ -1,6 +1,7 @@ package com.correx.infrastructure.tools import com.correx.core.tools.contract.Tool +import com.correx.infrastructure.tools.filesystem.FileDeleteTool import com.correx.infrastructure.tools.filesystem.FileEditTool import com.correx.infrastructure.tools.filesystem.FileReadTool import com.correx.infrastructure.tools.filesystem.FileWriteTool @@ -81,6 +82,14 @@ fun ToolConfig.buildTools(): List = buildList { workingDir = fileWrite.workingDir, ), ) + // file_delete is the explicit, separately-named deletion capability split out of + // file_write; it shares the writer's path jail and is gated by the same fileWrite toggle. + add( + FileDeleteTool( + allowedPaths = fileWrite.allowedPaths, + workingDir = fileWrite.workingDir, + ), + ) } if (fileEdit.enabled) { add( diff --git a/infrastructure/tools/src/test/kotlin/com/correx/infrastructure/tools/SandboxedToolExecutorFileMutationTest.kt b/infrastructure/tools/src/test/kotlin/com/correx/infrastructure/tools/SandboxedToolExecutorFileMutationTest.kt index 73eb2129..e96a3c04 100644 --- a/infrastructure/tools/src/test/kotlin/com/correx/infrastructure/tools/SandboxedToolExecutorFileMutationTest.kt +++ b/infrastructure/tools/src/test/kotlin/com/correx/infrastructure/tools/SandboxedToolExecutorFileMutationTest.kt @@ -65,7 +65,7 @@ class SandboxedToolExecutorFileMutationTest { private fun writeRequest(path: String, content: String) = ToolRequest( ToolInvocationId("inv"), SessionId("s"), StageId("st"), "file_write", - mapOf("operation" to "write", "path" to path, "content" to content), + mapOf("path" to path, "content" to content), ) private fun executor(tool: FileWriteTool, store: FakeArtifactStore, events: CapturingEventStore) = diff --git a/settings.gradle b/settings.gradle index 8026899a..9e2c929d 100644 --- a/settings.gradle +++ b/settings.gradle @@ -33,6 +33,7 @@ include ':infrastructure:persistence' include ':infrastructure:inference' include ':infrastructure:inference:commons' include ':infrastructure:inference:llama_cpp' +include ':infrastructure:inference:openai_compat' include ':infrastructure:router' include ':infrastructure:router:turbovec' include ':infrastructure:tools'