diff --git a/core/transitions/src/main/kotlin/com/correx/core/transitions/graph/StageConfig.kt b/core/transitions/src/main/kotlin/com/correx/core/transitions/graph/StageConfig.kt index 7f6827bb..5082c3b2 100644 --- a/core/transitions/src/main/kotlin/com/correx/core/transitions/graph/StageConfig.kt +++ b/core/transitions/src/main/kotlin/com/correx/core/transitions/graph/StageConfig.kt @@ -1,11 +1,10 @@ package com.correx.core.transitions.graph +import com.correx.core.artifacts.kind.TypedArtifactSlot import com.correx.core.events.types.ArtifactId import com.correx.core.inference.GenerationConfig import com.correx.core.inference.ModelCapability -import kotlinx.serialization.Serializable -@Serializable data class StageConfig( val requiredCapabilities: Set = emptySet(), val tokenBudget: Int = 4096, @@ -16,7 +15,7 @@ data class StageConfig( ), val allowedTools: Set = emptySet(), val maxRetries: Int = 3, - val produces: Set = emptySet(), + val produces: List = emptyList(), val needs: Set = emptySet(), val metadata: Map = emptyMap(), ) diff --git a/infrastructure/workflow/build.gradle b/infrastructure/workflow/build.gradle index 8fda310e..7b27ea38 100644 --- a/infrastructure/workflow/build.gradle +++ b/infrastructure/workflow/build.gradle @@ -7,6 +7,7 @@ dependencies { implementation(project(":core:transitions")) implementation(project(":core:inference")) implementation(project(":core:events")) + implementation(project(":core:artifacts")) implementation("com.fasterxml.jackson.dataformat:jackson-dataformat-toml:2.17.0") implementation("com.fasterxml.jackson.module:jackson-module-kotlin:2.17.0") testImplementation "org.junit.jupiter:junit-jupiter" diff --git a/infrastructure/workflow/src/main/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoader.kt b/infrastructure/workflow/src/main/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoader.kt index 08fc6f79..6ec6ef9b 100644 --- a/infrastructure/workflow/src/main/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoader.kt +++ b/infrastructure/workflow/src/main/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoader.kt @@ -1,5 +1,8 @@ package com.correx.infrastructure.workflow +import com.correx.core.artifacts.kind.ArtifactKindRegistry +import com.correx.core.artifacts.kind.DefaultArtifactKindRegistry +import com.correx.core.artifacts.kind.TypedArtifactSlot import com.correx.core.events.types.ArtifactId import com.correx.core.events.types.StageId import com.correx.core.events.types.TransitionId @@ -21,11 +24,16 @@ private data class WorkflowFile( val transitions: List = emptyList(), ) +private data class ProducesEntry( + val name: String = "", + val kind: String = "", +) + private data class StageSection( val id: String = "", val prompt: String? = null, - val systemPrompt: String? = null, - val produces: List = emptyList(), + @JsonProperty("system_prompt") val systemPrompt: String? = null, + val produces: List = emptyList(), val needs: List = emptyList(), @JsonProperty("allowed_tools") val allowedTools: List = emptyList(), @JsonProperty("token_budget") val tokenBudget: Int = 4096, @@ -50,7 +58,9 @@ private val mapper = TomlMapper.builder() .disable(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES) .build() -class TomlWorkflowLoader : WorkflowLoader { +class TomlWorkflowLoader( + private val registry: ArtifactKindRegistry = DefaultArtifactKindRegistry(), +) : WorkflowLoader { override fun load(path: Path): WorkflowGraph { val raw = path.readText() val file = runCatching { mapper.readValue(raw) } @@ -62,14 +72,18 @@ class TomlWorkflowLoader : WorkflowLoader { val startId = StageId(start) val stageMap = stages.associate { s -> StageId(s.id) to StageConfig( - produces = s.produces.map { ArtifactId(it) }.toSet(), + produces = s.produces.map { entry -> + val resolvedKind = registry.get(entry.kind) + ?: error("Unknown artifact kind '${entry.kind}' in stage '${s.id}'") + TypedArtifactSlot(name = ArtifactId(entry.name), kind = resolvedKind) + }, needs = s.needs.map { ArtifactId(it) }.toSet(), allowedTools = s.allowedTools.toSet(), tokenBudget = s.tokenBudget, maxRetries = s.maxRetries, metadata = buildMap { s.prompt?.let { put("prompt", it) } - ?: s.systemPrompt?.let { put("systemPrompt", it) } + s.systemPrompt?.let { put("systemPrompt", it) } }, ) } @@ -77,7 +91,7 @@ class TomlWorkflowLoader : WorkflowLoader { val declaredIds = stageMap.keys.map { it.value }.toSet() validate(start, declaredIds, transitions) - val allProduced = stageMap.values.flatMap { it.produces }.toSet() + val allProduced = stageMap.values.flatMap { it.produces }.map { it.name }.toSet() stageMap.forEach { (stageId, config) -> config.needs.forEach { needed -> if (needed !in allProduced) { diff --git a/infrastructure/workflow/src/test/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoaderTest.kt b/infrastructure/workflow/src/test/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoaderTest.kt index 5c26bab5..ea64a16b 100644 --- a/infrastructure/workflow/src/test/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoaderTest.kt +++ b/infrastructure/workflow/src/test/kotlin/com/correx/infrastructure/workflow/TomlWorkflowLoaderTest.kt @@ -17,7 +17,7 @@ class TomlWorkflowLoaderTest { [[stages]] id = "collect" prompt = "prompts/collect.md" - produces = ["system_stats"] + produces = [{ name = "system_stats", kind = "file_written" }] allowed_tools = ["ShellTool"] token_budget = 2048 max_retries = 2 @@ -26,7 +26,7 @@ class TomlWorkflowLoaderTest { id = "report" prompt = "prompts/report.md" needs = ["system_stats"] - produces = ["report"] + produces = [{ name = "report", kind = "file_written" }] token_budget = 2048 [[transitions]] @@ -53,10 +53,10 @@ class TomlWorkflowLoaderTest { assertEquals(2, graph.transitions.size) val collectStage = graph.stages[graph.start]!! - assertEquals(setOf("system_stats"), collectStage.produces.map { it.value }.toSet()) + assertEquals(setOf("system_stats"), collectStage.produces.map { it.name.value }.toSet()) assertEquals("prompts/collect.md", collectStage.metadata["prompt"]) - val reportStage = graph.stages.values.first { it.produces.any { a -> a.value == "report" } } + val reportStage = graph.stages.values.first { it.produces.any { a -> a.name.value == "report" } } assertEquals(setOf("system_stats"), reportStage.needs.map { it.value }.toSet()) } @@ -90,4 +90,11 @@ class TomlWorkflowLoaderTest { val path = Files.createTempFile("workflow", ".toml").also { it.writeText(bad) } assertThrows { loader.load(path) } } + + @Test + fun `unknown artifact kind throws error`() { + val bad = validToml.replace("kind = \"file_written\"", "kind = \"unknown_kind\"") + val path = Files.createTempFile("workflow", ".toml").also { it.writeText(bad) } + assertThrows { loader.load(path) } + } }