fix(kernel+infra): tool approval gate and executor cleanup

- SessionOrchestrator: wire T2/T3/T4 approval gate before tool execution;
  emit OrchestrationPausedEvent/ApprovalRequestedEvent, await decision,
  return rejection as ERROR context entry so LLM sees the denial;
  propagate tool ERROR entries as StageExecutionResult.Failure
- SandboxedToolExecutor: remove dead code and simplify
- InfrastructureModule: minor wiring cleanup
- LlamaCppInferenceProvider / build.gradle: related build fixes
This commit is contained in:
2026-05-19 00:05:08 +04:00
parent f32d400138
commit c0efa0ef03
5 changed files with 86 additions and 90 deletions
@@ -238,6 +238,10 @@ abstract class SessionOrchestrator(
toolRounds < MAX_TOOL_ROUNDS toolRounds < MAX_TOOL_ROUNDS
) { ) {
val toolEntries = dispatchToolCalls(sessionId, stageId, inferenceResult.response.toolCalls) val toolEntries = dispatchToolCalls(sessionId, stageId, inferenceResult.response.toolCalls)
val firstError = toolEntries.firstOrNull { it.content.startsWith("ERROR:") }
if (firstError != null) {
return StageExecutionResult.Failure(firstError.content.removePrefix("ERROR: "), retryable = true)
}
val allEntries = currentContext.layers.values.flatten() + toolEntries val allEntries = currentContext.layers.values.flatten() + toolEntries
currentContext = contextPackBuilder.build( currentContext = contextPackBuilder.build(
id = ContextPackId(UUID.randomUUID().toString()), id = ContextPackId(UUID.randomUUID().toString()),
@@ -297,6 +301,63 @@ abstract class SessionOrchestrator(
request = request, request = request,
), ),
) )
val requiresApproval = when (tier) {
Tier.T0, Tier.T1 -> false
Tier.T2, Tier.T3, Tier.T4 -> true
}
if (requiresApproval) {
val requestId = ApprovalRequestId(UUID.randomUUID().toString())
val domainRequest = DomainApprovalRequest(
id = requestId,
tier = tier,
validationReportId = ValidationReportId(UUID.randomUUID().toString()),
riskSummaryId = null,
timestamp = Clock.System.now(),
)
emit(sessionId, OrchestrationPausedEvent(sessionId, stageId, "APPROVAL_PENDING"))
emit(
sessionId,
ApprovalRequestedEvent(
requestId = requestId,
tier = tier,
validationReportId = domainRequest.validationReportId,
riskSummaryId = null,
sessionId = sessionId,
stageId = stageId,
projectId = null,
),
)
val deferred = CompletableDeferred<ApprovalDecision>()
pendingApprovals[requestId] = deferred
val decision = try {
deferred.await()
} finally {
pendingApprovals.remove(requestId)
}
emitDecisionResolved(sessionId, domainRequest, decision)
if (!decision.isApproved) {
val sourceId = toolCall.id ?: invocationId.value
val assistantEntry = ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "assistantToolCall",
sourceId = sourceId,
content = Json.encodeToString(ToolCallRequest.serializer(), toolCall),
tokenEstimate = toolCall.function.arguments.length / 4,
)
val reason = decision.reason ?: "approval denied"
val resultEntry = ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "toolResult",
sourceId = sourceId,
content = "ERROR: $reason",
tokenEstimate = reason.length / 4,
)
return@flatMap listOf(assistantEntry, resultEntry)
}
emit(sessionId, OrchestrationResumedEvent(sessionId, stageId))
}
val result = executor.execute(request) val result = executor.execute(request)
val sourceId = toolCall.id ?: invocationId.value val sourceId = toolCall.id ?: invocationId.value
val assistantEntry = ContextEntry( val assistantEntry = ContextEntry(
@@ -17,8 +17,6 @@ dependencies {
implementation project(':core:context') implementation project(':core:context')
implementation project(':infrastructure:inference:commons') implementation project(':infrastructure:inference:commons')
implementation("com.fasterxml.jackson.core:jackson-databind:2.17.0")
implementation("com.fasterxml.jackson.module:jackson-module-kotlin:2.17.0")
implementation "io.ktor:ktor-client-core:$ktor_version" implementation "io.ktor:ktor-client-core:$ktor_version"
implementation "io.ktor:ktor-client-cio:$ktor_version" implementation "io.ktor:ktor-client-cio:$ktor_version"
implementation "io.ktor:ktor-client-content-negotiation:$ktor_version" implementation "io.ktor:ktor-client-content-negotiation:$ktor_version"
@@ -25,7 +25,9 @@ import io.ktor.client.request.setBody
import io.ktor.http.ContentType import io.ktor.http.ContentType
import io.ktor.http.contentType import io.ktor.http.contentType
import io.ktor.serialization.kotlinx.json.json import io.ktor.serialization.kotlinx.json.json
import kotlinx.serialization.encodeToString
import kotlinx.serialization.json.Json import kotlinx.serialization.json.Json
import org.slf4j.LoggerFactory
private const val DEFAULT_REQUEST_TIMEOUT_MS = 60_000L private const val DEFAULT_REQUEST_TIMEOUT_MS = 60_000L
@@ -42,7 +44,23 @@ private fun defaultHttpClient(): HttpClient = HttpClient(CIO) {
install(HttpTimeout) { install(HttpTimeout) {
requestTimeoutMillis = DEFAULT_REQUEST_TIMEOUT_MS requestTimeoutMillis = DEFAULT_REQUEST_TIMEOUT_MS
} }
install(ContentNegotiation) {
json(
Json {
ignoreUnknownKeys = true
explicitNulls = false
encodeDefaults = true
},
)
}
} }
private val json = Json {
ignoreUnknownKeys = true
explicitNulls = false
encodeDefaults = true
}
private val log = LoggerFactory.getLogger(LlamaCppInferenceProvider::class.java)
@Suppress("TooGenericExceptionCaught", "MagicNumber") @Suppress("TooGenericExceptionCaught", "MagicNumber")
class LlamaCppInferenceProvider( class LlamaCppInferenceProvider(
@@ -80,12 +98,17 @@ class LlamaCppInferenceProvider(
tools = request.tools.takeIf { it.isNotEmpty() }, tools = request.tools.takeIf { it.isNotEmpty() },
) )
val encoded = json.encodeToString(body)
log.debug("sending request to llm: {}", encoded)
val response = httpClient.post("$baseUrl/v1/chat/completions") { val response = httpClient.post("$baseUrl/v1/chat/completions") {
contentType(ContentType.Application.Json) contentType(ContentType.Application.Json)
accept(ContentType.Application.Json) accept(ContentType.Application.Json)
setBody(body) setBody(encoded)
}.body<ChatCompletionResponse>() }.body<ChatCompletionResponse>()
log.debug("got response from llm: {}", response)
val message = response.choices.first().message val message = response.choices.first().message
val finishReason = when (response.choices.first().finishReason.lowercase()) { val finishReason = when (response.choices.first().finishReason.lowercase()) {
"length" -> FinishReason.Length "length" -> FinishReason.Length
@@ -1,6 +1,5 @@
package com.correx.infrastructure package com.correx.infrastructure
import com.correx.core.approvals.domain.ApprovalEngine
import com.correx.core.artifactstore.ArtifactStore import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.EventDispatcher import com.correx.core.events.EventDispatcher
import com.correx.core.events.stores.EventStore import com.correx.core.events.stores.EventStore
@@ -110,15 +109,11 @@ object InfrastructureModule {
fun createToolExecutor( fun createToolExecutor(
registry: ToolRegistry, registry: ToolRegistry,
approvalEngine: ApprovalEngine,
eventStore: EventStore,
eventDispatcher: EventDispatcher, eventDispatcher: EventDispatcher,
workDir: Path, workDir: Path,
): ToolExecutor = SandboxedToolExecutor( ): ToolExecutor = SandboxedToolExecutor(
delegate = DispatchingToolExecutor(registry), delegate = DispatchingToolExecutor(registry),
registry = registry, registry = registry,
approvalEngine = approvalEngine,
eventStore = eventStore,
eventDispatcher = eventDispatcher, eventDispatcher = eventDispatcher,
workDir = workDir, workDir = workDir,
) )
@@ -1,28 +1,15 @@
package com.correx.infrastructure.tools package com.correx.infrastructure.tools
import com.correx.core.approvals.Tier import com.correx.core.approvals.Tier
import com.correx.core.approvals.domain.ApprovalEngine
import com.correx.core.approvals.model.ApprovalContext
import com.correx.core.approvals.model.ApprovalGrant
import com.correx.core.approvals.model.ApprovalScopeIdentity
import com.correx.core.approvals.model.DomainApprovalRequest
import com.correx.core.events.EventDispatcher import com.correx.core.events.EventDispatcher
import com.correx.core.events.events.ApprovalGrantCreatedEvent
import com.correx.core.events.events.EventPayload import com.correx.core.events.events.EventPayload
import com.correx.core.events.events.StoredEvent
import com.correx.core.events.events.ToolExecutionCompletedEvent import com.correx.core.events.events.ToolExecutionCompletedEvent
import com.correx.core.events.events.ToolExecutionFailedEvent import com.correx.core.events.events.ToolExecutionFailedEvent
import com.correx.core.events.events.ToolExecutionRejectedEvent
import com.correx.core.events.events.ToolExecutionStartedEvent import com.correx.core.events.events.ToolExecutionStartedEvent
import com.correx.core.events.events.ToolReceipt import com.correx.core.events.events.ToolReceipt
import com.correx.core.events.events.ToolRequest import com.correx.core.events.events.ToolRequest
import com.correx.core.events.stores.EventStore
import com.correx.core.events.types.ApprovalRequestId
import com.correx.core.events.types.SessionId 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.events.types.ToolInvocationId
import com.correx.core.events.types.ValidationReportId
import com.correx.core.sessions.ApprovalMode
import com.correx.core.tools.contract.FileAffectingTool import com.correx.core.tools.contract.FileAffectingTool
import com.correx.core.tools.contract.Tool import com.correx.core.tools.contract.Tool
import com.correx.core.tools.contract.ToolExecutor import com.correx.core.tools.contract.ToolExecutor
@@ -41,8 +28,6 @@ import java.util.*
class SandboxedToolExecutor( class SandboxedToolExecutor(
private val delegate: ToolExecutor, private val delegate: ToolExecutor,
private val registry: ToolRegistry, private val registry: ToolRegistry,
private val approvalEngine: ApprovalEngine,
private val eventStore: EventStore,
private val eventDispatcher: EventDispatcher, private val eventDispatcher: EventDispatcher,
private val workDir: Path = Path.of("/tmp/correx-sandbox"), private val workDir: Path = Path.of("/tmp/correx-sandbox"),
) : ToolExecutor { ) : ToolExecutor {
@@ -61,26 +46,7 @@ class SandboxedToolExecutor(
recoverable = false, recoverable = false,
) )
// 2. approval check for T2-T4 // 2. emit started
val requiresApproval = when (tool.tier) {
Tier.T0, Tier.T1 -> false
Tier.T2, Tier.T3, Tier.T4 -> true
}
if (requiresApproval) {
val sessionEvents = eventStore.read(sessionId)
val decision = evaluateApproval(tool, sessionId, request.stageId, sessionEvents)
if (!decision.isApproved) {
emitRejected(sessionId, invocationId, toolName, tool.tier, decision.reason ?: "approval denied")
return@withContext ToolResult.Failure(
invocationId = invocationId,
reason = decision.reason ?: "approval denied",
recoverable = false,
)
}
}
// 3. emit started
emitStarted(sessionId, invocationId, toolName) emitStarted(sessionId, invocationId, toolName)
// 4. create working dir // 4. create working dir
@@ -121,45 +87,6 @@ class SandboxedToolExecutor(
} }
} }
// --- approval ---
private fun evaluateApproval(
tool: Tool,
sessionId: SessionId,
stageId: StageId,
sessionEvents: List<StoredEvent>,
) = approvalEngine.evaluate(
request = DomainApprovalRequest(
id = ApprovalRequestId(UUID.randomUUID().toString()),
tier = tool.tier,
validationReportId = ValidationReportId(UUID.randomUUID().toString()),
riskSummaryId = null,
timestamp = Clock.System.now(),
causationId = null,
correlationId = null,
),
context = ApprovalContext(
identity = ApprovalScopeIdentity(sessionId, stageId, null),
mode = ApprovalMode.PROMPT,
),
grants = extractGrants(sessionEvents),
now = Clock.System.now(),
)
private fun extractGrants(events: List<StoredEvent>): List<ApprovalGrant> =
events
.mapNotNull { it.payload as? ApprovalGrantCreatedEvent }
.map { e ->
ApprovalGrant(
id = e.grantId,
scope = e.scope,
permittedTiers = e.permittedTiers,
reason = e.reason,
timestamp = Clock.System.now(),
expiresAt = e.expiresAt,
)
}
// --- backup/restore --- // --- backup/restore ---
/** /**
@@ -209,14 +136,6 @@ class SandboxedToolExecutor(
// --- event emission --- // --- event emission ---
private suspend fun emitRejected(
sessionId: SessionId,
invocationId: ToolInvocationId,
toolName: String,
tier: Tier,
reason: String,
) = emit(sessionId, ToolExecutionRejectedEvent(invocationId, sessionId, toolName, tier, reason))
private suspend fun emitStarted( private suspend fun emitStarted(
sessionId: SessionId, sessionId: SessionId,
invocationId: ToolInvocationId, invocationId: ToolInvocationId,