feat: implement CreateGrant handler and grant-aware approval engine

- GlobalStreamHandler.handleCreateGrant validates scope (SESSION/STAGE)
  and sends ProtocolError to client on validation failure instead of
  silent log.warn + drop.
- Derive event stageId directly from GrantScope type, removing
  redundant stageIdForEvent tuple.
- SessionOrchestrator integrates ApprovalEngine to check active grants
  before requesting user approval for T2+ tool calls.
- ContextEntry gains EntryRole (SYSTEM/USER/ASSISTANT/TOOL) field for
  proper chat message role mapping.
- RouterContextBuilder.build is now suspend; uses Tokenizer for
  accurate token estimation with fallback to character-based estimate.
- LlamaCppInferenceProvider maps EntryRole to ChatMessage role instead
  of heuristic layer-based inference.
This commit is contained in:
2026-05-28 14:06:58 +04:00
parent 92bea6c2c4
commit e05532e7b2
11 changed files with 316 additions and 133 deletions
@@ -56,7 +56,8 @@ fun main() {
logStoresInfo(eventStore, artifactStore) logStoresInfo(eventStore, artifactStore)
val infraRegistry = InfrastructureModule.createProviderRegistry(listOf(buildLlamaProvider())) val llamaProvider = buildLlamaProvider()
val infraRegistry = InfrastructureModule.createProviderRegistry(listOf(llamaProvider))
val modelId = System.getenv("CORREX_MODEL_ID") ?: "default" val modelId = System.getenv("CORREX_MODEL_ID") ?: "default"
val modelPath = System.getenv("CORREX_MODEL_PATH") ?: "(not set)" val modelPath = System.getenv("CORREX_MODEL_PATH") ?: "(not set)"
val llamaUrl = System.getenv("CORREX_LLAMA_URL") ?: "http://127.0.0.1:10000" val llamaUrl = System.getenv("CORREX_LLAMA_URL") ?: "http://127.0.0.1:10000"
@@ -109,6 +110,7 @@ fun main() {
engines = engines, engines = engines,
retryCoordinator = DefaultRetryCoordinator(eventStore), retryCoordinator = DefaultRetryCoordinator(eventStore),
artifactStore = artifactStore, artifactStore = artifactStore,
tokenizer = llamaProvider.tokenizer,
) )
val defaultOrchestrationConfig = OrchestrationConfig( val defaultOrchestrationConfig = OrchestrationConfig(
sandboxRoot = sandboxRoot, sandboxRoot = sandboxRoot,
@@ -119,6 +121,7 @@ fun main() {
eventStore = eventStore, eventStore = eventStore,
inferenceRouter = inferenceRouter, inferenceRouter = inferenceRouter,
config = routerConfig, config = routerConfig,
tokenizer = llamaProvider.tokenizer,
) )
val module = ServerModule( val module = ServerModule(
orchestrator = orchestrator, orchestrator = orchestrator,
@@ -1,8 +1,12 @@
package com.correx.apps.server.protocol package com.correx.apps.server.protocol
import com.correx.core.approvals.GrantScope
import com.correx.core.approvals.Tier
import com.correx.core.events.types.ApprovalRequestId 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.router.ChatMode import com.correx.core.router.ChatMode
import kotlinx.datetime.Instant
import kotlinx.serialization.Serializable import kotlinx.serialization.Serializable
@Serializable @Serializable
@@ -11,6 +15,12 @@ enum class ApprovalDecision {
REJECT, REJECT,
} }
@Serializable
enum class GrantScopeDto {
SESSION,
STAGE,
}
@Serializable @Serializable
sealed class ClientMessage { sealed class ClientMessage {
@Serializable @Serializable
@@ -29,6 +39,16 @@ sealed class ClientMessage {
val steeringNote: String?, val steeringNote: String?,
) : ClientMessage() ) : ClientMessage()
@Serializable
data class CreateGrant(
val sessionId: SessionId,
val scope: GrantScopeDto,
val stageId: StageId? = null,
val permittedTiers: List<Tier>,
val reason: String,
val expiresAt: Instant? = null,
) : ClientMessage()
@Serializable @Serializable
data class Ping(val timestamp: Long) : ClientMessage() data class Ping(val timestamp: Long) : ClientMessage()
@@ -4,16 +4,20 @@ import com.correx.apps.server.ServerModule
import com.correx.apps.server.bridge.DomainEventMapper import com.correx.apps.server.bridge.DomainEventMapper
import com.correx.apps.server.bridge.SessionEventBridge import com.correx.apps.server.bridge.SessionEventBridge
import com.correx.apps.server.protocol.ClientMessage import com.correx.apps.server.protocol.ClientMessage
import com.correx.apps.server.protocol.GrantScopeDto
import com.correx.apps.server.protocol.ProtocolSerializer import com.correx.apps.server.protocol.ProtocolSerializer
import com.correx.apps.server.protocol.ProviderHealthDto import com.correx.apps.server.protocol.ProviderHealthDto
import com.correx.apps.server.protocol.ServerMessage import com.correx.apps.server.protocol.ServerMessage
import com.correx.apps.server.protocol.WorkflowDto import com.correx.apps.server.protocol.WorkflowDto
import com.correx.apps.server.protocol.StageToolDecl import com.correx.apps.server.protocol.StageToolDecl
import com.correx.apps.server.protocol.ToolDecl import com.correx.apps.server.protocol.ToolDecl
import com.correx.core.approvals.GrantScope
import com.correx.core.events.events.ApprovalGrantCreatedEvent
import com.correx.core.events.events.EventMetadata import com.correx.core.events.events.EventMetadata
import com.correx.core.events.events.NewEvent import com.correx.core.events.events.NewEvent
import com.correx.core.events.events.StoredEvent import com.correx.core.events.events.StoredEvent
import com.correx.core.events.events.WorkflowFailedEvent import com.correx.core.events.events.WorkflowFailedEvent
import com.correx.core.events.types.GrantId
import com.correx.core.events.stores.EventStore import com.correx.core.events.stores.EventStore
import com.correx.core.events.types.EventId import com.correx.core.events.types.EventId
import com.correx.core.events.types.SessionId import com.correx.core.events.types.SessionId
@@ -140,6 +144,7 @@ class GlobalStreamHandler(private val module: ServerModule) {
is ClientMessage.CancelSession -> module.orchestrator.cancel(msg.sessionId) is ClientMessage.CancelSession -> module.orchestrator.cancel(msg.sessionId)
is ClientMessage.ResumeSession -> session.send(encodeError("ResumeSession not supported")) is ClientMessage.ResumeSession -> session.send(encodeError("ResumeSession not supported"))
is ClientMessage.ApprovalResponse -> handleApprovalResponse(msg, sendFrame) is ClientMessage.ApprovalResponse -> handleApprovalResponse(msg, sendFrame)
is ClientMessage.CreateGrant -> handleCreateGrant(msg, sendFrame)
is ClientMessage.ChatInput -> Unit // handler deferred to subsequent task is ClientMessage.ChatInput -> Unit // handler deferred to subsequent task
} }
} }
@@ -163,6 +168,12 @@ class GlobalStreamHandler(private val module: ServerModule) {
module.approvalCoordinator.handleResponse(msg, scopeSessionId)?.let { sendFrame(it) } module.approvalCoordinator.handleResponse(msg, scopeSessionId)?.let { sendFrame(it) }
} }
private fun errorResponse(message: String) = ServerMessage.ProtocolError(
message = message,
sequence = null,
sessionSequence = null,
)
private fun encodeError(message: String): Frame.Text = private fun encodeError(message: String): Frame.Text =
Frame.Text( Frame.Text(
ProtocolSerializer.encodeServerMessage( ProtocolSerializer.encodeServerMessage(
@@ -174,6 +185,49 @@ class GlobalStreamHandler(private val module: ServerModule) {
), ),
) )
private suspend fun handleCreateGrant(
msg: ClientMessage.CreateGrant,
sendFrame: suspend (ServerMessage) -> Unit,
) {
val scope = when (msg.scope) {
GrantScopeDto.SESSION -> {
if (msg.stageId != null) {
sendFrame(errorResponse("CreateGrant: SESSION scope must not include stageId"))
return
}
GrantScope.SESSION
}
GrantScopeDto.STAGE -> {
val sid = msg.stageId ?: run {
sendFrame(errorResponse("CreateGrant: STAGE scope requires stageId"))
return
}
GrantScope.STAGE(sid)
}
}
val event = NewEvent(
metadata = EventMetadata(
eventId = EventId(UUID.randomUUID().toString()),
sessionId = msg.sessionId,
timestamp = Clock.System.now(),
schemaVersion = 1,
causationId = null,
correlationId = null,
),
payload = ApprovalGrantCreatedEvent(
grantId = GrantId(UUID.randomUUID().toString()),
scope = scope,
permittedTiers = msg.permittedTiers.toSet(),
reason = msg.reason,
expiresAt = msg.expiresAt,
sessionId = msg.sessionId,
stageId = (scope as? GrantScope.STAGE)?.stageId,
projectId = null,
),
)
module.eventStore.append(event)
}
private suspend fun handleStartSession( private suspend fun handleStartSession(
session: DefaultWebSocketServerSession, session: DefaultWebSocketServerSession,
msg: ClientMessage.StartSession, msg: ClientMessage.StartSession,
@@ -3,6 +3,8 @@ package com.correx.core.context.model
import com.correx.core.events.types.ContextEntryId import com.correx.core.events.types.ContextEntryId
import kotlinx.serialization.Serializable import kotlinx.serialization.Serializable
enum class EntryRole { SYSTEM, USER, ASSISTANT, TOOL }
@Serializable @Serializable
data class ContextEntry( data class ContextEntry(
val id: ContextEntryId, val id: ContextEntryId,
@@ -10,5 +12,6 @@ data class ContextEntry(
val content: String, val content: String,
val sourceType: String, val sourceType: String,
val sourceId: String, val sourceId: String,
val tokenEstimate: Int val tokenEstimate: Int,
val role: EntryRole = EntryRole.USER,
) )
@@ -3,6 +3,7 @@ package com.correx.core.kernel.orchestration
import com.correx.core.approvals.ApprovalOutcome import com.correx.core.approvals.ApprovalOutcome
import com.correx.core.approvals.ApprovalStatus import com.correx.core.approvals.ApprovalStatus
import com.correx.core.approvals.model.ApprovalDecision import com.correx.core.approvals.model.ApprovalDecision
import com.correx.core.inference.Tokenizer
import com.correx.core.artifacts.ArtifactState import com.correx.core.artifacts.ArtifactState
import com.correx.core.artifactstore.ArtifactStore import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.events.ApprovalDecisionResolvedEvent import com.correx.core.events.events.ApprovalDecisionResolvedEvent
@@ -38,7 +39,9 @@ class DefaultSessionOrchestrator(
engines: OrchestratorEngines, engines: OrchestratorEngines,
private val retryCoordinator: RetryCoordinator, private val retryCoordinator: RetryCoordinator,
artifactStore: ArtifactStore, artifactStore: ArtifactStore,
tokenizer: Tokenizer? = null,
) : SessionOrchestrator(repositories, engines, artifactStore), ApprovalGateway { ) : SessionOrchestrator(repositories, engines, artifactStore), ApprovalGateway {
override val tokenizer: Tokenizer? = tokenizer
override val cancellations: ConcurrentHashMap<SessionId, AtomicBoolean> = override val cancellations: ConcurrentHashMap<SessionId, AtomicBoolean> =
ConcurrentHashMap<SessionId, AtomicBoolean>() ConcurrentHashMap<SessionId, AtomicBoolean>()
@@ -3,7 +3,12 @@ package com.correx.core.kernel.orchestration
import com.correx.core.approvals.ApprovalOutcome import com.correx.core.approvals.ApprovalOutcome
import com.correx.core.approvals.ApprovalStatus import com.correx.core.approvals.ApprovalStatus
import com.correx.core.approvals.Tier import com.correx.core.approvals.Tier
import com.correx.core.approvals.DefaultApprovalRepository
import com.correx.core.approvals.domain.ApprovalEngine
import com.correx.core.approvals.isAtMost
import com.correx.core.approvals.model.ApprovalContext
import com.correx.core.approvals.model.ApprovalDecision import com.correx.core.approvals.model.ApprovalDecision
import com.correx.core.approvals.model.ApprovalScopeIdentity
import com.correx.core.approvals.model.DomainApprovalRequest import com.correx.core.approvals.model.DomainApprovalRequest
import com.correx.core.artifacts.ArtifactState import com.correx.core.artifacts.ArtifactState
import com.correx.core.artifacts.kind.JsonSchema import com.correx.core.artifacts.kind.JsonSchema
@@ -14,6 +19,7 @@ import com.correx.core.context.model.ContextEntry
import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextLayer
import com.correx.core.context.model.ContextPack import com.correx.core.context.model.ContextPack
import com.correx.core.context.model.TokenBudget import com.correx.core.context.model.TokenBudget
import com.correx.core.context.model.EntryRole
import com.correx.core.events.events.ApprovalDecisionResolvedEvent import com.correx.core.events.events.ApprovalDecisionResolvedEvent
import com.correx.core.events.events.ApprovalRequestedEvent import com.correx.core.events.events.ApprovalRequestedEvent
import com.correx.core.events.events.ArtifactCreatedEvent import com.correx.core.events.events.ArtifactCreatedEvent
@@ -55,8 +61,10 @@ import com.correx.core.inference.InferenceRequest
import com.correx.core.inference.InferenceResponse import com.correx.core.inference.InferenceResponse
import com.correx.core.inference.InferenceRouter import com.correx.core.inference.InferenceRouter
import com.correx.core.inference.ResponseFormat import com.correx.core.inference.ResponseFormat
import com.correx.core.inference.Tokenizer
import com.correx.core.inference.ToolCallRequest import com.correx.core.inference.ToolCallRequest
import com.correx.core.inference.ToolDefinition import com.correx.core.inference.ToolDefinition
import com.correx.core.sessions.ApprovalMode
import com.correx.core.inference.ToolFunction import com.correx.core.inference.ToolFunction
import com.correx.core.kernel.execution.WorkflowResult import com.correx.core.kernel.execution.WorkflowResult
import com.correx.core.risk.RiskAssessor import com.correx.core.risk.RiskAssessor
@@ -115,6 +123,9 @@ abstract class SessionOrchestrator(
private val toolRegistry: ToolRegistry? = engines.toolRegistry private val toolRegistry: ToolRegistry? = engines.toolRegistry
private val inferenceRepository: InferenceRepository = repositories.inferenceRepository private val inferenceRepository: InferenceRepository = repositories.inferenceRepository
internal val orchestrationRepository: OrchestrationRepository = repositories.orchestrationRepository internal val orchestrationRepository: OrchestrationRepository = repositories.orchestrationRepository
protected open val tokenizer: Tokenizer? = null
private val approvalEngine: ApprovalEngine = engines.approvalEngine
private val approvalRepository: DefaultApprovalRepository = repositories.approvalRepository
internal abstract val cancellations: ConcurrentHashMap<SessionId, AtomicBoolean> internal abstract val cancellations: ConcurrentHashMap<SessionId, AtomicBoolean>
internal val pendingApprovals: ConcurrentHashMap<ApprovalRequestId, CompletableDeferred<ApprovalDecision>> = internal val pendingApprovals: ConcurrentHashMap<ApprovalRequestId, CompletableDeferred<ApprovalDecision>> =
ConcurrentHashMap() ConcurrentHashMap()
@@ -162,7 +173,8 @@ abstract class SessionOrchestrator(
content = text, content = text,
sourceType = "systemPrompt", sourceType = "systemPrompt",
sourceId = stageId.value, sourceId = stageId.value,
tokenEstimate = text.length / 4, tokenEstimate = estimateTokens(text),
role = EntryRole.SYSTEM,
), ),
) )
} ?: config.defaultSystemPromptPath } ?: config.defaultSystemPromptPath
@@ -182,7 +194,8 @@ abstract class SessionOrchestrator(
content = text, content = text,
sourceType = "systemPrompt", sourceType = "systemPrompt",
sourceId = stageId.value, sourceId = stageId.value,
tokenEstimate = text.length / 4, tokenEstimate = estimateTokens(text),
role = EntryRole.SYSTEM,
), ),
) )
} ?: emptyList() } ?: emptyList()
@@ -206,7 +219,8 @@ abstract class SessionOrchestrator(
content = text, content = text,
sourceType = "agentPrompt", sourceType = "agentPrompt",
sourceId = stageId.value, sourceId = stageId.value,
tokenEstimate = text.length / 4, tokenEstimate = estimateTokens(text),
role = EntryRole.USER,
), ),
) )
} ?: emptyList() } ?: emptyList()
@@ -333,11 +347,15 @@ abstract class SessionOrchestrator(
request = request, request = request,
), ),
) )
val requiresApproval = when (tier) { if (tier.isAtMost(Tier.T1)) {
Tier.T0, Tier.T1 -> false // no approval needed
Tier.T2, Tier.T3, Tier.T4 -> true } else {
} val approvalState = approvalRepository.getApprovalState(sessionId)
if (requiresApproval) { val activeGrants = approvalState.grants.values.toList()
val approvalCtx = ApprovalContext(
identity = ApprovalScopeIdentity(sessionId, stageId, projectId = null),
mode = ApprovalMode.PROMPT,
)
val requestId = ApprovalRequestId(UUID.randomUUID().toString()) val requestId = ApprovalRequestId(UUID.randomUUID().toString())
val domainRequest = DomainApprovalRequest( val domainRequest = DomainApprovalRequest(
id = requestId, id = requestId,
@@ -348,59 +366,101 @@ abstract class SessionOrchestrator(
toolName = toolCall.function.name, toolName = toolCall.function.name,
preview = toolCall.function.arguments.take(200), preview = toolCall.function.arguments.take(200),
) )
emit(sessionId, OrchestrationPausedEvent(sessionId, stageId, "APPROVAL_PENDING")) val engineDecision = approvalEngine.evaluate(
emit( domainRequest, approvalCtx, activeGrants, Clock.System.now(),
sessionId,
ApprovalRequestedEvent(
requestId = requestId,
tier = tier,
validationReportId = domainRequest.validationReportId,
riskSummaryId = null,
sessionId = sessionId,
stageId = stageId,
projectId = null,
toolName = toolCall.function.name,
preview = toolCall.function.arguments.take(200),
),
) )
val deferred = CompletableDeferred<ApprovalDecision>() if (engineDecision.state == ApprovalStatus.COMPLETED) {
pendingApprovals[requestId] = deferred emitDecisionResolved(sessionId, domainRequest, engineDecision)
val decision = try { if (!engineDecision.isApproved) {
deferred.await() val rejectReason = engineDecision.reason ?: "denied"
} finally { emit(sessionId, ToolExecutionRejectedEvent(
pendingApprovals.remove(requestId) invocationId = invocationId,
} sessionId = sessionId,
emitDecisionResolved(sessionId, domainRequest, decision) toolName = toolCall.function.name,
if (!decision.isApproved) { tier = tier,
val rejectReason = decision.reason ?: "approval denied" reason = rejectReason,
emit(sessionId, ToolExecutionRejectedEvent( ))
invocationId = invocationId, val sourceId = toolCall.id ?: invocationId.value
sessionId = sessionId, return@flatMap listOf(
toolName = toolCall.function.name, ContextEntry(
tier = tier, id = ContextEntryId(UUID.randomUUID().toString()),
reason = rejectReason, layer = ContextLayer.L2,
)) sourceType = "assistantToolCall",
sourceId = sourceId,
content = Json.encodeToString(ToolCallRequest.serializer(), toolCall),
tokenEstimate = estimateTokens(toolCall.function.arguments),
role = EntryRole.ASSISTANT,
),
ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "toolResult",
sourceId = sourceId,
content = "ERROR: $rejectReason",
tokenEstimate = estimateTokens(rejectReason),
role = EntryRole.TOOL,
),
)
}
// grant auto-approved — fall through to execute
} else {
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,
toolName = toolCall.function.name,
preview = toolCall.function.arguments.take(200),
),
)
val deferred = CompletableDeferred<ApprovalDecision>()
pendingApprovals[requestId] = deferred
val userDecision = try {
deferred.await()
} finally {
pendingApprovals.remove(requestId)
}
emitDecisionResolved(sessionId, domainRequest, userDecision)
if (!userDecision.isApproved) {
val rejectReason = userDecision.reason ?: "approval denied"
emit(sessionId, ToolExecutionRejectedEvent(
invocationId = invocationId,
sessionId = sessionId,
toolName = toolCall.function.name,
tier = tier,
reason = rejectReason,
))
emit(sessionId, OrchestrationResumedEvent(sessionId, stageId))
val sourceId = toolCall.id ?: invocationId.value
return@flatMap listOf(
ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "assistantToolCall",
sourceId = sourceId,
content = Json.encodeToString(ToolCallRequest.serializer(), toolCall),
tokenEstimate = estimateTokens(toolCall.function.arguments),
role = EntryRole.ASSISTANT,
),
ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "toolResult",
sourceId = sourceId,
content = "ERROR: $rejectReason",
tokenEstimate = estimateTokens(rejectReason),
role = EntryRole.TOOL,
),
)
}
emit(sessionId, OrchestrationResumedEvent(sessionId, stageId)) emit(sessionId, OrchestrationResumedEvent(sessionId, stageId))
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 resultEntry = ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()),
layer = ContextLayer.L2,
sourceType = "toolResult",
sourceId = sourceId,
content = "ERROR: $rejectReason",
tokenEstimate = rejectReason.length / 4,
)
return@flatMap listOf(assistantEntry, resultEntry)
} }
emit(sessionId, OrchestrationResumedEvent(sessionId, stageId))
} }
val result = executor.execute(request) val result = executor.execute(request)
@@ -436,7 +496,8 @@ abstract class SessionOrchestrator(
sourceType = "assistantToolCall", sourceType = "assistantToolCall",
sourceId = sourceId, sourceId = sourceId,
content = Json.encodeToString(ToolCallRequest.serializer(), toolCall), content = Json.encodeToString(ToolCallRequest.serializer(), toolCall),
tokenEstimate = toolCall.function.arguments.length / 4, tokenEstimate = estimateTokens(toolCall.function.arguments),
role = EntryRole.ASSISTANT,
) )
val resultContent = when (result) { val resultContent = when (result) {
is ToolResult.Success -> result.output is ToolResult.Success -> result.output
@@ -451,13 +512,14 @@ abstract class SessionOrchestrator(
sourceType = "toolResult", sourceType = "toolResult",
sourceId = sourceId, sourceId = sourceId,
content = resultContent, content = resultContent,
tokenEstimate = resultContent.length / 4, tokenEstimate = estimateTokens(resultContent),
role = EntryRole.TOOL,
) )
listOf(assistantEntry, resultEntry) listOf(assistantEntry, resultEntry)
} }
} }
private fun buildSchemaEntries( private suspend fun buildSchemaEntries(
responseFormat: ResponseFormat, responseFormat: ResponseFormat,
stageId: StageId, stageId: StageId,
): List<ContextEntry> { ): List<ContextEntry> {
@@ -473,7 +535,8 @@ abstract class SessionOrchestrator(
content = instruction, content = instruction,
sourceType = "schemaInstruction", sourceType = "schemaInstruction",
sourceId = stageId.value, sourceId = stageId.value,
tokenEstimate = instruction.length / 4, tokenEstimate = estimateTokens(instruction),
role = EntryRole.SYSTEM,
), ),
) )
} }
@@ -742,6 +805,20 @@ abstract class SessionOrchestrator(
) )
} }
// --- token estimation ---
protected open suspend fun estimateTokens(content: String): Int {
val t = tokenizer
if (t != null) {
return runCatching { t.countTokens(content) }.getOrElse { fallbackTokenEstimate(content) }
}
return fallbackTokenEstimate(content)
}
private fun fallbackTokenEstimate(content: String): Int {
return (content.length / 4).coerceAtLeast(1)
}
// --- private functions --- // --- private functions ---
private suspend fun handleApproval( private suspend fun handleApproval(
@@ -4,23 +4,27 @@ import com.correx.core.context.model.CompressionMetadata
import com.correx.core.context.model.ContextEntry import com.correx.core.context.model.ContextEntry
import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextLayer
import com.correx.core.context.model.ContextPack import com.correx.core.context.model.ContextPack
import com.correx.core.context.model.EntryRole
import com.correx.core.context.model.TokenBudget import com.correx.core.context.model.TokenBudget
import com.correx.core.events.types.ContextEntryId import com.correx.core.events.types.ContextEntryId
import com.correx.core.events.types.ContextPackId import com.correx.core.events.types.ContextPackId
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.StageId
import com.correx.core.inference.Tokenizer
import com.correx.core.router.model.RouterConfig import com.correx.core.router.model.RouterConfig
import com.correx.core.router.model.RouterL2Entry import com.correx.core.router.model.RouterL2Entry
import com.correx.core.router.model.RouterState import com.correx.core.router.model.RouterState
import com.correx.core.router.model.RouterTurn import com.correx.core.router.model.RouterTurn
import com.correx.core.router.model.TurnRole
import java.util.* import java.util.*
interface RouterContextBuilder { interface RouterContextBuilder {
fun build(state: RouterState, budget: TokenBudget): ContextPack suspend fun build(state: RouterState, budget: TokenBudget): ContextPack
} }
class DefaultRouterContextBuilder( class DefaultRouterContextBuilder(
private val config: RouterConfig, private val config: RouterConfig,
private val tokenizer: Tokenizer? = null,
) : RouterContextBuilder { ) : RouterContextBuilder {
companion object { companion object {
@@ -28,40 +32,43 @@ class DefaultRouterContextBuilder(
"You are a routing assistant. Provide guidance based on workflow state and conversation context." "You are a routing assistant. Provide guidance based on workflow state and conversation context."
} }
override fun build(state: RouterState, budget: TokenBudget): ContextPack { override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack {
var remainingBudget = budget.limit var remainingBudget = budget.limit
val allEntries = mutableListOf<ContextEntry>() val allEntries = mutableListOf<ContextEntry>()
var droppedCount = 0 var droppedCount = 0
// === L0: System prompt (never dropped) ===
val systemPrompt = buildContextEntry( val systemPrompt = buildContextEntry(
sourceType = "systemPrompt", sourceType = "systemPrompt",
sourceId = "router-system", sourceId = "router-system",
content = SYSTEM_PROMPT, content = SYSTEM_PROMPT,
role = EntryRole.SYSTEM,
) )
remainingBudget -= systemPrompt.tokenEstimate remainingBudget -= systemPrompt.tokenEstimate
if (remainingBudget < 0) remainingBudget = 0 if (remainingBudget < 0) remainingBudget = 0
allEntries += systemPrompt allEntries += systemPrompt
// === L0: Workflow status (never dropped) ===
val workflowStatusEntry = buildContextEntry( val workflowStatusEntry = buildContextEntry(
sourceType = "workflowStatus", sourceType = "workflowStatus",
sourceId = state.currentStageId?.value ?: "none", sourceId = state.currentStageId?.value ?: "none",
content = buildWorkflowStatusContent(state), content = buildWorkflowStatusContent(state),
role = EntryRole.SYSTEM,
) )
remainingBudget -= workflowStatusEntry.tokenEstimate remainingBudget -= workflowStatusEntry.tokenEstimate
if (remainingBudget < 0) remainingBudget = 0 if (remainingBudget < 0) remainingBudget = 0
allEntries += workflowStatusEntry allEntries += workflowStatusEntry
// === L1: Conversation history (last N turns, capped at keepLast) ===
val recentTurns = state.conversationHistory.takeLast(config.conversationKeepLast) val recentTurns = state.conversationHistory.takeLast(config.conversationKeepLast)
for (turn in recentTurns) { for (turn in recentTurns) {
val content = buildConversationContent(turn) val role = when (turn.role) {
TurnRole.USER -> EntryRole.USER
TurnRole.ROUTER -> EntryRole.ASSISTANT
}
val entry = buildContextEntry( val entry = buildContextEntry(
sourceType = "conversation", sourceType = "conversation",
sourceId = "${turn.role.name}-${turn.hashCode()}", sourceId = "${turn.role.name}-${turn.hashCode()}",
content = content, content = turn.content,
layer = ContextLayer.L1, layer = ContextLayer.L1,
role = role,
) )
if (remainingBudget >= entry.tokenEstimate) { if (remainingBudget >= entry.tokenEstimate) {
remainingBudget -= entry.tokenEstimate remainingBudget -= entry.tokenEstimate
@@ -72,7 +79,6 @@ class DefaultRouterContextBuilder(
} }
} }
// === L2: Stage summaries from L2 memory (oldest-first eviction) ===
for (l2Entry in state.l2Memory) { for (l2Entry in state.l2Memory) {
val content = buildL2Content(l2Entry) val content = buildL2Content(l2Entry)
val entry = buildContextEntry( val entry = buildContextEntry(
@@ -80,6 +86,7 @@ class DefaultRouterContextBuilder(
sourceId = l2Entry.stageId.value, sourceId = l2Entry.stageId.value,
content = content, content = content,
layer = ContextLayer.L2, layer = ContextLayer.L2,
role = EntryRole.SYSTEM,
) )
if (remainingBudget >= entry.tokenEstimate) { if (remainingBudget >= entry.tokenEstimate) {
remainingBudget -= entry.tokenEstimate remainingBudget -= entry.tokenEstimate
@@ -90,7 +97,6 @@ class DefaultRouterContextBuilder(
} }
} }
// Assemble layers and pack
val layers = allEntries.groupBy { it.layer } val layers = allEntries.groupBy { it.layer }
val budgetUsed = allEntries.sumOf { it.tokenEstimate } val budgetUsed = allEntries.sumOf { it.tokenEstimate }
@@ -119,19 +125,16 @@ class DefaultRouterContextBuilder(
return s.lowercase().let { it[0].uppercaseChar() + it.drop(1) } return s.lowercase().let { it[0].uppercaseChar() + it.drop(1) }
} }
private fun buildConversationContent(turn: RouterTurn): String {
return "[${turn.role.name}] ${turn.content}"
}
private fun buildL2Content(l2Entry: RouterL2Entry): String { private fun buildL2Content(l2Entry: RouterL2Entry): String {
return "Stage ${l2Entry.stageId.value} (${l2Entry.outcome}): ${l2Entry.summary}" return "Stage ${l2Entry.stageId.value} (${l2Entry.outcome}): ${l2Entry.summary}"
} }
private fun buildContextEntry( private suspend fun buildContextEntry(
sourceType: String, sourceType: String,
sourceId: String, sourceId: String,
content: String, content: String,
layer: ContextLayer = ContextLayer.L0, layer: ContextLayer = ContextLayer.L0,
role: EntryRole = EntryRole.USER,
): ContextEntry { ): ContextEntry {
return ContextEntry( return ContextEntry(
id = ContextEntryId(UUID.randomUUID().toString()), id = ContextEntryId(UUID.randomUUID().toString()),
@@ -140,10 +143,19 @@ class DefaultRouterContextBuilder(
sourceType = sourceType, sourceType = sourceType,
sourceId = sourceId, sourceId = sourceId,
tokenEstimate = estimateTokens(content), tokenEstimate = estimateTokens(content),
role = role,
) )
} }
private fun estimateTokens(content: String): Int { private suspend fun estimateTokens(content: String): Int {
return (content.length / 2).coerceAtLeast(1) val t = tokenizer
if (t != null) {
return runCatching { t.countTokens(content) }.getOrElse { fallbackEstimate(content) }
}
return fallbackEstimate(content)
} }
}
private fun fallbackEstimate(content: String): Int {
return (content.length / 4).coerceAtLeast(1)
}
}
@@ -2,6 +2,7 @@ package com.correx.infrastructure.inference.llama.cpp
import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextLayer
import com.correx.core.context.model.ContextPack import com.correx.core.context.model.ContextPack
import com.correx.core.context.model.EntryRole
import com.correx.core.events.types.ProviderId import com.correx.core.events.types.ProviderId
import com.correx.core.inference.CapabilityScore import com.correx.core.inference.CapabilityScore
import com.correx.core.inference.FinishReason import com.correx.core.inference.FinishReason
@@ -149,17 +150,17 @@ class LlamaCppInferenceProvider(
.flatMap { it.value } .flatMap { it.value }
.joinToString("\n\n") { it.content } .joinToString("\n\n") { it.content }
.takeIf { it.isNotBlank() } .takeIf { it.isNotBlank() }
val conversationMessages = sorted.filter { it.key != ContextLayer.L0 }.flatMap { (layer, entries) -> val conversationMessages = sorted.filter { it.key != ContextLayer.L0 }.flatMap { (_, entries) ->
when (layer) { entries.map { entry ->
ContextLayer.L1 -> entries.map { ChatMessage("user", it.content) } ChatMessage(
ContextLayer.L2 -> entries.map { role = when (entry.role) {
ChatMessage( EntryRole.SYSTEM -> "system"
role = if (it.sourceType == "assistant") "assistant" else "user", EntryRole.ASSISTANT -> "assistant"
content = it.content, EntryRole.TOOL -> "tool"
) EntryRole.USER -> "user"
} },
content = entry.content,
else -> emptyList() )
} }
} }
val messages = buildList { val messages = buildList {
@@ -168,4 +169,4 @@ class LlamaCppInferenceProvider(
} }
return messages.ifEmpty { listOf(ChatMessage("user", "")) } return messages.ifEmpty { listOf(ChatMessage("user", "")) }
} }
} }
@@ -10,6 +10,7 @@ import com.correx.core.events.EventDispatcher
import com.correx.core.events.stores.EventStore import com.correx.core.events.stores.EventStore
import com.correx.core.inference.InferenceProvider import com.correx.core.inference.InferenceProvider
import com.correx.core.inference.InferenceRouter import com.correx.core.inference.InferenceRouter
import com.correx.core.inference.Tokenizer
import com.correx.core.router.DefaultRouterContextBuilder import com.correx.core.router.DefaultRouterContextBuilder
import com.correx.core.router.DefaultRouterFacade import com.correx.core.router.DefaultRouterFacade
import com.correx.core.router.DefaultRouterReducer import com.correx.core.router.DefaultRouterReducer
@@ -120,12 +121,13 @@ object InfrastructureModule {
eventStore: EventStore, eventStore: EventStore,
inferenceRouter: InferenceRouter, inferenceRouter: InferenceRouter,
config: RouterConfig = RouterConfig(), config: RouterConfig = RouterConfig(),
tokenizer: Tokenizer? = null,
): RouterFacade { ): RouterFacade {
val reducer = DefaultRouterReducer() val reducer = DefaultRouterReducer()
val projector = RouterProjector(reducer) val projector = RouterProjector(reducer)
val replayer = DefaultEventReplayer(eventStore, projector) val replayer = DefaultEventReplayer(eventStore, projector)
val repository = DefaultRouterRepository(replayer) val repository = DefaultRouterRepository(replayer)
val contextBuilder = DefaultRouterContextBuilder(config) val contextBuilder = DefaultRouterContextBuilder(config, tokenizer)
return DefaultRouterFacade(repository, contextBuilder, inferenceRouter, eventStore, config) return DefaultRouterFacade(repository, contextBuilder, inferenceRouter, eventStore, config)
} }
} }
@@ -1,5 +1,8 @@
import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextLayer
import com.correx.core.context.model.ContextPack
import com.correx.core.context.model.EntryRole
import com.correx.core.context.model.TokenBudget import com.correx.core.context.model.TokenBudget
import kotlinx.coroutines.runBlocking
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.StageId
import com.correx.core.router.DefaultRouterContextBuilder import com.correx.core.router.DefaultRouterContextBuilder
@@ -23,6 +26,10 @@ class RouterContextBuilderTest {
private val config = RouterConfig(conversationKeepLast = 3, tokenBudget = TokenBudget(limit = 200)) private val config = RouterConfig(conversationKeepLast = 3, tokenBudget = TokenBudget(limit = 200))
private val builder = DefaultRouterContextBuilder(config) private val builder = DefaultRouterContextBuilder(config)
private fun buildPack(state: RouterState, budget: TokenBudget): ContextPack = runBlocking {
buildPack(state, budget)
}
private val sessionId = SessionId("test-session") private val sessionId = SessionId("test-session")
private val stageId = StageId("stage-1") private val stageId = StageId("stage-1")
private val clock = Clock.System private val clock = Clock.System
@@ -41,7 +48,7 @@ class RouterContextBuilderTest {
RouterTurn(TurnRole.USER, "hi", clock.now()), RouterTurn(TurnRole.USER, "hi", clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
assertEquals(10000, pack.budgetLimit) assertEquals(10000, pack.budgetLimit)
assertTrue(pack.budgetUsed <= pack.budgetLimit) assertTrue(pack.budgetUsed <= pack.budgetLimit)
} }
@@ -63,7 +70,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 100)) val pack = buildPack(state, TokenBudget(limit = 100))
// L0 entries (system prompt + workflow status) always fit; L1/L2 should be dropped // L0 entries (system prompt + workflow status) always fit; L1/L2 should be dropped
assertTrue(pack.compressionMetadata.entriesDropped > 0) assertTrue(pack.compressionMetadata.entriesDropped > 0)
} }
@@ -83,7 +90,7 @@ class RouterContextBuilderTest {
) )
// L0 consumes ~60 tokens; budget 93 leaves ~33 for L2. // L0 consumes ~60 tokens; budget 93 leaves ~33 for L2.
// Each L2 entry is ~12 tokens; 2 fit (s1, s2), s3 is dropped // Each L2 entry is ~12 tokens; 2 fit (s1, s2), s3 is dropped
val pack = builder.build(state, TokenBudget(limit = 93)) val pack = buildPack(state, TokenBudget(limit = 93))
val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList()
// Oldest-first eviction: s1 and s2 fit, s3 is the one dropped // Oldest-first eviction: s1 and s2 fit, s3 is the one dropped
val remainingStageIds = l2Entries.map { it.sourceId }.toSet() val remainingStageIds = l2Entries.map { it.sourceId }.toSet()
@@ -127,7 +134,7 @@ class RouterContextBuilderTest {
workflowStatus = WorkflowStatus.IDLE, workflowStatus = WorkflowStatus.IDLE,
currentStageId = null, currentStageId = null,
) )
val pack = builder.build(state, TokenBudget(limit = 0)) val pack = buildPack(state, TokenBudget(limit = 0))
val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList() val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList()
assertTrue(l0Entries.any { it.sourceType == "systemPrompt" }) assertTrue(l0Entries.any { it.sourceType == "systemPrompt" })
} }
@@ -139,7 +146,7 @@ class RouterContextBuilderTest {
workflowStatus = WorkflowStatus.FAILED, workflowStatus = WorkflowStatus.FAILED,
currentStageId = StageId("failed-stage"), currentStageId = StageId("failed-stage"),
) )
val pack = builder.build(state, TokenBudget(limit = 0)) val pack = buildPack(state, TokenBudget(limit = 0))
val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList() val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList()
assertTrue(l0Entries.any { it.sourceType == "workflowStatus" }) assertTrue(l0Entries.any { it.sourceType == "workflowStatus" })
val workflowEntry = l0Entries.find { it.sourceType == "workflowStatus" } val workflowEntry = l0Entries.find { it.sourceType == "workflowStatus" }
@@ -156,7 +163,7 @@ class RouterContextBuilderTest {
conversationHistory = emptyList(), conversationHistory = emptyList(),
l2Memory = emptyList(), l2Memory = emptyList(),
) )
val pack = builder.build(state, TokenBudget(limit = 0)) val pack = buildPack(state, TokenBudget(limit = 0))
val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList() val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList()
assertEquals(2, l0Entries.size) // systemPrompt + workflowStatus assertEquals(2, l0Entries.size) // systemPrompt + workflowStatus
} }
@@ -174,7 +181,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s1"), "x".repeat(500), StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s1"), "x".repeat(500), StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 50)) val pack = buildPack(state, TokenBudget(limit = 50))
val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList() val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList()
val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList() val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList()
val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList()
@@ -225,7 +232,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s2"), "summary", StageOutcomeKind.FAILURE, clock.now()), RouterL2Entry(StageId("s2"), "summary", StageOutcomeKind.FAILURE, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList()
assertEquals(2, l2Entries.size) assertEquals(2, l2Entries.size)
assertEquals(0, pack.compressionMetadata.entriesDropped) assertEquals(0, pack.compressionMetadata.entriesDropped)
@@ -241,7 +248,7 @@ class RouterContextBuilderTest {
sessionId = sessionId, sessionId = sessionId,
workflowStatus = WorkflowStatus.IDLE, workflowStatus = WorkflowStatus.IDLE,
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l0Entries = pack.layers[ContextLayer.L0] val l0Entries = pack.layers[ContextLayer.L0]
assertNotNull(l0Entries) assertNotNull(l0Entries)
assertTrue(l0Entries!!.any { it.sourceType == "systemPrompt" && it.layer == ContextLayer.L0 }) assertTrue(l0Entries!!.any { it.sourceType == "systemPrompt" && it.layer == ContextLayer.L0 })
@@ -254,7 +261,7 @@ class RouterContextBuilderTest {
workflowStatus = WorkflowStatus.RUNNING, workflowStatus = WorkflowStatus.RUNNING,
currentStageId = stageId, currentStageId = stageId,
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l0Entries = pack.layers[ContextLayer.L0] val l0Entries = pack.layers[ContextLayer.L0]
assertTrue(l0Entries!!.any { it.sourceType == "workflowStatus" && it.layer == ContextLayer.L0 }) assertTrue(l0Entries!!.any { it.sourceType == "workflowStatus" && it.layer == ContextLayer.L0 })
} }
@@ -270,7 +277,7 @@ class RouterContextBuilderTest {
RouterTurn(TurnRole.ROUTER, "hi", clock.now()), RouterTurn(TurnRole.ROUTER, "hi", clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l1Entries = pack.layers[ContextLayer.L1] val l1Entries = pack.layers[ContextLayer.L1]
assertNotNull(l1Entries) assertNotNull(l1Entries)
assertEquals(2, l1Entries!!.size) assertEquals(2, l1Entries!!.size)
@@ -288,7 +295,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l2Entries = pack.layers[ContextLayer.L2] val l2Entries = pack.layers[ContextLayer.L2]
assertNotNull(l2Entries) assertNotNull(l2Entries)
assertEquals(1, l2Entries!!.size) assertEquals(1, l2Entries!!.size)
@@ -309,7 +316,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val presentLayers = pack.layers.keys val presentLayers = pack.layers.keys
assertTrue(presentLayers.containsAll(listOf(ContextLayer.L0, ContextLayer.L1, ContextLayer.L2))) assertTrue(presentLayers.containsAll(listOf(ContextLayer.L0, ContextLayer.L1, ContextLayer.L2)))
assertFalse(presentLayers.contains(ContextLayer.L3)) assertFalse(presentLayers.contains(ContextLayer.L3))
@@ -336,7 +343,7 @@ class RouterContextBuilderTest {
RouterTurn(TurnRole.USER, "five", Instant.parse("2026-01-05T00:00:00Z")), RouterTurn(TurnRole.USER, "five", Instant.parse("2026-01-05T00:00:00Z")),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList() val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList()
// conversationKeepLast=2, so only the last 2 turns should appear // conversationKeepLast=2, so only the last 2 turns should appear
assertEquals(2, l1Entries.size) assertEquals(2, l1Entries.size)
@@ -357,7 +364,7 @@ class RouterContextBuilderTest {
RouterTurn(TurnRole.ROUTER, "fourth", Instant.parse("2026-01-04T00:00:00Z")), RouterTurn(TurnRole.ROUTER, "fourth", Instant.parse("2026-01-04T00:00:00Z")),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList() val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList()
assertEquals(3, l1Entries.size) assertEquals(3, l1Entries.size)
// The last 3: second, third, fourth // The last 3: second, third, fourth
@@ -373,7 +380,7 @@ class RouterContextBuilderTest {
@Test @Test
fun `build with empty state produces L0 only`() { fun `build with empty state produces L0 only`() {
val state = RouterState() val state = RouterState()
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l0Entries = pack.layers[ContextLayer.L0] val l0Entries = pack.layers[ContextLayer.L0]
assertNotNull(l0Entries) assertNotNull(l0Entries)
assertEquals(2, l0Entries!!.size) assertEquals(2, l0Entries!!.size)
@@ -384,7 +391,7 @@ class RouterContextBuilderTest {
@Test @Test
fun `pack contains correct metadata on empty state`() { fun `pack contains correct metadata on empty state`() {
val state = RouterState() val state = RouterState()
val pack = builder.build(state, TokenBudget(limit = 5000)) val pack = buildPack(state, TokenBudget(limit = 5000))
assertEquals(5000, pack.budgetLimit) assertEquals(5000, pack.budgetLimit)
assertEquals(0, pack.compressionMetadata.entriesDropped) assertEquals(0, pack.compressionMetadata.entriesDropped)
assertEquals(listOf("L0Immutable", "Conversation"), pack.compressionMetadata.appliedStrategies) assertEquals(listOf("L0Immutable", "Conversation"), pack.compressionMetadata.appliedStrategies)
@@ -398,7 +405,7 @@ class RouterContextBuilderTest {
workflowStatus = WorkflowStatus.RUNNING, workflowStatus = WorkflowStatus.RUNNING,
currentStageId = stageId, currentStageId = stageId,
) )
val pack = builder.build(state, TokenBudget(limit = 1000)) val pack = buildPack(state, TokenBudget(limit = 1000))
assertEquals("test-session-router-pack", pack.id.value) assertEquals("test-session-router-pack", pack.id.value)
assertEquals(sessionId, pack.sessionId) assertEquals(sessionId, pack.sessionId)
assertEquals(stageId, pack.stageId) assertEquals(stageId, pack.stageId)
@@ -411,7 +418,7 @@ class RouterContextBuilderTest {
@Test @Test
fun `system prompt has expected content`() { fun `system prompt has expected content`() {
val state = RouterState() val state = RouterState()
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val systemEntry = pack.layers[ContextLayer.L0]?.find { it.sourceType == "systemPrompt" } val systemEntry = pack.layers[ContextLayer.L0]?.find { it.sourceType == "systemPrompt" }
assertNotNull(systemEntry) assertNotNull(systemEntry)
assertEquals( assertEquals(
@@ -427,7 +434,7 @@ class RouterContextBuilderTest {
workflowStatus = WorkflowStatus.COMPLETED, workflowStatus = WorkflowStatus.COMPLETED,
currentStageId = null, currentStageId = null,
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val workflowEntry = pack.layers[ContextLayer.L0]?.find { it.sourceType == "workflowStatus" } val workflowEntry = pack.layers[ContextLayer.L0]?.find { it.sourceType == "workflowStatus" }
assertNotNull(workflowEntry) assertNotNull(workflowEntry)
assertTrue(workflowEntry!!.content.contains("Status: Completed")) assertTrue(workflowEntry!!.content.contains("Status: Completed"))
@@ -444,10 +451,11 @@ class RouterContextBuilderTest {
RouterTurn(TurnRole.USER, "user message", clock.now()), RouterTurn(TurnRole.USER, "user message", clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList() val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList()
assertEquals(1, l1Entries.size) assertEquals(1, l1Entries.size)
assertEquals("[USER] user message", l1Entries[0].content) assertEquals("user message", l1Entries[0].content)
assertEquals(EntryRole.USER, l1Entries[0].role)
} }
@Test @Test
@@ -460,7 +468,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("stage-x"), "completed with 3 items", StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("stage-x"), "completed with 3 items", StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList()
assertEquals(1, l2Entries.size) assertEquals(1, l2Entries.size)
val entry = l2Entries[0] val entry = l2Entries[0]
@@ -487,7 +495,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
val allEntries = pack.layers.values.flatten() val allEntries = pack.layers.values.flatten()
val computedSum = allEntries.sumOf { it.tokenEstimate } val computedSum = allEntries.sumOf { it.tokenEstimate }
assertEquals(computedSum, pack.budgetUsed) assertEquals(computedSum, pack.budgetUsed)
@@ -509,7 +517,7 @@ class RouterContextBuilderTest {
RouterL2Entry(StageId("s2"), "x".repeat(500), StageOutcomeKind.SUCCESS, clock.now()), RouterL2Entry(StageId("s2"), "x".repeat(500), StageOutcomeKind.SUCCESS, clock.now()),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 65)) val pack = buildPack(state, TokenBudget(limit = 65))
// L0 consumes ~60 tokens, leaving 5 — both L2 entries (each ~259 tokens) dropped // L0 consumes ~60 tokens, leaving 5 — both L2 entries (each ~259 tokens) dropped
// conversationKeepLast=0 means conversation entry is not included // conversationKeepLast=0 means conversation entry is not included
assertEquals(2, pack.compressionMetadata.entriesDropped) assertEquals(2, pack.compressionMetadata.entriesDropped)
@@ -548,7 +556,7 @@ class RouterContextBuilderTest {
), ),
), ),
) )
val pack = builder.build(state, TokenBudget(limit = 10000)) val pack = buildPack(state, TokenBudget(limit = 10000))
// L0: system prompt + workflow status // L0: system prompt + workflow status
val l0 = pack.layers[ContextLayer.L0]!! val l0 = pack.layers[ContextLayer.L0]!!
@@ -577,7 +585,7 @@ class RouterContextBuilderTest {
val state = RouterState( val state = RouterState(
workflowStatus = WorkflowStatus.IDLE, workflowStatus = WorkflowStatus.IDLE,
) )
val pack = builder.build(state, TokenBudget(limit = 1000)) val pack = buildPack(state, TokenBudget(limit = 1000))
assertTrue(pack.id.value.contains("unknown")) assertTrue(pack.id.value.contains("unknown"))
} }
} }
@@ -109,7 +109,7 @@ class RouterFacadeTest {
) )
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter("response"), inferenceRouter = mockInferenceRouter("response"),
eventStore = mockStore, eventStore = mockStore,
@@ -129,7 +129,7 @@ class RouterFacadeTest {
RouterState(sessionId = sessionId, workflowStatus = WorkflowStatus.IDLE, currentStageId = null) RouterState(sessionId = sessionId, workflowStatus = WorkflowStatus.IDLE, currentStageId = null)
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter("response"), inferenceRouter = mockInferenceRouter("response"),
eventStore = mockStore, eventStore = mockStore,
@@ -157,7 +157,7 @@ class RouterFacadeTest {
) )
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack { override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack {
capturedStates.add(state) capturedStates.add(state)
return emptyContextPack() return emptyContextPack()
} }
@@ -190,7 +190,7 @@ class RouterFacadeTest {
RouterState(sessionId = sessionId, workflowStatus = WorkflowStatus.RUNNING) RouterState(sessionId = sessionId, workflowStatus = WorkflowStatus.RUNNING)
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack { override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack {
capturedStates.add(state) capturedStates.add(state)
return emptyContextPack() return emptyContextPack()
} }
@@ -218,7 +218,7 @@ class RouterFacadeTest {
fun `state is passed through to context builder`(): Unit = runBlocking { fun `state is passed through to context builder`(): Unit = runBlocking {
val capturedState = mutableListOf<RouterState>() val capturedState = mutableListOf<RouterState>()
val mockContextBuilder = object : RouterContextBuilder { val mockContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack { override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack {
capturedState.add(state) capturedState.add(state)
return emptyContextPack() return emptyContextPack()
} }
@@ -288,7 +288,7 @@ class RouterFacadeTest {
) )
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter, inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(), eventStore = mockEventStore(),
@@ -319,7 +319,7 @@ class RouterFacadeTest {
) )
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter, inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(), eventStore = mockEventStore(),
@@ -345,7 +345,7 @@ class RouterFacadeTest {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState() override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter, inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(), eventStore = mockEventStore(),
@@ -373,7 +373,7 @@ class RouterFacadeTest {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState() override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter, inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(), eventStore = mockEventStore(),
@@ -394,7 +394,7 @@ class RouterFacadeTest {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState() override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack { override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack {
val pack = ContextPack( val pack = ContextPack(
id = ContextPackId("test-pack"), id = ContextPackId("test-pack"),
sessionId = state.sessionId ?: SessionId("unknown"), sessionId = state.sessionId ?: SessionId("unknown"),
@@ -449,7 +449,7 @@ class RouterFacadeTest {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState() override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = object : InferenceRouter { inferenceRouter = object : InferenceRouter {
override suspend fun route( override suspend fun route(
@@ -497,7 +497,7 @@ class RouterFacadeTest {
) )
}, },
routerContextBuilder = object : RouterContextBuilder { routerContextBuilder = object : RouterContextBuilder {
override fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack() override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
}, },
inferenceRouter = mockInferenceRouter("inference response"), inferenceRouter = mockInferenceRouter("inference response"),
eventStore = eventStore, eventStore = eventStore,