feat(workspace): per-session workspace-scoped tools, policy, and undo

Compute the effective tool registry/executor and plane-2 WorkspacePolicy per run
from config.workspace (effectivesFor), threaded through the orchestrators; a null
workspace falls back to the boot instances byte-for-byte. Wire the concrete
WorkspaceToolRegistryProvider in Main (buildToolConfigForWorkspace). Session undo
unions the session's recorded workspace into its jail roots. (Axis 2 Phase A, tasks 3 + 7.)
This commit is contained in:
2026-06-02 19:57:51 +04:00
parent 7d7e524756
commit b2f60d09bb
11 changed files with 720 additions and 36 deletions
@@ -28,6 +28,9 @@ import com.correx.core.kernel.orchestration.OrchestrationProjector
import com.correx.core.kernel.orchestration.OrchestrationRepository
import com.correx.core.kernel.orchestration.OrchestratorEngines
import com.correx.core.kernel.orchestration.OrchestratorRepositories
import com.correx.core.kernel.orchestration.WorkspaceContext
import com.correx.core.kernel.orchestration.WorkspaceTools
import com.correx.core.kernel.orchestration.WorkspaceToolRegistryProvider
import com.correx.core.kernel.retry.DefaultRetryCoordinator
import com.correx.core.risk.DefaultRiskAssessor
import com.correx.core.router.l3.L3MetadataRehydrator
@@ -57,6 +60,7 @@ import com.correx.infrastructure.inference.commons.ResourceProbe
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.FileEditConfig
import com.correx.infrastructure.tools.FileReadConfig
import com.correx.infrastructure.tools.FileWriteConfig
@@ -188,6 +192,14 @@ fun main() {
),
)
val wsToolRegistryProvider = WorkspaceToolRegistryProvider { workspace ->
val wsRegistry = InfrastructureModule.createToolRegistry(
buildToolConfigForWorkspace(workspace, shellAllowedExecutables, toolsConfig),
)
val wsExecutor = DispatchingToolExecutor(wsRegistry)
WorkspaceTools(registry = wsRegistry, executor = wsExecutor)
}
val engines = OrchestratorEngines(
transitionResolver = DefaultTransitionResolver { condition, ctx -> condition.evaluate(ctx) },
contextPackBuilder = DefaultContextPackBuilder(DefaultContextCompressor()),
@@ -210,6 +222,7 @@ fun main() {
toolExecutor = toolExecutor,
toolCallAssessor = toolCallAssessor,
workspacePolicy = workspacePolicy,
workspaceToolRegistryProvider = wsToolRegistryProvider,
)
val orchestrator = DefaultSessionOrchestrator(
repositories = repositories,
@@ -239,10 +252,8 @@ fun main() {
)
val sessionUndoService = com.correx.apps.server.undo.SessionUndoService(
eventStore = eventStore,
reverser = com.correx.infrastructure.tools.filesystem.FileMutationReverser(
artifactStore = artifactStore,
allowedRoots = setOf(workspaceRoot, workingDir),
),
artifactStore = artifactStore,
bootRoots = setOf(workspaceRoot, workingDir),
)
val module = ServerModule(
orchestrator = orchestrator,
@@ -436,6 +447,32 @@ private fun buildToolConfig(
)
}
private fun buildToolConfigForWorkspace(
workspace: WorkspaceContext,
shellAllowedExecutables: Set<String>,
toolsConfig: com.correx.core.config.ToolsConfig,
): ToolConfig = ToolConfig(
shell = ShellConfig(
enabled = toolsConfig.shellEnabled,
allowedExecutables = shellAllowedExecutables,
workingDir = workspace.workingDir,
),
fileRead = FileReadConfig(
enabled = toolsConfig.fileReadEnabled,
allowedPaths = workspace.allowedPaths,
),
fileWrite = FileWriteConfig(
enabled = toolsConfig.fileWriteEnabled,
allowedPaths = workspace.allowedPaths,
workingDir = workspace.workingDir,
),
fileEdit = FileEditConfig(
enabled = toolsConfig.fileEditEnabled,
allowedPaths = workspace.allowedPaths,
workingDir = workspace.workingDir,
),
)
private fun DefaultProviderRegistry.asServerRegistry(): ProviderRegistry = object : ProviderRegistry {
override fun listAll() = this@asServerRegistry.listAll()
override suspend fun healthCheckAll() = this@asServerRegistry.healthCheckAll()
@@ -1,11 +1,14 @@
package com.correx.apps.server.undo
import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.events.FileWrittenEvent
import com.correx.core.events.events.SessionWorkspaceBoundEvent
import com.correx.core.events.stores.EventStore
import com.correx.core.events.types.SessionId
import com.correx.infrastructure.tools.filesystem.FileMutationReverser
import com.correx.infrastructure.tools.filesystem.RevertResult
import kotlinx.serialization.Serializable
import java.nio.file.Path
@Serializable
data class UndoSummary(
@@ -21,13 +24,34 @@ data class UndoSummary(
* Reverses every file mutation a session performed, newest-first, using only the
* event log + CAS (the reverser). Newest-first so a path written multiple times in
* the session ends at its pre-session content.
*
* Jail roots for each undo = [bootRoots] the workspace recorded in the session's
* [SessionWorkspaceBoundEvent] (if present). When no workspace event exists the
* effective roots are identical to [bootRoots], preserving pre-existing behaviour.
*/
class SessionUndoService(
private val eventStore: EventStore,
private val reverser: FileMutationReverser,
private val artifactStore: ArtifactStore,
private val bootRoots: Set<Path>,
) {
suspend fun undo(sessionId: SessionId): UndoSummary {
val mutations = eventStore.read(sessionId)
val events = eventStore.read(sessionId)
val sessionRoots: Set<Path> = events
.mapNotNull { it.payload as? SessionWorkspaceBoundEvent }
.lastOrNull()
?.let { bound ->
buildSet {
add(Path.of(bound.workspaceRoot))
bound.allowedPaths.forEach { add(Path.of(it)) }
}
}
?: emptySet()
val effectiveRoots = bootRoots + sessionRoots
val reverser = FileMutationReverser(artifactStore, effectiveRoots)
val mutations = events
.mapNotNull { it.payload as? FileWrittenEvent }
.asReversed()
@@ -44,7 +44,6 @@ import com.correx.infrastructure.tools.FileReadConfig
import com.correx.infrastructure.tools.FileWriteConfig
import com.correx.infrastructure.tools.ShellConfig
import com.correx.infrastructure.tools.ToolConfig
import com.correx.infrastructure.tools.filesystem.FileMutationReverser
import io.ktor.client.HttpClient
import io.ktor.client.engine.mock.MockEngine
import io.ktor.client.engine.mock.respond
@@ -124,10 +123,8 @@ fun buildTestServerModule(
val sessionUndoService = SessionUndoService(
eventStore = eventStore,
reverser = FileMutationReverser(
artifactStore = artifactStore,
allowedRoots = setOf(tempDir),
),
artifactStore = artifactStore,
bootRoots = setOf(tempDir),
)
val providerRegistry: ProviderRegistry = object : ProviderRegistry {
@@ -5,13 +5,13 @@ import com.correx.core.events.events.EventMetadata
import com.correx.core.events.events.EventPayload
import com.correx.core.events.events.FileWrittenEvent
import com.correx.core.events.events.NewEvent
import com.correx.core.events.events.SessionWorkspaceBoundEvent
import com.correx.core.events.events.StoredEvent
import com.correx.core.events.stores.EventStore
import com.correx.core.events.types.ArtifactId
import com.correx.core.events.types.EventId
import com.correx.core.events.types.SessionId
import com.correx.core.events.types.ToolInvocationId
import com.correx.infrastructure.tools.filesystem.FileMutationReverser
import kotlinx.coroutines.flow.Flow
import kotlinx.coroutines.flow.emptyFlow
import kotlinx.coroutines.runBlocking
@@ -20,6 +20,7 @@ import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Test
import java.nio.file.Files
import java.nio.file.Path
class SessionUndoServiceTest {
@@ -83,8 +84,7 @@ class SessionUndoServiceTest {
timestampMs = 0L,
)
val fakeEventStore = FakeEventStore(listOf(storedEvent(event, sessionId, 1L)))
val reverser = FileMutationReverser(fakeStore, setOf(dir))
val service = SessionUndoService(fakeEventStore, reverser)
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(dir))
val summary = service.undo(sessionId)
@@ -113,8 +113,7 @@ class SessionUndoServiceTest {
timestampMs = 0L,
)
val fakeEventStore = FakeEventStore(listOf(storedEvent(event, sessionId, 1L)))
val reverser = FileMutationReverser(fakeStore, setOf(dir))
val service = SessionUndoService(fakeEventStore, reverser)
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(dir))
val summary = service.undo(sessionId)
@@ -124,4 +123,212 @@ class SessionUndoServiceTest {
assertEquals(0, summary.failed)
assertFalse(Files.exists(target))
}
/**
* When the session log contains a [SessionWorkspaceBoundEvent] the workspace root
* is unioned into the jail roots, so a file written under that workspace (which is
* NOT in the boot roots) is successfully reverted rather than rejected.
*
* The sibling assertion that proves the workspace exclusion: a service built with
* boot roots that do NOT include [workspaceDir] would return rejected=1 instead of
* reverted=1. We verify this via a separate call using an unrelated boot root.
*/
@Test
fun `undo uses session workspace root from SessionWorkspaceBoundEvent`(): Unit = runBlocking {
val bootDir = Files.createTempDirectory("undo-boot").toRealPath()
val workspaceDir = Files.createTempDirectory("undo-workspace").toRealPath()
val target = workspaceDir.resolve("ws-file.txt")
Files.writeString(target, "CURRENT")
val fakeStore = FakeArtifactStore()
val preImageHash = fakeStore.put("ORIGINAL".toByteArray()).value
val sessionId = SessionId("session-workspace")
val workspaceBound = SessionWorkspaceBoundEvent(
sessionId = sessionId,
workspaceRoot = workspaceDir.toString(),
allowedPaths = emptyList(),
)
val fileWritten = FileWrittenEvent(
invocationId = ToolInvocationId("inv-ws"),
sessionId = sessionId,
path = target.toString(),
preImageHash = preImageHash,
postImageHash = null,
preExisted = true,
timestampMs = 0L,
)
val eventList = listOf(
storedEvent(workspaceBound, sessionId, 1L),
storedEvent(fileWritten, sessionId, 2L),
)
val fakeEventStore = FakeEventStore(eventList)
// Boot roots do NOT include workspaceDir — the workspace event unlocks it.
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(bootDir))
val summary = service.undo(sessionId)
assertEquals(1, summary.reverted)
assertEquals(0, summary.rejected, "workspace file must not be jailed out when bound event is present")
assertEquals("ORIGINAL", Files.readString(target))
// Prove that without the workspace event the revert would have been jailed:
// a service with boot-only roots (excluding workspaceDir) and no bound event rejects the file.
val storeNoWorkspace = FakeArtifactStore().also { it.blobs[preImageHash] = "ORIGINAL".toByteArray() }
val noWorkspaceEventStore = FakeEventStore(listOf(storedEvent(fileWritten, sessionId, 1L)))
val serviceBootOnly = SessionUndoService(noWorkspaceEventStore, storeNoWorkspace, setOf(bootDir))
val summaryBootOnly = serviceBootOnly.undo(sessionId)
assertEquals(1, summaryBootOnly.rejected, "boot-only roots must jail the workspace-dir file when no bound event exists")
}
/**
* A session with NO [SessionWorkspaceBoundEvent] uses boot roots exactly as before.
* Files inside boot roots are reverted; files outside are rejected.
*/
@Test
fun `undo with no SessionWorkspaceBoundEvent falls back to boot roots only`(): Unit = runBlocking {
val bootDir = Files.createTempDirectory("undo-boot-compat").toRealPath()
val outsideDir = Files.createTempDirectory("undo-outside").toRealPath()
val fakeStore = FakeArtifactStore()
val preImageHash = fakeStore.put("PREV".toByteArray()).value
val sessionId = SessionId("session-no-workspace")
// File inside boot root — should be reverted.
val insideTarget = bootDir.resolve("inside.txt")
Files.writeString(insideTarget, "NEW")
// File outside boot root and no bound event — should be rejected.
val outsideTarget = outsideDir.resolve("outside.txt")
Files.writeString(outsideTarget, "NEW")
val insideEvent = FileWrittenEvent(
invocationId = ToolInvocationId("inv-in"),
sessionId = sessionId,
path = insideTarget.toString(),
preImageHash = preImageHash,
postImageHash = null,
preExisted = true,
timestampMs = 0L,
)
val outsideEvent = FileWrittenEvent(
invocationId = ToolInvocationId("inv-out"),
sessionId = sessionId,
path = outsideTarget.toString(),
preImageHash = preImageHash,
postImageHash = null,
preExisted = true,
timestampMs = 0L,
)
val fakeEventStore = FakeEventStore(
listOf(storedEvent(insideEvent, sessionId, 1L), storedEvent(outsideEvent, sessionId, 2L)),
)
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(bootDir))
val summary = service.undo(sessionId)
assertEquals(1, summary.reverted, "file inside boot root must be reverted")
assertEquals(1, summary.rejected, "file outside boot root must be rejected when no workspace event exists")
}
/**
* When [SessionWorkspaceBoundEvent.allowedPaths] carries additional paths beyond
* [SessionWorkspaceBoundEvent.workspaceRoot], those paths are also included in the
* effective jail roots.
*/
@Test
fun `undo includes allowedPaths from SessionWorkspaceBoundEvent`(): Unit = runBlocking {
val bootDir = Files.createTempDirectory("undo-boot-allowed").toRealPath()
val workspaceDir = Files.createTempDirectory("undo-ws-root").toRealPath()
val extraAllowedDir = Files.createTempDirectory("undo-extra-allowed").toRealPath()
val target = extraAllowedDir.resolve("extra.txt")
Files.writeString(target, "STALE")
val fakeStore = FakeArtifactStore()
val preImageHash = fakeStore.put("FRESH".toByteArray()).value
val sessionId = SessionId("session-allowed-paths")
val workspaceBound = SessionWorkspaceBoundEvent(
sessionId = sessionId,
workspaceRoot = workspaceDir.toString(),
allowedPaths = listOf(extraAllowedDir.toString()),
)
val fileWritten = FileWrittenEvent(
invocationId = ToolInvocationId("inv-extra"),
sessionId = sessionId,
path = target.toString(),
preImageHash = preImageHash,
postImageHash = null,
preExisted = true,
timestampMs = 0L,
)
val fakeEventStore = FakeEventStore(
listOf(storedEvent(workspaceBound, sessionId, 1L), storedEvent(fileWritten, sessionId, 2L)),
)
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(bootDir))
val summary = service.undo(sessionId)
assertEquals(1, summary.reverted)
assertEquals(0, summary.rejected, "extra allowed path must not be jailed")
assertEquals("FRESH", Files.readString(target))
}
/**
* When multiple [SessionWorkspaceBoundEvent]s appear the LAST one wins.
*/
@Test
fun `undo uses last SessionWorkspaceBoundEvent when multiple are present`(): Unit = runBlocking {
val bootDir = Files.createTempDirectory("undo-boot-multi").toRealPath()
val firstWorkspace = Files.createTempDirectory("undo-ws-first").toRealPath()
val lastWorkspace = Files.createTempDirectory("undo-ws-last").toRealPath()
val target = lastWorkspace.resolve("last-ws-file.txt")
Files.writeString(target, "MODIFIED")
val fakeStore = FakeArtifactStore()
val preImageHash = fakeStore.put("ORIGINAL".toByteArray()).value
val sessionId = SessionId("session-multi-workspace")
val firstBound = SessionWorkspaceBoundEvent(
sessionId = sessionId,
workspaceRoot = firstWorkspace.toString(),
allowedPaths = emptyList(),
)
val lastBound = SessionWorkspaceBoundEvent(
sessionId = sessionId,
workspaceRoot = lastWorkspace.toString(),
allowedPaths = emptyList(),
)
val fileWritten = FileWrittenEvent(
invocationId = ToolInvocationId("inv-multi"),
sessionId = sessionId,
path = target.toString(),
preImageHash = preImageHash,
postImageHash = null,
preExisted = true,
timestampMs = 0L,
)
val fakeEventStore = FakeEventStore(
listOf(
storedEvent(firstBound, sessionId, 1L),
storedEvent(lastBound, sessionId, 2L),
storedEvent(fileWritten, sessionId, 3L),
),
)
val service = SessionUndoService(fakeEventStore, fakeStore, setOf(bootDir))
val summary = service.undo(sessionId)
assertEquals(1, summary.reverted)
assertEquals(0, summary.rejected, "file under last workspace must be reverted")
assertEquals("ORIGINAL", Files.readString(target))
}
}
@@ -50,6 +50,7 @@ class DefaultSessionOrchestrator(
graph: WorkflowGraph,
config: OrchestrationConfig,
): WorkflowResult {
val effectives = effectivesFor(config)
log.debug("[Orchestrator] session={} workflow={} start={}", sessionId.value, graph.id, graph.start.value)
emitWorkflowStarted(sessionId, graph, config)
@@ -57,14 +58,14 @@ class DefaultSessionOrchestrator(
val enriched = base.enrich()
// Execute the start stage before entering the step loop
return when (val result = enterStage(enriched, graph.start)) {
is StepResult.Continue -> step(result.ctx)
return when (val result = enterStage(enriched, graph.start, effectives)) {
is StepResult.Continue -> step(result.ctx, effectives)
is StepResult.Terminal -> result.result
}
}
@Suppress("LongMethod")
private tailrec suspend fun step(ctx: EnrichedExecutionContext): WorkflowResult {
private tailrec suspend fun step(ctx: EnrichedExecutionContext, effectives: RunEffectives): WorkflowResult {
log.debug(
"[Orchestrator] step session={} stage={} stageCount={}",
ctx.sessionId.value, ctx.currentStageId.value, ctx.stageCount,
@@ -106,8 +107,8 @@ class DefaultSessionOrchestrator(
enriched.sessionId.value, enriched.currentStageId.value, decision::class.simpleName,
)
return when (decision) {
is TransitionDecision.Move -> when (val r = executeMove(enriched, decision)) {
is StepResult.Continue -> step(r.ctx)
is TransitionDecision.Move -> when (val r = executeMove(enriched, decision, effectives)) {
is StepResult.Continue -> step(r.ctx, effectives)
is StepResult.Terminal -> r.result
}
@@ -144,13 +145,14 @@ class DefaultSessionOrchestrator(
graph: WorkflowGraph,
config: OrchestrationConfig,
): WorkflowResult {
val effectives = effectivesFor(config)
val stageId = orchestrationRepository.getState(sessionId).currentStageId
?: return WorkflowResult.Failed(sessionId, "resume: no currentStageId", retryExhausted = false)
log.info("[Orchestrator] resuming session={} workflow={} stage={}", sessionId.value, graph.id, stageId.value)
val base = ExecutionContext(graph, sessionId, 0, stageId, config, null, null)
val enriched = base.enrich()
return when (val result = enterStage(enriched, stageId)) {
is StepResult.Continue -> step(result.ctx.copy(currentStageId = stageId))
return when (val result = enterStage(enriched, stageId, effectives)) {
is StepResult.Continue -> step(result.ctx.copy(currentStageId = stageId), effectives)
is StepResult.Terminal -> result.result
}
}
@@ -191,6 +193,7 @@ class DefaultSessionOrchestrator(
private suspend fun executeMove(
ctx: EnrichedExecutionContext,
decision: TransitionDecision.Move,
effectives: RunEffectives,
): StepResult {
val nextStageId = decision.to
@@ -201,7 +204,7 @@ class DefaultSessionOrchestrator(
)
}
return when (val result = enterStage(ctx, nextStageId)) {
return when (val result = enterStage(ctx, nextStageId, effectives)) {
is StepResult.Continue -> StepResult.Continue(
result.ctx.copy(
currentStageId = advanceStage(ctx.sessionId, ctx.currentStageId, decision),
@@ -216,13 +219,14 @@ class DefaultSessionOrchestrator(
private suspend fun enterStage(
ctx: EnrichedExecutionContext,
stageId: StageId,
effectives: RunEffectives,
): StepResult {
log.debug("[Orchestrator] executeStage session=${ctx.sessionId.value} stage=${stageId.value}")
while (true) {
if (isCancelled(ctx.sessionId)) {
return StepResult.Terminal(handleCancellation(ctx.sessionId, stageId))
}
when (val result = executeStage(ctx.sessionId, stageId, ctx.graph, ctx.session, ctx.config)) {
when (val result = executeStage(ctx.sessionId, stageId, ctx.graph, ctx.session, ctx.config, effectives)) {
is StageExecutionResult.Success ->
return StepResult.Continue(ctx.copy(stageCount = ctx.stageCount + 1))
@@ -27,4 +27,5 @@ data class OrchestratorEngines(
val toolCallAssessor: ToolCallAssessor? = null,
val workspacePolicy: WorkspacePolicy? = null,
val worldProbe: WorldProbe = FileSystemWorldProbe(),
val workspaceToolRegistryProvider: WorkspaceToolRegistryProvider? = null,
)
@@ -128,7 +128,7 @@ class ReplayOrchestrator(
session: Session,
): ReplayStepResult {
log.debug("[Orchestrator] executeStage session=${ctx.sessionId.value} stage=${stageId.value}")
return when (val result = executeStage(ctx.sessionId, stageId, ctx.graph, session, ctx.config)) {
return when (val result = executeStage(ctx.sessionId, stageId, ctx.graph, session, ctx.config, effectivesFor(ctx.config))) {
is StageExecutionResult.Success -> ReplayStepResult.Continue(
ctx.copy(stageCount = ctx.stageCount + 1),
)
@@ -151,6 +151,7 @@ class ReplayOrchestrator(
stageConfig: StageConfig,
timeoutMs: Long,
responseFormat: ResponseFormat,
effectives: RunEffectives,
): InferenceResult = when (strategy) {
is ReplayStrategy.SkipInference -> {
// bypass router entirely — use recorded artifact
@@ -179,7 +180,7 @@ class ReplayOrchestrator(
},
)
}
else -> super.runInference(sessionId, stageId, contextPack, stageConfig, timeoutMs, responseFormat)
else -> super.runInference(sessionId, stageId, contextPack, stageConfig, timeoutMs, responseFormat, effectives)
}
override suspend fun mapValidationOutcome(
@@ -143,6 +143,7 @@ abstract class SessionOrchestrator(
private val toolCallAssessor: ToolCallAssessor? = engines.toolCallAssessor
private val workspacePolicy: WorkspacePolicy? = engines.workspacePolicy
private val worldProbe: WorldProbe = engines.worldProbe
private val workspaceToolRegistryProvider: WorkspaceToolRegistryProvider? = engines.workspaceToolRegistryProvider
private val inferenceRepository: InferenceRepository = repositories.inferenceRepository
internal val orchestrationRepository: OrchestrationRepository = repositories.orchestrationRepository
protected open val tokenizer: Tokenizer? = null
@@ -165,6 +166,24 @@ abstract class SessionOrchestrator(
abstract suspend fun cancel(sessionId: SessionId)
// --- per-run effective registry + policy ---
internal data class RunEffectives(
val registry: ToolRegistry?,
val executor: ToolExecutor?,
val policy: WorkspacePolicy?,
)
internal fun effectivesFor(config: OrchestrationConfig): RunEffectives {
val wsTools = config.workspace?.let { workspaceToolRegistryProvider?.forWorkspace(it) }
val registry = wsTools?.registry ?: toolRegistry
val executor = wsTools?.executor ?: toolExecutor
val policy = config.workspace
?.let { WorkspacePolicy(it.workspaceRoot, it.privilegedLocations) }
?: workspacePolicy
return RunEffectives(registry, executor, policy)
}
// --- stage execution ---
@Suppress("CyclomaticComplexMethod")
@@ -174,6 +193,7 @@ abstract class SessionOrchestrator(
graph: WorkflowGraph,
session: Session,
config: OrchestrationConfig,
effectives: RunEffectives,
): StageExecutionResult {
val stageConfig = requireNotNull(graph.stages[stageId]) {
"Stage '${stageId.value}' not declared in workflow graph"
@@ -272,7 +292,7 @@ abstract class SessionOrchestrator(
var currentContext = contextPack
var inferenceResult = runInference(
sessionId, stageId, currentContext, stageConfig, config.stageTimeoutMs, responseFormat,
sessionId, stageId, currentContext, stageConfig, config.stageTimeoutMs, responseFormat, effectives,
)
var toolRounds = 0
while (
@@ -287,7 +307,7 @@ abstract class SessionOrchestrator(
break
}
val toolEntries = dispatchToolCalls(sessionId, stageId,
inferenceResult.response.toolCalls, stageConfig)
inferenceResult.response.toolCalls, stageConfig, effectives)
val fatalEntry = toolEntries.firstOrNull { it.content.startsWith("FATAL:") }
if (fatalEntry != null) {
emitProcessResultEvents(sessionId, stageId, stageConfig)
@@ -314,7 +334,7 @@ abstract class SessionOrchestrator(
)
emitContextTruncationIfNeeded(sessionId, stageId, currentContext)
inferenceResult = runInference(
sessionId, stageId, currentContext, stageConfig, config.stageTimeoutMs, responseFormat,
sessionId, stageId, currentContext, stageConfig, config.stageTimeoutMs, responseFormat, effectives,
)
toolRounds++
}
@@ -326,7 +346,7 @@ abstract class SessionOrchestrator(
val validationCtx = ValidationContext(
graph = graph,
sessionState = session.state,
availableTools = toolRegistry?.all()?.map { it.name }?.toSet()
availableTools = effectives.registry?.all()?.map { it.name }?.toSet()
)
when (val outcome = mapValidationOutcome(sessionId, stageId, validationCtx)) {
is StageExecutionResult.Success -> {
@@ -353,8 +373,9 @@ abstract class SessionOrchestrator(
stageId: StageId,
toolCalls: List<ToolCallRequest>,
stageConfig: StageConfig,
effectives: RunEffectives,
): List<ContextEntry> {
val executor = toolExecutor ?: return emptyList()
val executor = effectives.executor ?: return emptyList()
val processResultSlots = stageConfig.produces.filter { it.kind.id == "process_result" }
return toolCalls.flatMap { toolCall ->
val invocationId = ToolInvocationId(UUID.randomUUID().toString())
@@ -372,7 +393,7 @@ abstract class SessionOrchestrator(
toolName = toolCall.function.name,
parameters = parameters,
)
val tool = toolRegistry?.resolve(toolCall.function.name)
val tool = effectives.registry?.resolve(toolCall.function.name)
val tier = tool?.tier ?: Tier.T2
emit(
sessionId,
@@ -386,7 +407,7 @@ abstract class SessionOrchestrator(
),
)
val plane2Risk: RiskSummary? = runPlane2Assessment(
sessionId, stageId, invocationId, toolCall.function.name, request, tool,
sessionId, stageId, invocationId, toolCall.function.name, request, tool, effectives,
)?.let { assessment ->
if (assessment.recommendedAction == RiskAction.BLOCK) {
emit(
@@ -607,9 +628,10 @@ abstract class SessionOrchestrator(
toolName: String,
request: ToolRequest,
tool: Tool?,
effectives: RunEffectives,
): RiskSummary? {
val assessor = toolCallAssessor ?: return null
val policy = workspacePolicy ?: return null
val policy = effectives.policy ?: return null
val assessment = assessor.assess(
ToolCallAssessmentInput(
request = request,
@@ -771,6 +793,7 @@ abstract class SessionOrchestrator(
stageConfig: StageConfig,
timeoutMs: Long,
responseFormat: ResponseFormat = ResponseFormat.Text,
effectives: RunEffectives = RunEffectives(toolRegistry, toolExecutor, workspacePolicy),
): InferenceResult {
val provider = inferenceRouter.route(stageId, stageConfig.requiredCapabilities, stageConfig.modelId)
log.debug(
@@ -786,7 +809,7 @@ abstract class SessionOrchestrator(
generationConfig = stageConfig.generationConfig,
responseFormat = responseFormat,
tools = stageConfig.allowedTools
.mapNotNull { toolRegistry?.resolve(it) }
.mapNotNull { effectives.registry?.resolve(it) }
.map { tool ->
ToolDefinition(
function = ToolFunction(
@@ -0,0 +1,13 @@
package com.correx.core.kernel.orchestration
import com.correx.core.tools.contract.ToolExecutor
import com.correx.core.tools.registry.ToolRegistry
data class WorkspaceTools(
val registry: ToolRegistry,
val executor: ToolExecutor,
)
fun interface WorkspaceToolRegistryProvider {
fun forWorkspace(workspace: WorkspaceContext): WorkspaceTools
}
+2
View File
@@ -25,4 +25,6 @@ dependencies {
testImplementation(project(":testing:contracts"))
testImplementation(testFixtures(project(":testing:contracts")))
testImplementation("org.jetbrains.kotlinx:kotlinx-coroutines-test:1.9.0")
testImplementation(project(":infrastructure:tools"))
testImplementation(project(":infrastructure:tools:filesystem"))
}
@@ -0,0 +1,375 @@
import com.correx.core.approvals.ApprovalOutcome
import com.correx.core.approvals.ApprovalProjector
import com.correx.core.approvals.ApprovalStatus
import com.correx.core.approvals.DefaultApprovalReducer
import com.correx.core.approvals.DefaultApprovalRepository
import com.correx.core.approvals.domain.ApprovalEngine
import com.correx.core.approvals.model.ApprovalContext
import com.correx.core.approvals.model.ApprovalDecision
import com.correx.core.approvals.model.ApprovalGrant
import com.correx.core.approvals.model.DomainApprovalRequest
import com.correx.core.artifacts.DefaultArtifactReducer
import com.correx.core.events.events.ToolCallAssessedEvent
import com.correx.core.events.execution.RetryPolicy
import com.correx.core.events.risk.RiskAction
import com.correx.core.events.types.ProviderId
import com.correx.core.events.types.SessionId
import com.correx.core.events.types.StageId
import com.correx.core.events.types.TransitionId
import com.correx.core.inference.CapabilityScore
import com.correx.core.inference.FinishReason
import com.correx.core.inference.InferenceProvider
import com.correx.core.inference.InferenceRepository
import com.correx.core.inference.InferenceRequest
import com.correx.core.inference.InferenceResponse
import com.correx.core.inference.InferenceRouter
import com.correx.core.inference.InferenceState
import com.correx.core.inference.ModelCapability
import com.correx.core.inference.ProviderHealth
import com.correx.core.inference.Tokenizer
import com.correx.core.inference.ToolCallFunction
import com.correx.core.inference.ToolCallRequest
import com.correx.core.inference.TokenUsage
import com.correx.core.kernel.orchestration.DefaultOrchestrationReducer
import com.correx.core.kernel.orchestration.DefaultSessionOrchestrator
import com.correx.core.kernel.orchestration.OrchestrationConfig
import com.correx.core.kernel.orchestration.OrchestrationProjector
import com.correx.core.kernel.orchestration.OrchestrationRepository
import com.correx.core.kernel.orchestration.OrchestratorEngines
import com.correx.core.kernel.orchestration.OrchestratorRepositories
import com.correx.core.kernel.orchestration.WorkspaceContext
import com.correx.core.kernel.orchestration.WorkspaceTools
import com.correx.core.kernel.orchestration.WorkspaceToolRegistryProvider
import com.correx.core.kernel.retry.DefaultRetryCoordinator
import com.correx.core.risk.DefaultRiskAssessor
import com.correx.core.sessions.DefaultSessionReducer
import com.correx.core.sessions.DefaultSessionRepository
import com.correx.core.sessions.SessionProjector
import com.correx.core.sessions.projections.replay.DefaultEventReplayer
import com.correx.core.sessions.projections.replay.EventReplayer
import com.correx.core.toolintent.ToolCallAssessor
import com.correx.core.toolintent.WorkspacePolicy
import com.correx.core.toolintent.rules.PathContainmentRule
import com.correx.core.tools.registry.ToolRegistry
import com.correx.core.transitions.graph.StageConfig
import com.correx.core.transitions.graph.TransitionEdge
import com.correx.core.transitions.graph.WorkflowGraph
import com.correx.core.transitions.resolution.DefaultTransitionResolver
import com.correx.core.validation.pipeline.ValidationPipeline
import com.correx.infrastructure.persistence.InMemoryEventStore
import com.correx.infrastructure.persistence.artifact.LiveArtifactRepository
import com.correx.infrastructure.tools.DefaultToolRegistry
import com.correx.infrastructure.tools.DispatchingToolExecutor
import com.correx.infrastructure.tools.FileEditConfig
import com.correx.infrastructure.tools.FileReadConfig
import com.correx.infrastructure.tools.FileWriteConfig
import com.correx.infrastructure.tools.ShellConfig
import com.correx.infrastructure.tools.ToolConfig
import com.correx.infrastructure.tools.buildTools
import com.correx.testing.contracts.fixtures.artifactstore.NoopArtifactStore
import com.correx.testing.fixtures.context.ContextFixtures
import com.correx.testing.fixtures.cyclePolicyMissingValidator
import com.correx.testing.fixtures.inference.MockTokenizer
import kotlinx.coroutines.runBlocking
import kotlinx.datetime.Instant
import java.nio.file.Files
import kotlin.test.Test
import kotlin.test.assertFalse
import kotlin.test.assertNotNull
import kotlin.test.assertTrue
@Suppress("LongMethod")
class WorkspaceScopedToolRegistryTest {
private inner class FileWriteCallingProvider(
private val targetPath: String,
private val content: String,
) : InferenceProvider {
override val id = ProviderId("file-write-caller")
override val name = "file-write-caller"
override val tokenizer: Tokenizer = MockTokenizer()
private var callCount = 0
override suspend fun infer(request: InferenceRequest): InferenceResponse {
callCount++
return if (callCount == 1) {
InferenceResponse(
requestId = request.requestId,
text = "",
finishReason = FinishReason.ToolCall,
tokensUsed = TokenUsage(10, 5),
latencyMs = 0,
toolCalls = listOf(
ToolCallRequest(
id = "tc-write",
function = ToolCallFunction(
name = "file_write",
arguments = """{"path":"$targetPath","operation":"write","content":"$content"}""",
),
),
),
)
} else {
InferenceResponse(
requestId = request.requestId,
text = "done",
finishReason = FinishReason.Stop,
tokensUsed = TokenUsage(10, 5),
latencyMs = 0,
)
}
}
override suspend fun healthCheck(): ProviderHealth = ProviderHealth.Healthy
override fun capabilities(): Set<CapabilityScore> = setOf(CapabilityScore(ModelCapability.General, 1.0))
}
private fun buildToolsForWorkspace(workspace: WorkspaceContext): WorkspaceTools {
val registry = DefaultToolRegistry.build(
ToolConfig(
shell = ShellConfig(enabled = false),
fileRead = FileReadConfig(enabled = false),
fileWrite = FileWriteConfig(
enabled = true,
allowedPaths = workspace.allowedPaths,
workingDir = workspace.workingDir,
),
fileEdit = FileEditConfig(enabled = false),
).buildTools(),
)
return WorkspaceTools(registry = registry, executor = DispatchingToolExecutor(registry))
}
private fun buildOrchestrator(
workspaceContext: WorkspaceContext?,
bootRegistry: ToolRegistry,
provider: InferenceProvider,
): Pair<DefaultSessionOrchestrator, InMemoryEventStore> {
val eventStore = InMemoryEventStore()
val artifactStore = NoopArtifactStore()
val assessor = ToolCallAssessor(listOf(PathContainmentRule()))
val bootPolicyRoot = workspaceContext?.workspaceRoot
?: java.nio.file.Path.of(System.getProperty("java.io.tmpdir"))
val bootPolicy = WorkspacePolicy(bootPolicyRoot)
val inferenceRouter = object : InferenceRouter {
override suspend fun route(stageId: StageId, requiredCapabilities: Set<ModelCapability>) = provider
}
val approvalRepository = DefaultApprovalRepository(
DefaultEventReplayer(eventStore, ApprovalProjector(DefaultApprovalReducer())),
)
val repositories = OrchestratorRepositories(
eventStore = eventStore,
inferenceRepository = InferenceRepository(object : EventReplayer<InferenceState> {
override fun rebuild(sessionId: SessionId) = InferenceState()
}),
orchestrationRepository = OrchestrationRepository(
DefaultEventReplayer(eventStore, OrchestrationProjector(DefaultOrchestrationReducer())),
),
sessionRepository = DefaultSessionRepository(
DefaultEventReplayer(eventStore, SessionProjector(DefaultSessionReducer())),
),
artifactRepository = LiveArtifactRepository(eventStore, DefaultArtifactReducer()),
approvalRepository = approvalRepository,
)
val workspaceProvider = WorkspaceToolRegistryProvider { ws -> buildToolsForWorkspace(ws) }
val autoApproveEngine = object : ApprovalEngine {
override fun evaluate(
request: DomainApprovalRequest,
context: ApprovalContext,
grants: List<ApprovalGrant>,
now: Instant,
) = ApprovalDecision(
id = null,
requestId = request.id,
outcome = ApprovalOutcome.AUTO_APPROVED,
state = ApprovalStatus.COMPLETED,
tier = request.tier,
contextSnapshot = context,
resolutionTimestamp = now,
reason = "test-auto-approve",
)
}
val engines = OrchestratorEngines(
transitionResolver = DefaultTransitionResolver { _, _ -> true },
contextPackBuilder = ContextFixtures.simpleBuilder(),
inferenceRouter = inferenceRouter,
validationPipeline = ValidationPipeline(validators = listOf(cyclePolicyMissingValidator())),
approvalEngine = autoApproveEngine,
riskAssessor = DefaultRiskAssessor(),
toolRegistry = bootRegistry,
toolExecutor = DispatchingToolExecutor(bootRegistry),
toolCallAssessor = assessor,
workspacePolicy = bootPolicy,
workspaceToolRegistryProvider = workspaceProvider,
)
return DefaultSessionOrchestrator(
repositories = repositories,
engines = engines,
retryCoordinator = DefaultRetryCoordinator(eventStore),
artifactStore = artifactStore,
) to eventStore
}
private fun singleStageGraph(): WorkflowGraph = WorkflowGraph(
id = "ws-test",
stages = mapOf(
StageId("A") to StageConfig(allowedTools = setOf("file_write")),
),
transitions = setOf(
TransitionEdge(TransitionId("t1"), StageId("A"), StageId("done"), condition = { true }),
),
start = StageId("A"),
)
@Test
fun `write inside workspace A succeeds and file exists`(): Unit = runBlocking {
val workspaceA = Files.createTempDirectory("ws-a")
val targetFile = workspaceA.resolve("output.txt")
val workspace = WorkspaceContext(
workspaceRoot = workspaceA,
workingDir = workspaceA,
allowedPaths = setOf(workspaceA),
)
val provider = FileWriteCallingProvider(targetFile.toString(), "hello workspace A")
val bootRegistry = DefaultToolRegistry.build(emptyList<com.correx.core.tools.contract.Tool>())
val (orchestrator, _) = buildOrchestrator(workspace, bootRegistry, provider)
val config = OrchestrationConfig(
retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0),
workspace = workspace,
)
orchestrator.run(SessionId("ws-a-ok"), singleStageGraph(), config)
assertTrue(Files.exists(targetFile), "file_write inside workspace A must create the file")
}
@Test
fun `write outside workspace A is rejected by jail`(): Unit = runBlocking {
val workspaceA = Files.createTempDirectory("ws-a-jail")
val outsideDir = Files.createTempDirectory("ws-outside")
val targetFile = outsideDir.resolve("escape.txt")
val workspace = WorkspaceContext(
workspaceRoot = workspaceA,
workingDir = workspaceA,
allowedPaths = setOf(workspaceA),
)
val provider = FileWriteCallingProvider(targetFile.toString(), "should be blocked")
val bootRegistry = DefaultToolRegistry.build(emptyList<com.correx.core.tools.contract.Tool>())
val (orchestrator, _) = buildOrchestrator(workspace, bootRegistry, provider)
val config = OrchestrationConfig(
retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0),
workspace = workspace,
)
orchestrator.run(SessionId("ws-a-jail"), singleStageGraph(), config)
assertFalse(Files.exists(targetFile), "file_write outside workspace A must NOT create the file")
}
@Test
fun `two sessions with different workspaces are independently jailed`(): Unit = runBlocking {
val workspaceA = Files.createTempDirectory("ws-aa")
val workspaceB = Files.createTempDirectory("ws-bb")
val fileA = workspaceA.resolve("a.txt")
val fileB = workspaceB.resolve("b.txt")
val ctxA = WorkspaceContext(
workspaceRoot = workspaceA,
workingDir = workspaceA,
allowedPaths = setOf(workspaceA),
)
val ctxB = WorkspaceContext(
workspaceRoot = workspaceB,
workingDir = workspaceB,
allowedPaths = setOf(workspaceB),
)
val bootRegistry = DefaultToolRegistry.build(emptyList<com.correx.core.tools.contract.Tool>())
val (orchA, _) = buildOrchestrator(ctxA, bootRegistry, FileWriteCallingProvider(fileA.toString(), "session A"))
val (orchB, _) = buildOrchestrator(ctxB, bootRegistry, FileWriteCallingProvider(fileB.toString(), "session B"))
val configA = OrchestrationConfig(retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0), workspace = ctxA)
val configB = OrchestrationConfig(retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0), workspace = ctxB)
orchA.run(SessionId("session-a"), singleStageGraph(), configA)
orchB.run(SessionId("session-b"), singleStageGraph(), configB)
assertTrue(Files.exists(fileA), "session A must write to workspace A")
assertTrue(Files.exists(fileB), "session B must write to workspace B")
assertFalse(Files.exists(workspaceB.resolve("a.txt")), "session A must NOT write to workspace B")
assertFalse(Files.exists(workspaceA.resolve("b.txt")), "session B must NOT write to workspace A")
}
@Test
fun `null workspace uses boot registry — regression guard`(): Unit = runBlocking {
val bootWorkspace = Files.createTempDirectory("ws-boot")
val targetFile = bootWorkspace.resolve("boot.txt")
val bootRegistry = DefaultToolRegistry.build(
ToolConfig(
shell = ShellConfig(enabled = false),
fileRead = FileReadConfig(enabled = false),
fileWrite = FileWriteConfig(
enabled = true,
allowedPaths = setOf(bootWorkspace),
workingDir = bootWorkspace,
),
fileEdit = FileEditConfig(enabled = false),
).buildTools(),
)
val provider = FileWriteCallingProvider(targetFile.toString(), "boot content")
val (orchestrator, _) = buildOrchestrator(null, bootRegistry, provider)
val config = OrchestrationConfig(
retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0),
workspace = null,
)
orchestrator.run(SessionId("null-workspace"), singleStageGraph(), config)
assertTrue(Files.exists(targetFile), "null-workspace run must use boot registry — file must be written")
}
@Test
fun `plane-2 assessment uses session workspace policy not boot policy`(): Unit = runBlocking {
val workspaceA = Files.createTempDirectory("ws-plane2")
val outsideDir = Files.createTempDirectory("ws-plane2-outside")
val targetFile = outsideDir.resolve("plane2.txt")
val workspace = WorkspaceContext(
workspaceRoot = workspaceA,
workingDir = workspaceA,
allowedPaths = setOf(workspaceA),
)
val provider = FileWriteCallingProvider(targetFile.toString(), "test")
val bootRegistry = DefaultToolRegistry.build(emptyList<com.correx.core.tools.contract.Tool>())
val (orchestrator, eventStore) = buildOrchestrator(workspace, bootRegistry, provider)
val config = OrchestrationConfig(
retryPolicy = RetryPolicy(maxAttempts = 1, backoffMs = 0),
workspace = workspace,
)
val sessionId = SessionId("plane2-session")
orchestrator.run(sessionId, singleStageGraph(), config)
val assessed = eventStore.read(sessionId)
.mapNotNull { it.payload as? ToolCallAssessedEvent }
.firstOrNull()
assertNotNull(assessed, "ToolCallAssessedEvent must be emitted when path is outside session workspace")
assertTrue(
assessed.disposition != RiskAction.PROCEED,
"plane-2 disposition for out-of-workspace write must not be PROCEED, got ${assessed.disposition}",
)
}
}