diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chain/AbstractChain.java b/spring-ai-core/src/main/java/org/springframework/ai/chain/AbstractChain.java index f714685f9..5300062f6 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chain/AbstractChain.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chain/AbstractChain.java @@ -26,27 +26,27 @@ public abstract class AbstractChain implements Chain { // TODO validation of input/outputs @Override - public Map apply(Map inputMap) { - Map inputMapToUse = preProcess(inputMap); - Map outputMap = doApply(inputMapToUse); - Map outputMapToUse = postProcess(inputMapToUse, outputMap); - return outputMapToUse; + public AiOutput apply(AiInput aiInput) { + AiInput aiInputToUse = preProcess(aiInput); + AiOutput aiOutput = doApply(aiInputToUse); + Map outputMapToUse = postProcess(aiInput, aiOutput); + return new AiOutput(outputMapToUse); } @Override - public Map preProcess(Map inputMap) { - validateInputs(inputMap); - return inputMap; + public AiInput preProcess(AiInput aiInput) { + validateInputs(aiInput.getInputData()); + return aiInput; } - protected abstract Map doApply(Map inputMap); + protected abstract AiOutput doApply(AiInput aiInput); @Override - public Map postProcess(Map inputMap, Map outputMap) { - validateOutputs(outputMap); + public Map postProcess(AiInput aiInput, AiOutput aiOutput) { + validateOutputs(aiOutput.getOutputData()); Map combindedMap = new HashMap<>(); - combindedMap.putAll(inputMap); - combindedMap.putAll(outputMap); + combindedMap.putAll(aiInput.getInputData()); + combindedMap.putAll(aiOutput.getOutputData()); return combindedMap; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chain/AiInput.java b/spring-ai-core/src/main/java/org/springframework/ai/chain/AiInput.java new file mode 100644 index 000000000..0e80cad58 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chain/AiInput.java @@ -0,0 +1,18 @@ +package org.springframework.ai.chain; + +import java.util.List; +import java.util.Map; + +public class AiInput { + + private Map inputData; + + public AiInput(Map inputData) { + this.inputData = inputData; + } + + Map getInputData() { + return inputData; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chain/AiOutput.java b/spring-ai-core/src/main/java/org/springframework/ai/chain/AiOutput.java new file mode 100644 index 000000000..7b62ac406 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chain/AiOutput.java @@ -0,0 +1,18 @@ +package org.springframework.ai.chain; + +import java.util.List; +import java.util.Map; + +public class AiOutput { + + private final Map outputData; + + public AiOutput(Map outputData) { + this.outputData = outputData; + } + + Map getOutputData() { + return this.outputData; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chain/Chain.java b/spring-ai-core/src/main/java/org/springframework/ai/chain/Chain.java index 44328b733..68cf34d97 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chain/Chain.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chain/Chain.java @@ -4,14 +4,14 @@ import java.util.List; import java.util.Map; import java.util.function.Function; -public interface Chain extends Function, Map> { +public interface Chain extends Function { List getInputKeys(); List getOutputKeys(); - Map preProcess(Map inputMap); + AiInput preProcess(AiInput aiInput); - Map postProcess(Map inputMap, Map outputMap); + Map postProcess(AiInput aiInput, AiOutput aiOutput); } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chain/ChainTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chain/ChainTests.java index 86c406534..8e4d70b76 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chain/ChainTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chain/ChainTests.java @@ -14,15 +14,17 @@ class ChainTests { @Test void badInputs() { Chain chain = new FakeChain(); - assertThatThrownBy(() -> chain.apply(Map.of("foobar", "baz"))).isInstanceOf(IllegalArgumentException.class) + AiInput aiInput = new AiInput(Map.of("foobar", "baz")); + assertThatThrownBy(() -> chain.apply(aiInput)).isInstanceOf(IllegalArgumentException.class) .hasMessageContaining("Missing some input keys"); } @Test void correctInputs() { Chain chain = new FakeChain(); - Map output = chain.apply(Map.of("foo", "bar")); - assertThat(output).containsEntry("foo", "bar").containsEntry("bar", "baz"); + AiInput aiInput = new AiInput(Map.of("foo", "bar")); + AiOutput aiOutput = chain.apply(aiInput); + assertThat(aiOutput.getOutputData()).containsEntry("foo", "bar").containsEntry("bar", "baz"); } class FakeChain extends AbstractChain { @@ -53,12 +55,12 @@ class ChainTests { } @Override - protected Map doApply(Map inputMap) { + protected AiOutput doApply(AiInput aiInput) { if (beCorrect) { - return Map.of("bar", "baz"); + return new AiOutput(Map.of("bar", "baz")); } else { - return Map.of("baz", "bar"); + return new AiOutput(Map.of("baz", "bar")); } }