diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractPromptTemplate.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractPromptTemplate.java index c18e7dda5..754ee3776 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractPromptTemplate.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractPromptTemplate.java @@ -23,8 +23,12 @@ public abstract class AbstractPromptTemplate implements PromptInput { private Optional outputParser = Optional.empty(); - public AbstractPromptTemplate(Optional outputParser) { - this.outputParser = outputParser; + public AbstractPromptTemplate() { + this.outputParser = Optional.empty(); + } + + public AbstractPromptTemplate(OutputParser outputParser) { + this.outputParser = Optional.of(outputParser); } @Override @@ -32,8 +36,9 @@ public abstract class AbstractPromptTemplate implements PromptInput { return this.outputParser; } - public abstract String formatAsString(Map inputVariables); - - public abstract PromptValue formatAsPrompt(Map inputVariables); + public PromptValue formatAsPrompt(Map inputVariables) { + String formattedPrompt = formatAsString(inputVariables); + return new StringPromptValue(formattedPrompt); + } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractStringPromptTemplate.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractStringPromptTemplate.java deleted file mode 100644 index cb6013121..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/AbstractStringPromptTemplate.java +++ /dev/null @@ -1,34 +0,0 @@ -/* - * Copyright 2023 the original author or authors. - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * https://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -package org.springframework.ai.core.prompts; - -import java.util.Map; -import java.util.Optional; - -public abstract class AbstractStringPromptTemplate extends AbstractPromptTemplate implements PromptTemplateInput { - - public AbstractStringPromptTemplate(Optional outputParser) { - super(outputParser); - } - - @Override - public PromptValue formatAsPrompt(Map inputVariables) { - String formattedPrompt = formatAsString(inputVariables); - return new StringPromptValue(formattedPrompt); - } - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptInput.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptInput.java index 3bcae2b56..d3741c4a6 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptInput.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptInput.java @@ -16,15 +16,27 @@ package org.springframework.ai.core.prompts; +import java.util.Map; import java.util.Optional; -import java.util.Set; public interface PromptInput { - Set getInputVariables(); + // *Output* Optional getOutputParser(); + + // *Input* + + // This is the handoff point. These methods provide the "input", then the + // "Template" is "rendered" and the output is then used to construct a "Message" that gets sent to the "LLM" model. + + // Maybe should be called String renderAsString() and renderAsPrompt + // View in spring mvc has render(Map model, HttpServletRequest request, HttpServletResponse response) method. + String formatAsString(Map inputVariables); + + PromptValue formatAsPrompt(Map inputVariables); + // Leave out Partial Input Variables for now } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplate.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplate.java index 7e023a64c..17409ed26 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplate.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplate.java @@ -28,63 +28,38 @@ import org.antlr.runtime.TokenStream; import org.stringtemplate.v4.ST; import org.stringtemplate.v4.compiler.STLexer; -public class PromptTemplate extends AbstractStringPromptTemplate implements PromptTemplateInput { +public class PromptTemplate extends AbstractPromptTemplate implements PromptTemplateInput { private String template; private TemplateFormat templateFormat = TemplateFormat.FSTRING; public PromptTemplate(String template) { - super(Optional.empty()); - this.template = template; - } - - public PromptTemplate(String template, boolean validate) { - super(Optional.empty()); - validateTemplate(validate); + super(); this.template = template; } public PromptTemplate(String template, TemplateFormat templateFormat) { - super(Optional.empty()); + super(); this.template = template; this.templateFormat = templateFormat; } - public PromptTemplate(String template, TemplateFormat templateFormat, boolean validate) { - super(Optional.empty()); - validateTemplate(validate); - this.template = template; - this.templateFormat = templateFormat; - } - public PromptTemplate(String template, Optional outputParser) { + public PromptTemplate(String template, OutputParser outputParser) { super(outputParser); this.template = template; } - public PromptTemplate(String template, Optional outputParser, boolean validate) { + public PromptTemplate(String template, OutputParser outputParser, TemplateFormat templateFormat) { super(outputParser); - validateTemplate(validate); - this.template = template; - } - - public PromptTemplate(String template, Optional outputParser, TemplateFormat templateFormat) { - super(outputParser); - this.template = template; - this.templateFormat = templateFormat; - } - - public PromptTemplate(String template, Optional outputParser, TemplateFormat templateFormat, - boolean validate) { - super(outputParser); - validateTemplate(validate); this.template = template; this.templateFormat = templateFormat; } @Override public String formatAsString(Map inputVariables) { + validate(); // Only "F-String" for now ST st = new ST(this.template, '{', '}'); for (Entry stringObjectEntry : inputVariables.entrySet()) { @@ -103,8 +78,7 @@ public class PromptTemplate extends AbstractStringPromptTemplate implements Prom return this.templateFormat; } - @Override - public Set getInputVariables() { + protected Set getInputVariables() { ST st = new ST(this.template, '{', '}'); TokenStream tokens = st.impl.tokens; return IntStream.range(0, tokens.range()) @@ -114,12 +88,10 @@ public class PromptTemplate extends AbstractStringPromptTemplate implements Prom .collect(Collectors.toSet()); } - private void validateTemplate(boolean validate) { - if (!validate) { - return; - } + public void validate() { try { ST st = new ST(this.template, '{', '}'); + // TODO is doing this test even necessary, if it parsed correctly in the ctor, there should be no issues. Set inputVariables = getInputVariables(); for (String inputVariable : inputVariables) { st.add(inputVariable, "foo"); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplateInput.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplateInput.java index 30990cff4..92a12e16c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplateInput.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/prompts/PromptTemplateInput.java @@ -22,4 +22,6 @@ public interface PromptTemplateInput extends PromptInput { TemplateFormat getTemplateFormat(); + // *Validation* + void validate(); } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/core/prompts/PromptTests.java b/spring-ai-core/src/test/java/org/springframework/ai/core/prompts/PromptTests.java index fd4e5fb26..77abe671e 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/core/prompts/PromptTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/core/prompts/PromptTests.java @@ -16,6 +16,8 @@ package org.springframework.ai.core.prompts; +import java.util.ArrayList; +import java.util.List; import java.util.Set; import org.assertj.core.api.Assertions; @@ -59,7 +61,8 @@ class PromptTests { void testBadTemplateString() { String template = "This is a {foo test"; Assertions.assertThatThrownBy(() -> { - PromptTemplate promptTemplate = new PromptTemplate(template, true); + PromptTemplate promptTemplate = new PromptTemplate(template); + promptTemplate.validate(); }).isInstanceOf(IllegalArgumentException.class).hasMessage("The template string is not valid."); }