diff --git a/.editorconfig b/.editorconfig index 42a72e773b..9da7cf71b7 100644 --- a/.editorconfig +++ b/.editorconfig @@ -52,6 +52,7 @@ ij_java_blank_lines_after_package = 1 ij_java_blank_lines_around_class = 1 ij_java_blank_lines_around_field = 0 ij_java_blank_lines_around_field_in_interface = 0 +ij_java_blank_lines_around_field_with_annotations = 0 ij_java_blank_lines_around_initializer = 1 ij_java_blank_lines_around_method = 1 ij_java_blank_lines_around_method_in_interface = 1 @@ -116,7 +117,7 @@ ij_java_generate_final_locals = true ij_java_generate_final_parameters = true ij_java_generate_use_type_annotation_before_type = true ij_java_if_brace_force = never -ij_java_imports_layout = *, |, javax.**, java.**, |, $* +ij_java_imports_layout = $*,|,javax.**,java.**,* ij_java_indent_case_from_switch = true ij_java_insert_inner_class_imports = false ij_java_insert_override_annotation = true diff --git a/check-split-packages.sh b/check-split-packages.sh index 9da053c3e2..9e2ad460b1 100755 --- a/check-split-packages.sh +++ b/check-split-packages.sh @@ -12,8 +12,8 @@ echo "🔍 Scanning all JARs under: $ROOT_DIR (excluding test JARs)" # Get all JAR files first, excluding test JARs jar_files=() while IFS= read -r jar; do - # Skip JAR files ending with "-tests.jar" - if [[ "$jar" != *"-tests.jar" ]]; then + # Skip JAR files ending with "-tests.jar" and anything in the integration-tests directory + if [[ "$jar" != *"-tests.jar" ]] && [[ "$jar" != "./integration-tests/"* ]]; then jar_files+=("$jar") fi done < <(find "$ROOT_DIR" -type f -name "*.jar") diff --git a/docs/docs/tutorials/classification.md b/docs/docs/tutorials/classification.md index 9cd8386bb5..bc3ebdf32d 100644 --- a/docs/docs/tutorials/classification.md +++ b/docs/docs/tutorials/classification.md @@ -1,5 +1,5 @@ --- -sidebar_position: 12 +sidebar_position: 13 --- # Classification diff --git a/docs/docs/tutorials/embedding-stores.md b/docs/docs/tutorials/embedding-stores.md index d8067ff0d6..72965031be 100644 --- a/docs/docs/tutorials/embedding-stores.md +++ b/docs/docs/tutorials/embedding-stores.md @@ -1,5 +1,5 @@ --- -sidebar_position: 13 +sidebar_position: 14 --- # Embedding (Vector) Stores diff --git a/docs/docs/tutorials/guardrails.md b/docs/docs/tutorials/guardrails.md new file mode 100644 index 0000000000..82fd21d5ec --- /dev/null +++ b/docs/docs/tutorials/guardrails.md @@ -0,0 +1,539 @@ +--- +sidebar_position: 12 +toc_max_heading_level: 5 +--- + +import useBaseUrl from '@docusaurus/useBaseUrl'; +import ThemedImage from '@theme/ThemedImage'; +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Guardrails + +Guardrails are mechanisms that let you validate the input and output of the LLM to ensure it meets your expectations. You can do some of the following things with guardrails: +- Verify the user input is not out of scope +- Ensure the input meets some criteria before calling the LLM (i.e. guard against a [prompt injection attack](https://genai.owasp.org/llmrisk/llm01-prompt-injection/)) +- Ensure the output format is correct (i.e. it is a JSON document with the correct schema) +- Ensure the LLM output is coherent with business rules and constraints (i.e. if this is a chatbot of company X, the response should not contain any reference to a competitor Y). +- Detect hallucinations + +Those are just examples. You can do many other things with guardrails. + +:::note +Guardrails are only available when using [AI Services](/tutorials/ai-services). They are a higher-level construct that can not be applied to a `ChatModel` or `StreamingChatModel`. +::: + +; + +The implementation was originally done in the [Quarkus LangChain4j extension](https://docs.quarkiverse.io/quarkus-langchain4j/dev/) and was backported here. + +## Implementing Guardrails + +Ideally, guardrail implementations should follow the [single responsibility principle](https://en.wikipedia.org/wiki/Single-responsibility_principle), meaning that each guardrail class should validate one thing. Then, chain guardrails together to guard against multiple things. + +The order of guardrails in the chain is important. The first guardrail in the chain to fail will trigger the overall failure. Ensure guardrails that catch the most failures are early in the chain, whereas more specific guardrails that may fail very infrequently are towards the end of the chain. + +Also keep in mind that guardrails can themselves call other services or even invoke other LLM interactions. If these kinds of guardrails have an execution penalty or monetary cost associated with them, make sure you take that into account. You might want to put more expensive guardrails towards the end of the chain. + +:::note +The term _expensive_ can mean that something takes some time to execute or has a monetary value associated with it. +::: + +## Input Guardrails + +Input guardrails are functions invoked before the LLM is called. Failing an input guardrail prevents the LLM from being called. Input guardrails are the last step prior to calling the LLM. They are invoked _after_ any [RAG](/tutorials/rag) operations have happened. + +### Implementing Input Guardrails + +Input guardrails are implemented by implementing the [`InputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java) interface. The `InputGuardrail` interface has two variants of the `validate` method, at least one of which needs to be implemented: + +```java +InputGuardrailResult validate(UserMessage userMessage); +InputGuardrailResult validate(InputGuardrailRequest params); +``` + +The first variant is used for simple guardrails, or when the guardrail only needs access to the [`UserMessage`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/data/message/UserMessage.java). + +The second variant is for more complex guardrails that need more information, such as the chat memory/history, user message template, augmentation results, or variables that were passed to the template. See [`InputGuardrailRequest`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailRequest.java) for more information. + +Some examples of things you could do: +- Check that there are enough documents in the augmentation results +- Ensure the user is not asking the same question multiple times +- Mitigate potential prompt injection attack + +Input guardrails can be used whether the operation is synchronous or asynchronous/streaming. + +### Input Guardrail Outcomes + +Input guardrails can have the following outcomes. There are helper methods on the `InputGuardrail` interface that can provide the outcomes: + +| Outcome | Helper method on `InputGuardrail` | Description | +|:------------------------------------|:--------------------------------------------------|:---------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| **_success_** | `success()` | - The input is valid.
- The next guardrail in the chain is executed.
- The LLM is called if the last guardrail passes. | +| **_success with alternate result_** | `successWith(String)` | Similar to **_success_** except the user message is altered before proceeding to the next step (next guardrail in the chain or calling the LLM). | +| **_failure_** | `failure(String)` or `failure(String, Throwable)` | - The input is invalid but the next guardrails in the chain continue to be executed in order to accumulate all possible validation problems.
- The LLM is not called. | +| **_fatal_** | `fatal(String)` or `fatal(String, Throwable)` | - The input is invalid and execution is halted with an `InputGuardrailException`.
- The LLM is not called. | + +### Declaring Input Guardrails + +There are several ways to declare input guardrails, listed here in order of precedence: +1. [`InputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java) implementation class names or instances set directly on the [`AiServices`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java) builder. +2. [`@InputGuardrails` annotations](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java) placed on an individual [AI Service](/tutorials/ai-services) method. +3. [`@InputGuardrails` annotation](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java) placed on an [AI Service](/tutorials/ai-services) class. +Regardless of how they are declared, input guardrails are always executed in the order they appear in the list. + +#### `AiServices` builder + +[`InputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java) implementation class names or instances set directly on the [`AiServices`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java) builder have the highest precedence, meaning if it is declared in any other ways, the one declared directly on the builder will be the one used. + +```java +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .inputGuardrailClasses(FirstInputGuardrail.class, SecondInputGuardrail.class) + .build(); +``` + +or + +```java +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .inputGuardrails(new FirstInputGuardrail(), new SecondInputGuardrail()) + .build(); +``` + +In the first scenario, classes that implement `InputGuardrail` are passed. New instances of these classes are created dynamically using reflection. + +:::info +The way classes are converted to instances can be customized. For example, frameworks that use dependency injection (like [Quarkus](https://quarkus.io) or [Spring](https://spring.io)) can use [extension points](#extension-points) to provide instances based on how they manage class instances rather than creating new instances via reflection each time. +::: + +#### Annotation on individual AI Service methods + +[`@InputGuardrails` annotations](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java) placed on an individual [AI Service](/tutorials/ai-services) methods have the next highest precedence. + +```java +public interface Assistant { + @InputGuardrails({ FirstInputGuardrail.class, SecondInputGuardrail.class }) + String chat(String question); + + String doSomethingElse(String question); +} + +var assistant = AiServices.create(Assistant.class, chatModel); +``` + +In this example, only the `chat` method has guardrails. +- On the `chat` method, `FirstInputGuardrail` is invoked first. +- Only if it is successful will the LLM be called. +- `SecondInputGuardrail` will only be invoked if `FirstInputGuardrail` does not result in a **_fatal_** result. +- Either `FirstInputGuardrail` or `SecondInputGuardrail` could re-write the user message. +- If `FirstInputGuardrail` re-writes the user message, then `SecondInputGuardrail` will receive the new user message as input. + +The `doSomethingElse` method does not have any guardrails. + +#### Annotation on the AI Service class + +[`@InputGuardrails` annotation](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java) placed on an [AI Service](/tutorials/ai-services) class has the lowest precedence. + +```java +@InputGuardrails({ FirstInputGuardrail.class, SecondInputGuardrail.class }) +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.create(Assistant.class, chatModel); +``` + +In this example, both the `chat` and `doSomethingElse` methods have the guardrails. +- Just like in the previous example, `FirstInputGuardrail` is invoked first. +- Only if it is successful will the LLM be called. +- `SecondInputGuardrail` will only be invoked if `FirstInputGuardrail` does not result in a **_fatal_** result. +- Either `FirstInputGuardrail` or `SecondInputGuardrail` could re-write the user message. +- If `FirstInputGuardrail` re-writes the user message, then `SecondInputGuardrail` will receive the new user message as input. + +### Unit Testing Input Guardrails + +There are some unit testing utilities based on [AssertJ](https://assertj.github.io/doc/) in the `langchain4j-test` module. + + + + ```xml + + dev.langchain4j + langchain4j-test + test + + ``` + + + ```groovy + testImplementation 'dev.langchain4j:langchain4j-test' + ``` + + + ```kotlin + testImplementation("dev.langchain4j:langchain4j-test") + ``` + + + +Once you have the dependency, you can perform these kinds of validations: + +```java +import static dev.langchain4j.test.guardrail.GuardrailAssertions.assertThat; + +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.GuardrailResult.Result; + +class Tests { + MyInputGuardrail inputGuardrail = new MyInputGuardrail(); + + @Test + void test() { + var userMessage = UserMessage.from("Some user message"); + var result = inputGuardrail.validate(userMessage); + + // These are just some examples of what you can do + assertThat(result) + .isSuccessful() + .hasResult(Result.FATAL) + .hasFailures() + .hasSingleFailureWithMessage("Prompt injection detected") + .assertSingleFailureSatisfied(failure -> assertThat(failure)...) + .withFailures()..... + } +} +``` + +:::info +See the [`GuardrailAssertions`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailAssertions.java) and [`InputGuardrailResultAssert`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/InputGuardrailResultAssert.java) classes for more details. +::: + +## Output Guardrails + +Output guardrails are functions executed after the LLM has produced its output. Failing an output guardrail allows for more advanced scenarios, such as [retrying](#retry) or [reprompting](#reprompt), to help improve the response. They are invoked _after_ all other operations, including function/tool calls, have happened. + +### Implementing Output Guardrails + +Similar to input guardrails, output guardrails are implemented by implementing the [`OutputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java) interface. The `OutputGuardrail` interface has two variants of the `validate` method, at least one of which needs to be implemented: + +```java +OutputGuardrailResult validate(AiMessage responseFromLLM); +OutputGuardrailResult validate(OutputGuardrailRequest params); +``` + +The first variant is used for simple guardrails, or when the guardrail only needs access to the resulting [`AiMessage`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/data/message/AiMessage.java). + +The second variant is for more complex guardrails that need more information, such as the entire chat response, chat memory/history, user message template, or variables that were passed to the template. See [`OutputGuardrailRequest`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailRequest.java) for more information. + +Some examples of things you could do: +- Ensure the output format is correct (i.e. it is a JSON document with the correct schema) +- Detect an LLM hallucination +- Validate that the LLM response contains certain information + +### Output Guardrail Outcomes + +Output guardrails can have the following outcomes. There are helper methods on the `OutputGuardrail` interface that can provide the outcomes: + +| Outcome | Helper method on `OutputGuardrail` | Description | +|:---------------------------|:--------------------------------------------------------------------|:----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| **_success_** | `success()` | - The output is valid.
- The next guardrail in the chain is executed. If the last guardrail passes the output is returned to the caller. | +| **_success with rewrite_** | `successWith(String)` or `successWith(String, Object)` | -Similar to **_success_** except the output isn't valid in its original form and has been rewritten to make it valid.
- The next guardrail is executed against the rewritten output. If the last guardrail passes the output is returned to the caller. | +| **_failure_** | `failure(String)` or `failure(String, Throwable)` | - The output is invalid but the next guardrails in the chain continue to be executed in order to accumulate all possible validation problems.
- The validation failure is returned to the user as an `OutputGuardrailException`. | +| **_fatal_** | `fatal(String)` or `fatal(String, Throwable)` | The output is invalid and execution is halted with an `OutputGuardrailException` thrown to the caller. | +| **_fatal with retry_** | `retry(String)` or `retry(String, Throwable)` | - Similar to **_fatal_** except the LLM is called again with the same prompt and chat history as the original call.
- If the failure persists after a [configurable number of retries](#configuration) then execution is halted with an `OutputGuardrailException` thrown to the caller.
- If the guardrail passes after a retry, the entire chain of guardrails are re-executed from the beginning. | +| **_fatal with reprompt_** | `reprompt(String, String)` or `reprompt(String, Throwable, String)` | - Similar to **_fatal with retry_** except the LLM is called again with a new prompt supplied by the guardrail.
- In this situation, the guardrail supplies an additional message to append to the previous user message, then sends a new request to the LLM with the new user message and original chat history.
- If the failure persists after a [configurable number of retries](#configuration) then execution is halted with an `OutputGuardrailException` thrown to the caller.
- If the guardrail passes after a reprompt, the entire chain of guardrails are re-executed from the beginning. | + +### Declaring Output Guardrails + +There are several ways to declare output guardrails, listed here in order of precedence: +1. [`OutputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java) implementation class names or instances set directly on the [`AiServices`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java) builder. +2. [`@OutputGuardrails` annotations](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java) placed on an individual [AI Service](/tutorials/ai-services) method. +3. [`@OutputGuardrails` annotation](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java) placed on an [AI Service](/tutorials/ai-services) class. + +Regardless of how they are declared, output guardrails are always executed in the order they appear in the list. + +#### `AiServices` builder + +[`OutputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java) implementation class names or instances set directly on the [`AiServices`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java) builder have the highest precedence, meaning if it is declared in any other ways, the one declared on the builder will be the one used. + +```java +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .outputGuardrailClasses(FirstOutputGuardrail.class, SecondOutputGuardrail.class) + .build(); +``` + +or + +```java +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .outputGuardrails(new FirstOutputGuardrail(), new SecondOutputGuardrail()) + .build(); +``` + +In the first scenario, classes that implement `OutputGuardrail` are passed. New instances of these classes are created dynamically using reflection. + +:::info +The way classes are converted to instances can be customized. For example, frameworks that use dependency injection (like [Quarkus](https://quarkus.io) or [Spring](https://spring.io)) can use [extension points](#extension-points) to provide instances based on how they manage class instances rather than creating new instances via reflection each time. +::: + +#### Annotation on individual AI Service methods + +[`@OutputGuardrails` annotations](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java) placed on ndividual [AI Service](/tutorials/ai-services) methods have the next highest precendence. + +```java +public interface Assistant { + @OutputGuardrails({ FirstOutputGuardrail.class, SecondOutputGuardrail.class }) + String chat(String question); + + String doSomethingElse(String question); +} + +var assistant = AiServices.create(Assistant.class, chatModel); +``` + +In this example, only the `chat` method has guardrails. +- On the `chat` method, `FirstOutputGuardrail` is invoked first. +- Only if it is successful will the result be returned to the caller. `SecondOutputGuardrail` will only be invoked if `FirstOutputGuardrail` does not result in a **_fatal_**, **_fatal with retry_**, or **_fatal with reprompt_** result. +- `SecondOutputGuardrail` will receive the output of `FirstOutputGuardrail`. +- If `SecondOutputGuardrail` succeeds after a retry or reprompt, then both `FirstOutputGuardrail` and `SecondOutputGuardrail` are re-executed. + +The `doSomethingElse` method does not have any guardrails. + +#### Annotation on the AI Service class + +[`@OutputGuardrails` annotation](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java) placed on an [AI Service](/tutorials/ai-services) class has the lowest precedence. + +```java +@OutputGuardrails({ FirstOutputGuardrail.class, SecondOutputGuardrail.class }) +public interface Assistant { + String chat(String question); + String doSomethingElse(String question); +} + +var assistant = AiServices.create(Assistant.class, chatModel); +``` + +In this example, both the `chat` and `doSomethingElse` methods have the guardrails. +- Just like in the previous example, `FirstOutputGuardrail` is invoked first. +- Only if it is successful will the result be returned to the caller. `SecondOutputGuardrail` will only be invoked if `FirstOutputGuardrail` does not result in a **_fatal_**, **_fatal with retry_**, or **_fatal with reprompt_** result. +- `SecondOutputGuardrail` will receive the output of `FirstOutputGuardrail`. +- If `SecondOutputGuardrail` succeeds after a retry or reprompt, then both `FirstOutputGuardrail` and `SecondOutputGuardrail` are re-executed. + +#### Configuration + +Output guardrails have the following additional configuration that can be supplied: + +| Configuration | Description | +|:--------------|:-----------------------------------------------------------------------------------------------------------------------------------------------------------| +| `maxRetries` | - The maximum number of retries for an output guardrail when performing a retry or reprompt.
- Defaults to `2`.
- Set to `0` to disable retries. | + +##### Annotation on individual AI Service methods + +```java +public interface MethodLevelAssistant { + @OutputGuardrails( + value = { FirstOutputGuardrail.class, SecondOutputGuardrail.class }, + maxRetries = 10 + ) + String chat(String question); +} + +var assistant = AiServices.create(MethodLevelAssistant.class, chatModel); +``` + +##### Annotation on the AI Service class + +```java +@OutputGuardrails( + value = { FirstOutputGuardrail.class, SecondOutputGuardrail.class }, + maxRetries = 10 +) +public interface ClassLevelAssistant { + String chat(String question); +} + +var assistant = AiServices.create(ClassLevelAssistant.class, chatModel); +``` + +##### `AiServices` builder + +```java +public interface Assistant { + String chat(String message); +} + +var outputGuardrailsConfig = OutputGuardrailsConfig.builder() + .maxRetries(10) + .build(); + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .outputGuardrailsConfig(outputGuardrailsConfig) + .outputGuardrailClasss(FirstOutputGuardrail.class, SecondOutputGuardrail.class) + .build(); +``` + +### Output Guardrails on Streaming Responses + +Output guardrails can also work for operations with streaming responses: + +```java +public interface StreamingAssistant { + @OutputGuardrails({ FirstOutputGuardrail.class, SecondOutputGuardrail.class }) + TokenStream streamingChat(String message); +} +``` + +In this scenario, the output guardrails will be executed once the entire stream is complete, or more specifically, when `TokenStream.onCompleteResponse` is called. `onPartialResponse` will be buffered and replayed once the guardrails succeed. + +In the situation where a **_retry_** or **_reprompt_** in the chain eventually succeeds, then the entire chain is re-executed _synchronously_. Each guardrail will be re-executed one after the other in the original order. Once the chain completes the result is passed into `TokenStream.onCompleteResponse`. + +### Out-of-the-box Output Guardrails + +There are several common use cases where implementations of an output guardrail are provided by LangChain4j: + +| Guardrail class | Description | +|:----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|:-------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| [`JsonExtractorOutputGuardrail`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrail.java) | An output guardrail that will check whether or not a response can be successfully deserialized from JSON to an object of a certain type.
- Uses a [Jackson ObjectMapper](https://github.com/FasterXML/jackson-databind) to try and deserialize an object.
- The LLM is reprompted if the response can't be deserialized into the expected object type.
- Can be used as-is, or can be extended and customized (there are several `protected` methods that can be overridden to customize behavior). | + +### Unit Testing Output Guardrails + +There are some unit testing utilities based on [AssertJ](https://assertj.github.io/doc/) in the `langchain4j-test` module. + + + + ```xml + + dev.langchain4j + langchain4j-test + test + + ``` + + + ```groovy + testImplementation 'dev.langchain4j:langchain4j-test' + ``` + + + ```kotlin + testImplementation("dev.langchain4j:langchain4j-test") + ``` + + + +Once you have the dependency, you can perform these kinds of validations: + +```java +import static dev.langchain4j.test.guardrail.GuardrailAssertions.assertThat; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.GuardrailResult.Result; + +class Tests { + MyOutputGuardrail outputGuardrail = new MyOutputGuardrail(); + + @Test + void test() { + var aiMessage = AiMessage.from("Some output"); + var result = outputGuardrail.validate(aiMessage); + + // These are just some examples of what you can do + assertThat(result) + .isSuccessful() + .hasResult(Result.FATAL) + .hasFailures() + .hasSingleFailureWithMessage("Hallucination detected!") + .hasSingleFailureWithMessageAndReprompt("Hallucination detected!", "Please LLM don't hallucinate!") + .assertSingleFailureSatisfied(failure -> assertThat(failure)...) + .withFailures()..... + } +} +``` + +:::info +See the [`GuardrailAssertions`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailAssertions.java) and [`OutputGuardrailResultAssert`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/OutputGuardrailResultAssert.java) classes for more details. +::: + +## Mixing and matching + +You can mix and match input and output guardrails however you like! + +```java +public class MyObjectJsonOutputGuardrail extends JsonExtractorOutputGuardrail { + public MyObjectJsonOutputGuardrail() { + super(MyObject.class); + } +} + +@InputGuardrails({ FirstInputGuardrail.class, SecondInputGuardrail.class }) +@OutputGuardrails(value = SomeOutputGuardrail.class, maxRetries = 5) +public interface Assistant { + String chat(String message); + + @InputGuardrails(PromptInjectionGuardrail.class) + @OutputGuardrails(MyObjectJsonOutputGuardrail.class) + MyObject chatAndReturnJson(String message); +} + +var outputGuardrailsConfig = OutputGuardrailsConfig.builder() + .maxRetries(10) + .build(); + +var assistant = AiServices.builder(Assistant.class) + .chatModel(chatModel) + .inputGuardrails(new AnotherInputGuardrail()) + .outputGuardrailsConfig(outputGuardrailsConfig) + .build(); +``` + +In this example, all the methods on the `Assistant` have a single input guardrail, `AnotherInputGuardrail`, because it is set on the `AiServices` builder. Additionally, all the output guardrails have a `maxRetries` value == `10`, because the config is also set on the `AiServices` builder. + +The `chat` method has a single output guardrail, `SomeOutputGuardrail`, with a `maxRetries` value == `10`. + +The `chatAndReturnJson` method a single output guardrail, `MyObjectJsonOutputGuardrail` with a `maxRetries` value == `10`. + +## Extension points + +The guardrail system was built in a composable way so it can be extended and reused in other downstream frameworks (such as [Quarkus](https://quarkus.io) or [Spring Boot](https://spring.io/projects/spring-boot)). This section describes some of the extension points or "hooks" that are provided. + +All of these extension points utilize the [Java Service Provider Interface (Java SPI)](https://www.baeldung.com/java-spi). + +| Extension point interface | Purpose | +|--------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|-------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| [`ClassInstanceFactory`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassInstanceFactory.java) | Provides instanceos of classes.
- Intended to delegate instance creation/retrieval to some other means.
- If not provided, uses reflection to create an instance using the default constructor.
- Other frameworks (like Quarkus or Spring) may use their own bean containers to provide instances of classes. Those frameworks would provide an implementation.
- A Quarkus implementation may look something like [`CDIClassInstanceFactory`](https://github.com/langchain4j/langchain4j/blob/main/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/CDIClassInstanceFactory.java)
- A Spring implementation may look something like [`ApplicationContextClassInstanceFactory`](https://github.com/langchain4j/langchain4j/blob/main/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/ApplicationContextClassInstanceFactory.java) | +| [`ClassMetadataProviderFactory`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassMetadataProviderFactory.java) | Provides access to class metadata.
- Used to scan the methods on `AiService` interfaces, and find and process the `@InputGuardrails`/`@OutputGuardrails` annotations.
- [`ReflectionBasedClassMetadataProviderFactory`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j/src/main/java/dev/langchain4j/classloading/ReflectionBasedClassMetadataProviderFactory.java) is the default implementation if no others are found, providing class metadata using reflection. | +| [`InputGuardrailsConfigBuilderFactory`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/InputGuardrailsConfigBuilderFactory.java) | - SPI for overriding and/or extending the default [`InputGuardrailsConfigBuilder`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/InputGuardrailsConfigBuilder.java)
- Other frameworks may provide their own implementation with extra additional configuration for input guardrails.
- Would also allow other frameworks to drive input guardrail configuration via some other mechanism (i.e. a properties file). | +| [`OutputGuardrailsConfigBuilderFactory`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/OutputGuardrailsConfigBuilderFactory.java) | - SPI for overriding and/or extending the default [`OutputGuardrailsConfigBuilder`](https://github.com/langchain4j/langchain4j/blob/main/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/OutputGuardrailsConfigBuilder.java)
- Other frameworks may provide their own implementation with extra additional configuration for output guardrails.
- Would also allow other frameworks to drive output guardrail configuration via some other mechanism (i.e. a properties file). | + diff --git a/docs/static/img/guardrails-dark-bg.png b/docs/static/img/guardrails-dark-bg.png new file mode 100644 index 0000000000..e2f53eb58d Binary files /dev/null and b/docs/static/img/guardrails-dark-bg.png differ diff --git a/docs/static/img/guardrails-light-bg.png b/docs/static/img/guardrails-light-bg.png new file mode 100644 index 0000000000..54c5b9df63 Binary files /dev/null and b/docs/static/img/guardrails-light-bg.png differ diff --git a/integration-tests/README.md b/integration-tests/README.md new file mode 100644 index 0000000000..c0c7418603 --- /dev/null +++ b/integration-tests/README.md @@ -0,0 +1,6 @@ +This contains other full "projects" that can use various LangChain4j features independently but yet aren't necessarily "integration tests". +Think of these are separate applications that may be testing some kind of functionality within LangChain4j. + +Think of where `ServiceLoader`s may be invoked - creating a `src/test/META-INF/services` for the service in one of the modules would then override the service being loaded for all tests, which isn't what's intended. + +Instead, we can create isolated projects in here for that. diff --git a/integration-tests/integration-tests-class-instance-loader/README.md b/integration-tests/integration-tests-class-instance-loader/README.md new file mode 100644 index 0000000000..24d6d9bc00 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/README.md @@ -0,0 +1 @@ +Some integration tests for the [`ClassInstanceLoader`](../../langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java). diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/README.md b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/README.md new file mode 100644 index 0000000000..1568bd92d8 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/README.md @@ -0,0 +1 @@ +Some integration tests for the [`ClassInstanceLoader`](../../../langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java), implemented with Quarkus. diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/pom.xml b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/pom.xml new file mode 100644 index 0000000000..fc9048bd75 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/pom.xml @@ -0,0 +1,81 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../../pom.xml + + + langchain4j-integration-tests-class-instance-loader-quarkus + LangChain4j :: Integration Tests :: Class Instance Loader :: Quarkus + Tests for the Class Instance Loading abstraction using Quarkus + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + quarkus-bom + io.quarkus.platform + 3.22.3 + + + + + + ${quarkus.platform.group-id} + ${quarkus.platform.artifact-id} + ${quarkus.platform.version} + pom + import + + + + + + + io.quarkus + quarkus-arc + + + + dev.langchain4j + langchain4j-core + 1.1.0-SNAPSHOT + + + + io.quarkus + quarkus-junit5 + test + + + + + + + ${quarkus.platform.group-id} + quarkus-maven-plugin + ${quarkus.platform.version} + true + + + + build + generate-code + generate-code-tests + + + + + + + diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Application.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Application.java new file mode 100644 index 0000000000..bfcdc0c3f8 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Application.java @@ -0,0 +1,20 @@ +package com.example; + +import io.quarkus.runtime.QuarkusApplication; +import io.quarkus.runtime.annotations.QuarkusMain; + +@QuarkusMain +public class Application implements QuarkusApplication { + private final Class1 class1; + private final Class2 class2; + + public Application(final Class1 class1, final Class2 class2) { + this.class1 = class1; + this.class2 = class2; + } + + @Override + public int run(final String... args) { + return 0; + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/CDIClassInstanceFactory.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/CDIClassInstanceFactory.java new file mode 100644 index 0000000000..f91adaa677 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/CDIClassInstanceFactory.java @@ -0,0 +1,13 @@ +package com.example; + +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import io.quarkus.logging.Log; +import jakarta.enterprise.inject.spi.CDI; + +public class CDIClassInstanceFactory implements ClassInstanceFactory { + @Override + public T getInstanceOfClass(Class clazz) { + Log.infof("Getting instance of class %s from CDI.", clazz.getName()); + return CDI.current().select(clazz).get(); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class1.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class1.java new file mode 100644 index 0000000000..c52ca4fdd0 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class1.java @@ -0,0 +1,6 @@ +package com.example; + +import jakarta.enterprise.context.ApplicationScoped; + +@ApplicationScoped +public class Class1 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class2.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class2.java new file mode 100644 index 0000000000..fc6b0954cc --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class2.java @@ -0,0 +1,6 @@ +package com.example; + +import jakarta.enterprise.context.ApplicationScoped; + +@ApplicationScoped +public class Class2 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class3.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class3.java new file mode 100644 index 0000000000..c8fb8e358d --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/java/com/example/Class3.java @@ -0,0 +1,3 @@ +package com.example; + +public class Class3 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory new file mode 100644 index 0000000000..1cea374fad --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory @@ -0,0 +1 @@ +com.example.CDIClassInstanceFactory diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/test/java/dev/langchain4j/classinstance/quarkus/ClassInstanceLoaderTests.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/test/java/dev/langchain4j/classinstance/quarkus/ClassInstanceLoaderTests.java new file mode 100644 index 0000000000..919c6e8c2a --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus/src/test/java/dev/langchain4j/classinstance/quarkus/ClassInstanceLoaderTests.java @@ -0,0 +1,48 @@ +package dev.langchain4j.classinstance.quarkus; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.CDIClassInstanceFactory; +import com.example.Class1; +import com.example.Class2; +import com.example.Class3; +import dev.langchain4j.classinstance.ClassInstanceLoader; +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import io.quarkus.test.junit.QuarkusTest; +import jakarta.enterprise.inject.UnsatisfiedResolutionException; +import jakarta.enterprise.inject.spi.CDI; +import java.util.ServiceLoader; +import org.junit.jupiter.api.Test; + +@QuarkusTest +class ClassInstanceLoaderTests { + @Test + void serviceLoaderFindsCorrectFactory() { + assertThat(ServiceLoader.load(ClassInstanceFactory.class).findFirst()) + .get() + .isInstanceOf(CDIClassInstanceFactory.class); + } + + @Test + void loadsClassInstances() { + var instance1 = ClassInstanceLoader.getClassInstance(Class1.class); + var instance2 = ClassInstanceLoader.getClassInstance(Class1.class); + var instance3 = ClassInstanceLoader.getClassInstance(Class2.class); + + assertThat(instance1).isNotNull().isInstanceOf(Class1.class); + assertThat(instance2) + .isNotNull() + .isInstanceOf(Class1.class) + .isEqualTo(instance1) + .isEqualTo(CDI.current().select(Class1.class).get()); + assertThat(instance3).isNotNull().isInstanceOf(Class2.class); + } + + @Test + void correctServiceLoader() { + assertThatExceptionOfType(UnsatisfiedResolutionException.class) + .isThrownBy(() -> ClassInstanceLoader.getClassInstance(Class3.class)) + .withMessage("No bean found for required type [class %s] and qualifiers [[]]", Class3.class.getName()); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/README.md b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/README.md new file mode 100644 index 0000000000..cf3cc1400d --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/README.md @@ -0,0 +1 @@ +Some integration tests for the [`ClassInstanceLoader`](../../../langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java), implemented with Spring. diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/pom.xml b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/pom.xml new file mode 100644 index 0000000000..7a0fd5ecd4 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/pom.xml @@ -0,0 +1,70 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../../pom.xml + + + langchain4j-integration-tests-class-instance-loader-spring + LangChain4j :: Integration Tests :: Class Instance Loader :: Spring + Tests for the Class Instance Loading abstraction using Spring + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + 2.0.16 + 1.5.16 + + + + + + + org.springframework.boot + spring-boot-dependencies + 3.4.5 + pom + import + + + + + + + org.springframework.boot + spring-boot-starter + + + + dev.langchain4j + langchain4j-core + 1.1.0-SNAPSHOT + + + + org.springframework.boot + spring-boot-starter-test + test + + + + + + + org.springframework.boot + spring-boot-maven-plugin + + + + diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/Application.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/Application.java new file mode 100644 index 0000000000..d43ac04352 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/Application.java @@ -0,0 +1,11 @@ +package com.example; + +import org.springframework.boot.SpringApplication; +import org.springframework.boot.autoconfigure.SpringBootApplication; + +@SpringBootApplication +public class Application { + public static void main(String[] args) { + SpringApplication.run(Application.class, args); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/ApplicationContextProvider.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/ApplicationContextProvider.java new file mode 100644 index 0000000000..245e4a9a8d --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/ApplicationContextProvider.java @@ -0,0 +1,20 @@ +package com.example; + +import org.springframework.beans.BeansException; +import org.springframework.context.ApplicationContext; +import org.springframework.context.ApplicationContextAware; +import org.springframework.stereotype.Component; + +@Component +public class ApplicationContextProvider implements ApplicationContextAware { + private static ApplicationContext context; + + @Override + public void setApplicationContext(ApplicationContext applicationContext) throws BeansException { + context = applicationContext; + } + + public static ApplicationContext getApplicationContext() { + return context; + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/ApplicationContextClassInstanceFactory.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/ApplicationContextClassInstanceFactory.java new file mode 100644 index 0000000000..08415774dd --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/ApplicationContextClassInstanceFactory.java @@ -0,0 +1,11 @@ +package com.example.classes; + +import com.example.ApplicationContextProvider; +import dev.langchain4j.spi.classloading.ClassInstanceFactory; + +public class ApplicationContextClassInstanceFactory implements ClassInstanceFactory { + @Override + public T getInstanceOfClass(Class clazz) { + return ApplicationContextProvider.getApplicationContext().getBean(clazz); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class1.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class1.java new file mode 100644 index 0000000000..d12eeea112 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class1.java @@ -0,0 +1,6 @@ +package com.example.classes; + +import org.springframework.stereotype.Component; + +@Component +public class Class1 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class2.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class2.java new file mode 100644 index 0000000000..78f3456334 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class2.java @@ -0,0 +1,6 @@ +package com.example.classes; + +import org.springframework.stereotype.Component; + +@Component +public class Class2 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class3.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class3.java new file mode 100644 index 0000000000..824d711253 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/java/com/example/classes/Class3.java @@ -0,0 +1,3 @@ +package com.example.classes; + +public class Class3 {} diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory new file mode 100644 index 0000000000..794326152f --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory @@ -0,0 +1 @@ +com.example.classes.ApplicationContextClassInstanceFactory diff --git a/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassInstanceLoaderTests.java b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassInstanceLoaderTests.java new file mode 100644 index 0000000000..fc72958f4d --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassInstanceLoaderTests.java @@ -0,0 +1,53 @@ +package dev.langchain4j.classinstance.spring; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.Application; +import com.example.classes.ApplicationContextClassInstanceFactory; +import com.example.classes.Class1; +import com.example.classes.Class2; +import com.example.classes.Class3; +import dev.langchain4j.classinstance.ClassInstanceLoader; +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import java.util.ServiceLoader; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.NoSuchBeanDefinitionException; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.ApplicationContext; + +@SpringBootTest(classes = Application.class) +class ClassInstanceLoaderTests { + @Autowired + ApplicationContext applicationContext; + + @Test + void serviceLoaderFindsCorrectFactory() { + assertThat(ServiceLoader.load(ClassInstanceFactory.class).findFirst()) + .get() + .isInstanceOf(ApplicationContextClassInstanceFactory.class); + } + + @Test + void loadsClassInstances() { + var instance1 = ClassInstanceLoader.getClassInstance(Class1.class); + var instance2 = ClassInstanceLoader.getClassInstance(Class1.class); + var instance3 = ClassInstanceLoader.getClassInstance(Class2.class); + + assertThat(instance1).isNotNull().isExactlyInstanceOf(Class1.class); + assertThat(instance2) + .isNotNull() + .isExactlyInstanceOf(Class1.class) + .isEqualTo(instance1) + .isEqualTo(this.applicationContext.getBean(Class1.class)); + assertThat(instance3).isNotNull().isExactlyInstanceOf(Class2.class); + } + + @Test + void correctServiceLoader() { + assertThatExceptionOfType(NoSuchBeanDefinitionException.class) + .isThrownBy(() -> ClassInstanceLoader.getClassInstance(Class3.class)) + .withMessage("No qualifying bean of type '%s' available", Class3.class.getName()); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/pom.xml b/integration-tests/integration-tests-class-instance-loader/pom.xml new file mode 100644 index 0000000000..e9bf4f5e11 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/pom.xml @@ -0,0 +1,32 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../pom.xml + + + langchain4j-integration-tests-class-instance-loader + LangChain4j :: Integration Tests :: Class Instance Loader + Tests for the Class Instance Loading abstraction + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + + dev.langchain4j + langchain4j-core + 1.1.0-SNAPSHOT + + + diff --git a/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/Classes.java b/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/Classes.java new file mode 100644 index 0000000000..f62f025daa --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/Classes.java @@ -0,0 +1,29 @@ +package com.example.classloading; + +public final class Classes { + private static final Classes INSTANCE = new Classes(); + + private Classes() {} + + public static Classes getInstance() { + return INSTANCE; + } + + public T getInstance(Class clazz) { + if (clazz == Class1.class) { + return (T) new Class1(); + } + + if (clazz == Class2.class) { + return (T) new Class2(); + } + + throw new IllegalArgumentException("Unknown class: %s".formatted(clazz.getName())); + } + + public static class Class1 {} + + public static class Class2 {} + + public static class Class3 {} +} diff --git a/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/GenericClassInstanceFactory.java b/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/GenericClassInstanceFactory.java new file mode 100644 index 0000000000..5ba6978790 --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/src/main/java/com/example/classloading/GenericClassInstanceFactory.java @@ -0,0 +1,10 @@ +package com.example.classloading; + +import dev.langchain4j.spi.classloading.ClassInstanceFactory; + +public class GenericClassInstanceFactory implements ClassInstanceFactory { + @Override + public T getInstanceOfClass(Class clazz) { + return Classes.getInstance().getInstance(clazz); + } +} diff --git a/integration-tests/integration-tests-class-instance-loader/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory b/integration-tests/integration-tests-class-instance-loader/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory new file mode 100644 index 0000000000..58abe7bd7b --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory @@ -0,0 +1 @@ +com.example.classloading.GenericClassInstanceFactory diff --git a/integration-tests/integration-tests-class-instance-loader/src/test/java/dev/langchain4j/classinstance/generic/ClassInstanceLoaderTests.java b/integration-tests/integration-tests-class-instance-loader/src/test/java/dev/langchain4j/classinstance/generic/ClassInstanceLoaderTests.java new file mode 100644 index 0000000000..e4e5557eac --- /dev/null +++ b/integration-tests/integration-tests-class-instance-loader/src/test/java/dev/langchain4j/classinstance/generic/ClassInstanceLoaderTests.java @@ -0,0 +1,41 @@ +package dev.langchain4j.classinstance.generic; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.classloading.Classes; +import com.example.classloading.GenericClassInstanceFactory; +import dev.langchain4j.classinstance.ClassInstanceLoader; +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import java.util.ServiceLoader; +import org.junit.jupiter.api.Test; + +class ClassInstanceLoaderTests { + @Test + void serviceLoaderFindsCorrectFactory() { + assertThat(ServiceLoader.load(ClassInstanceFactory.class).findFirst()) + .get() + .isInstanceOf(GenericClassInstanceFactory.class); + } + + @Test + void loadsClassInstances() { + var instance1 = ClassInstanceLoader.getClassInstance(Classes.Class1.class); + var instance2 = ClassInstanceLoader.getClassInstance(Classes.Class1.class); + var instance3 = ClassInstanceLoader.getClassInstance(Classes.Class2.class); + + assertThat(instance1).isNotNull().isExactlyInstanceOf(Classes.Class1.class); + assertThat(instance2) + .isNotNull() + .isExactlyInstanceOf(Classes.Class1.class) + .isNotEqualTo(instance1); + assertThat(instance3).isNotNull().isExactlyInstanceOf(Classes.Class2.class); + } + + @Test + void correctServiceLoader() { + assertThatExceptionOfType(IllegalArgumentException.class) + .isThrownBy(() -> ClassInstanceLoader.getClassInstance(Classes.Class3.class)) + .withMessage("Unknown class: %s", Classes.Class3.class.getName()); + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/README.md b/integration-tests/integration-tests-class-metadata-provider/README.md new file mode 100644 index 0000000000..32fa63f2ef --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/README.md @@ -0,0 +1 @@ +Some integration tests for the [`ClassMetadataProvider`](../../langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassMetadataProvider.java). diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/README.md b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/README.md new file mode 100644 index 0000000000..4a441815b2 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/README.md @@ -0,0 +1 @@ +Some integration tests for the [`ClassMetadataProvider`](../../../langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassMetadataProvider.java), implemented with Spring. diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/pom.xml b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/pom.xml new file mode 100644 index 0000000000..b47d5487b5 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/pom.xml @@ -0,0 +1,70 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../../pom.xml + + + langchain4j-integration-tests-class-metadata-provider-spring + LangChain4j :: Integration Tests :: Class Metadata Provider :: Spring + Tests for the Class Metadata Loading abstraction using Spring + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + 2.0.16 + 1.5.16 + + + + + + + org.springframework.boot + spring-boot-dependencies + 3.4.5 + pom + import + + + + + + + org.springframework.boot + spring-boot-starter + + + + dev.langchain4j + langchain4j + 1.1.0-SNAPSHOT + + + + org.springframework.boot + spring-boot-starter-test + test + + + + + + + org.springframework.boot + spring-boot-maven-plugin + + + + diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/Application.java b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/Application.java new file mode 100644 index 0000000000..d43ac04352 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/Application.java @@ -0,0 +1,11 @@ +package com.example; + +import org.springframework.boot.SpringApplication; +import org.springframework.boot.autoconfigure.SpringBootApplication; + +@SpringBootApplication +public class Application { + public static void main(String[] args) { + SpringApplication.run(Application.class, args); + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/Class1.java b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/Class1.java new file mode 100644 index 0000000000..db4a6cb7cc --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/Class1.java @@ -0,0 +1,19 @@ +package com.example.classes; + +import dev.langchain4j.Experimental; +import org.springframework.stereotype.Component; + +@Experimental("This is plain and boring!") +@Component +public class Class1 { + public void hello() {} + + @Experimental("Just trying things out") + public String goodbye() { + return "Goodbye!"; + } + + public static String wave() { + return "Wave!"; + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/SpringClassMetadataProviderFactory.java b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/SpringClassMetadataProviderFactory.java new file mode 100644 index 0000000000..86f6841a38 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/java/com/example/classes/SpringClassMetadataProviderFactory.java @@ -0,0 +1,29 @@ +package com.example.classes; + +import dev.langchain4j.spi.classloading.ClassMetadataProviderFactory; +import java.lang.annotation.Annotation; +import java.lang.reflect.Method; +import java.lang.reflect.Modifier; +import java.util.Optional; +import java.util.stream.Stream; +import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.util.ReflectionUtils; + +public class SpringClassMetadataProviderFactory implements ClassMetadataProviderFactory { + @Override + public Optional getAnnotation(Method method, Class annotationClass) { + return Optional.ofNullable(AnnotationUtils.findAnnotation(method, annotationClass)); + } + + @Override + public Optional getAnnotation(Class clazz, Class annotationClass) { + return Optional.ofNullable(AnnotationUtils.findAnnotation(clazz, annotationClass)); + } + + @Override + public Iterable getNonStaticMethodsOnClass(Class clazz) { + return Stream.of(ReflectionUtils.getDeclaredMethods(clazz)) + .filter(method -> !Modifier.isStatic(method.getModifiers())) + .toList(); + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory new file mode 100644 index 0000000000..0ae370bf4a --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory @@ -0,0 +1 @@ +com.example.classes.SpringClassMetadataProviderFactory diff --git a/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassMetadataProviderTests.java b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassMetadataProviderTests.java new file mode 100644 index 0000000000..65766587f0 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring/src/test/java/dev/langchain4j/classinstance/spring/ClassMetadataProviderTests.java @@ -0,0 +1,61 @@ +package dev.langchain4j.classinstance.spring; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.Application; +import com.example.classes.Class1; +import com.example.classes.SpringClassMetadataProviderFactory; +import dev.langchain4j.Experimental; +import dev.langchain4j.classloading.ClassMetadataProvider; +import java.lang.annotation.Target; +import java.lang.reflect.Method; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; +import org.junit.jupiter.api.Test; +import org.springframework.boot.test.context.SpringBootTest; + +@SpringBootTest(classes = Application.class) +class ClassMetadataProviderTests { + @Test + void serviceLoaderFindsCorrectFactory() { + assertThat(ClassMetadataProvider.getClassMetadataProviderFactory()) + .isInstanceOf(SpringClassMetadataProviderFactory.class); + } + + @Test + void loadsThingsCorrectly() { + var factory = ClassMetadataProvider.getClassMetadataProviderFactory(); + + assertThat(factory.getAnnotation(Class1.class, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is plain and boring!"); + + assertThat(factory.getAnnotation(Class1.class, Target.class)).isEmpty(); + + var methods = factory.getNonStaticMethodsOnClass(Class1.class); + + assertThat(methods).hasSize(2).extracting(Method::getName).containsExactlyInAnyOrder("hello", "goodbye"); + + var methodsByName = StreamSupport.stream(methods.spliterator(), false) + .collect(Collectors.toMap(Method::getName, method -> method)); + + var helloMethod = methodsByName.get("hello"); + var goodbyeMethod = methodsByName.get("goodbye"); + + assertThat(helloMethod).isNotNull(); + + assertThat(goodbyeMethod).isNotNull(); + + assertThat(factory.getAnnotation(goodbyeMethod, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("Just trying things out"); + + assertThat(factory.getAnnotation(goodbyeMethod, Target.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Experimental.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Target.class)).isEmpty(); + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/pom.xml b/integration-tests/integration-tests-class-metadata-provider/pom.xml new file mode 100644 index 0000000000..c4090cb50e --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/pom.xml @@ -0,0 +1,32 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../pom.xml + + + langchain4j-integration-tests-class-metadata-provider + LangChain4j :: Integration Tests :: Class Metadata Provider + Tests for the Class Metadata Loading abstraction + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + + dev.langchain4j + langchain4j + 1.1.0-SNAPSHOT + + + diff --git a/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/Class1.java b/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/Class1.java new file mode 100644 index 0000000000..b2421af503 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/Class1.java @@ -0,0 +1,17 @@ +package com.example.classloading; + +import dev.langchain4j.Experimental; + +@Experimental("This is plain and boring!") +public class Class1 { + public void hello() {} + + @Experimental("Just trying things out") + public String goodbye() { + return "Goodbye!"; + } + + public static String wave() { + return "Wave!"; + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/GenericClassMetadataProviderFactory.java b/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/GenericClassMetadataProviderFactory.java new file mode 100644 index 0000000000..a8b7476626 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/src/main/java/com/example/classloading/GenericClassMetadataProviderFactory.java @@ -0,0 +1,27 @@ +package com.example.classloading; + +import dev.langchain4j.spi.classloading.ClassMetadataProviderFactory; +import java.lang.annotation.Annotation; +import java.lang.reflect.Method; +import java.lang.reflect.Modifier; +import java.util.Optional; +import java.util.stream.Stream; + +public class GenericClassMetadataProviderFactory implements ClassMetadataProviderFactory { + @Override + public Optional getAnnotation(Method method, Class annotationClass) { + return Optional.ofNullable(method.getAnnotation(annotationClass)); + } + + @Override + public Optional getAnnotation(Class clazz, Class annotationClass) { + return Optional.ofNullable(clazz.getAnnotation(annotationClass)); + } + + @Override + public Iterable getNonStaticMethodsOnClass(Class clazz) { + return Stream.of(clazz.getDeclaredMethods()) + .filter(method -> !Modifier.isStatic(method.getModifiers())) + .toList(); + } +} diff --git a/integration-tests/integration-tests-class-metadata-provider/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory b/integration-tests/integration-tests-class-metadata-provider/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory new file mode 100644 index 0000000000..00dd72a045 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassMetadataProviderFactory @@ -0,0 +1 @@ +com.example.classloading.GenericClassMetadataProviderFactory diff --git a/integration-tests/integration-tests-class-metadata-provider/src/test/java/dev/langchain4j/classinstance/generic/ClassMetadataProviderTests.java b/integration-tests/integration-tests-class-metadata-provider/src/test/java/dev/langchain4j/classinstance/generic/ClassMetadataProviderTests.java new file mode 100644 index 0000000000..4a03ae11e0 --- /dev/null +++ b/integration-tests/integration-tests-class-metadata-provider/src/test/java/dev/langchain4j/classinstance/generic/ClassMetadataProviderTests.java @@ -0,0 +1,58 @@ +package dev.langchain4j.classinstance.generic; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.classloading.Class1; +import com.example.classloading.GenericClassMetadataProviderFactory; +import dev.langchain4j.Experimental; +import dev.langchain4j.classloading.ClassMetadataProvider; +import java.lang.annotation.Target; +import java.lang.reflect.Method; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; +import org.junit.jupiter.api.Test; + +class ClassMetadataProviderTests { + @Test + void serviceLoaderFindsCorrectFactory() { + assertThat(ClassMetadataProvider.getClassMetadataProviderFactory()) + .isInstanceOf(GenericClassMetadataProviderFactory.class); + } + + @Test + void loadsThingsCorrectly() { + var factory = ClassMetadataProvider.getClassMetadataProviderFactory(); + + assertThat(factory.getAnnotation(Class1.class, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is plain and boring!"); + + assertThat(factory.getAnnotation(Class1.class, Target.class)).isEmpty(); + + var methods = factory.getNonStaticMethodsOnClass(Class1.class); + + assertThat(methods).hasSize(2).extracting(Method::getName).containsExactlyInAnyOrder("hello", "goodbye"); + + var methodsByName = StreamSupport.stream(methods.spliterator(), false) + .collect(Collectors.toMap(Method::getName, method -> method)); + + var helloMethod = methodsByName.get("hello"); + var goodbyeMethod = methodsByName.get("goodbye"); + + assertThat(helloMethod).isNotNull(); + + assertThat(goodbyeMethod).isNotNull(); + + assertThat(factory.getAnnotation(goodbyeMethod, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("Just trying things out"); + + assertThat(factory.getAnnotation(goodbyeMethod, Target.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Experimental.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Target.class)).isEmpty(); + } +} diff --git a/integration-tests/integration-tests-guardrails/README.md b/integration-tests/integration-tests-guardrails/README.md new file mode 100644 index 0000000000..496b2541bf --- /dev/null +++ b/integration-tests/integration-tests-guardrails/README.md @@ -0,0 +1 @@ +Some integration tests for Guardrails on AiServices diff --git a/integration-tests/integration-tests-guardrails/pom.xml b/integration-tests/integration-tests-guardrails/pom.xml new file mode 100644 index 0000000000..d95687846d --- /dev/null +++ b/integration-tests/integration-tests-guardrails/pom.xml @@ -0,0 +1,38 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-integration-tests-parent + 1.1.0-beta7-SNAPSHOT + ../pom.xml + + + langchain4j-integration-tests-guardrails + LangChain4j :: Integration Tests :: Guardrails + Tests for Guardrails + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + + dev.langchain4j + langchain4j + 1.1.0-SNAPSHOT + + + + org.junit.jupiter + junit-jupiter-params + test + + + diff --git a/integration-tests/integration-tests-guardrails/src/main/java/com/example/InputGuardrailValidation.java b/integration-tests/integration-tests-guardrails/src/main/java/com/example/InputGuardrailValidation.java new file mode 100644 index 0000000000..7047a5dad8 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/main/java/com/example/InputGuardrailValidation.java @@ -0,0 +1,38 @@ +package com.example; + +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.InputGuardrailResult; +import java.util.Map; + +public class InputGuardrailValidation implements InputGuardrail { + private static final InputGuardrailValidation INSTANCE = new InputGuardrailValidation(); + private InputGuardrailRequest params; + + private InputGuardrailValidation() {} + + public static InputGuardrailValidation getInstance() { + return INSTANCE; + } + + public InputGuardrailResult validate(InputGuardrailRequest params) { + this.params = params; + return success(); + } + + public void reset() { + this.params = null; + } + + public String spyUserMessageTemplate() { + return params.requestParams().userMessageTemplate(); + } + + public String spyUserMessageText() { + return params.userMessage().singleText(); + } + + public Map spyVariables() { + return params.requestParams().variables(); + } +} diff --git a/integration-tests/integration-tests-guardrails/src/main/java/com/example/OutputGuardrailValidation.java b/integration-tests/integration-tests-guardrails/src/main/java/com/example/OutputGuardrailValidation.java new file mode 100644 index 0000000000..c1dcead102 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/main/java/com/example/OutputGuardrailValidation.java @@ -0,0 +1,34 @@ +package com.example; + +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import java.util.Map; + +public class OutputGuardrailValidation implements OutputGuardrail { + private static final OutputGuardrailValidation INSTANCE = new OutputGuardrailValidation(); + private OutputGuardrailRequest params; + + private OutputGuardrailValidation() {} + + public static OutputGuardrailValidation getInstance() { + return INSTANCE; + } + + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.params = params; + return success(); + } + + public void reset() { + this.params = null; + } + + public String spyUserMessageTemplate() { + return params.requestParams().userMessageTemplate(); + } + + public Map spyVariables() { + return params.requestParams().variables(); + } +} diff --git a/integration-tests/integration-tests-guardrails/src/main/java/com/example/SingletonClassInstanceFactory.java b/integration-tests/integration-tests-guardrails/src/main/java/com/example/SingletonClassInstanceFactory.java new file mode 100644 index 0000000000..e5b70b5320 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/main/java/com/example/SingletonClassInstanceFactory.java @@ -0,0 +1,49 @@ +package com.example; + +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import java.lang.reflect.InvocationTargetException; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; + +/** + * A factory for providing class instances that are singletons + */ +public class SingletonClassInstanceFactory implements ClassInstanceFactory { + private static final ConcurrentMap, Object> INSTANCES = new ConcurrentHashMap<>(5); + + public static T getInstance(Class clazz) { + if (clazz == InputGuardrailValidation.class) { + return (T) InputGuardrailValidation.getInstance(); + } + + if (clazz == OutputGuardrailValidation.class) { + return (T) OutputGuardrailValidation.getInstance(); + } + + return getClassInstance(clazz); + } + + public static void clearInstances() { + INSTANCES.clear(); + } + + @Override + public T getInstanceOfClass(Class clazz) { + return getInstance(clazz); + } + + private static T getClassInstance(Class clazz) { + return (T) INSTANCES.computeIfAbsent(clazz, SingletonClassInstanceFactory::createNewClassInstance); + } + + private static T createNewClassInstance(Class clazz) { + try { + return clazz.getDeclaredConstructor().newInstance(); + } catch (InstantiationException + | IllegalAccessException + | InvocationTargetException + | NoSuchMethodException e) { + throw new RuntimeException(e); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory b/integration-tests/integration-tests-guardrails/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory new file mode 100644 index 0000000000..1147091831 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/main/resources/META-INF/services/dev.langchain4j.spi.classloading.ClassInstanceFactory @@ -0,0 +1 @@ +com.example.SingletonClassInstanceFactory diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/BaseGuardrailTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/BaseGuardrailTests.java new file mode 100644 index 0000000000..e71c6e7d54 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/BaseGuardrailTests.java @@ -0,0 +1,73 @@ +package dev.langchain4j.service.guardrail; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.memory.chat.MessageWindowChatMemory; +import dev.langchain4j.service.AiServices; +import dev.langchain4j.service.TokenStream; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Function; +import java.util.function.Supplier; +import org.junit.jupiter.api.BeforeEach; + +public abstract class BaseGuardrailTests { + @BeforeEach + void beforeEach() { + SingletonClassInstanceFactory.clearInstances(); + } + + static String execute(Supplier aiServiceInvocation) throws InterruptedException { + var latch = new CountDownLatch(1); + var value = new AtomicReference(); + + aiServiceInvocation + .get() + .onError(t -> { + throw new RuntimeException(t); + }) + .onPartialResponse(token -> {}) + .onCompleteResponse(response -> { + value.set(response.aiMessage().text()); + latch.countDown(); + }) + .start(); + + latch.await(10, TimeUnit.SECONDS); + + return value.get(); + } + + static T createAiService(Class clazz, Function, AiServices> builderCustomizer) { + return createAiService(clazz, List.of(), List.of(), builderCustomizer); + } + + static T createAiService(Class clazz) { + return createAiService(clazz, Function.identity()); + } + + static T createAiService( + Class clazz, + List> inputGuardrailClasses, + List> outputGuardrailClasses) { + + return createAiService(clazz, inputGuardrailClasses, outputGuardrailClasses, Function.identity()); + } + + static T createAiService( + Class clazz, + List> inputGuardrailClasses, + List> outputGuardrailClasses, + Function, AiServices> builderCustomizer) { + + var builder = AiServices.builder(clazz) + .chatMemoryProvider(memoryId -> MessageWindowChatMemory.withMaxMessages(10)) + .inputGuardrailClasses(inputGuardrailClasses) + .outputGuardrailClasses(outputGuardrailClasses); + + return builderCustomizer.apply(builder).build(); + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/EchoChatModel.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/EchoChatModel.java new file mode 100644 index 0000000000..fd5164e6d4 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/EchoChatModel.java @@ -0,0 +1,19 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; + +/** + * A {@link ChatModel} that echoes out the {@link UserMessage} + */ +public class EchoChatModel implements ChatModel { + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + var userMessage = ((UserMessage) chatRequest.messages().get(0)).singleText(); + + return ChatResponse.builder().aiMessage(AiMessage.from(userMessage)).build(); + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputAndOutputGuardrailsTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputAndOutputGuardrailsTests.java new file mode 100644 index 0000000000..999fc82390 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputAndOutputGuardrailsTests.java @@ -0,0 +1,211 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailException; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class InputAndOutputGuardrailsTests extends BaseGuardrailTests { + MyAiService service = MyAiService.create(); + + @Test + void ok() { + var okIn = SingletonClassInstanceFactory.getInstance(MyOkInputGuardrail.class); + var okOut = SingletonClassInstanceFactory.getInstance(MyOkOutputGuardrail.class); + assertThat(okIn.getSpy()).isEqualTo(0); + assertThat(okOut.getSpy()).isEqualTo(0); + service.bothOk("1", "foo"); + assertThat(okIn.getSpy()).isEqualTo(1); + assertThat(okOut.getSpy()).isEqualTo(1); + } + + @Test + void inKo() { + var koIn = SingletonClassInstanceFactory.getInstance(MyKoInputGuardrail.class); + var okOut = SingletonClassInstanceFactory.getInstance(MyOkOutputGuardrail.class); + assertThat(koIn.getSpy()).isEqualTo(0); + assertThat(okOut.getSpy()).isEqualTo(0); + + assertThatExceptionOfType(InputGuardrailException.class) + .isThrownBy(() -> service.inKo("2", "foo")) + .withCauseExactlyInstanceOf(ValidationException.class) + .havingRootCause() + .withMessage("boom"); + assertThat(koIn.getSpy()).isEqualTo(1); + assertThat(okOut.getSpy()).isEqualTo(0); + } + + @Test + void outKo() { + var okIn = SingletonClassInstanceFactory.getInstance(MyOkInputGuardrail.class); + var koOut = SingletonClassInstanceFactory.getInstance(MyKoOutputGuardrail.class); + assertThat(okIn.getSpy()).isEqualTo(0); + assertThat(koOut.getSpy()).isEqualTo(0); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> service.outKo("2", "foo")) + .withCauseExactlyInstanceOf(ValidationException.class) + .havingRootCause() + .withMessage("boom"); + assertThat(okIn.getSpy()).isEqualTo(1); + assertThat(koOut.getSpy()).isEqualTo(1); + } + + @Test + void retry() { + var okIn = SingletonClassInstanceFactory.getInstance(MyOkInputGuardrail.class); + var koOutWithRetry = SingletonClassInstanceFactory.getInstance(MyKoWithRetryOutputGuardrail.class); + assertThat(okIn.getSpy()).isEqualTo(0); + assertThat(koOutWithRetry.getSpy()).isEqualTo(0); + service.outKoWithRetry("2", "foo"); + assertThat(okIn.getSpy()).isEqualTo(1); + assertThat(koOutWithRetry.getSpy()).isEqualTo(2); + } + + @Test + void reprompt() { + var okIn = SingletonClassInstanceFactory.getInstance(MyOkInputGuardrail.class); + var koOutWithReprompt = SingletonClassInstanceFactory.getInstance(MyKoWithRepromprOutputGuardrail.class); + assertThat(okIn.getSpy()).isEqualTo(0); + assertThat(koOutWithReprompt.getSpy()).isEqualTo(0); + service.outKoWithReprompt("2", "foo"); + assertThat(okIn.getSpy()).isEqualTo(1); + assertThat(koOutWithReprompt.getSpy()).isEqualTo(2); + } + + public interface MyAiService { + @InputGuardrails(MyOkInputGuardrail.class) + @OutputGuardrails(MyOkOutputGuardrail.class) + String bothOk(@MemoryId String id, @UserMessage String message); + + @InputGuardrails(MyKoInputGuardrail.class) + @OutputGuardrails(MyOkOutputGuardrail.class) + String inKo(@MemoryId String id, @UserMessage String message); + + @InputGuardrails(MyOkInputGuardrail.class) + @OutputGuardrails(MyKoOutputGuardrail.class) + String outKo(@MemoryId String id, @UserMessage String message); + + @InputGuardrails(MyOkInputGuardrail.class) + @OutputGuardrails(MyKoWithRetryOutputGuardrail.class) + String outKoWithRetry(@MemoryId String id, @UserMessage String message); + + @InputGuardrails(MyOkInputGuardrail.class) + @OutputGuardrails(MyKoWithRepromprOutputGuardrail.class) + String outKoWithReprompt(@MemoryId String id, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class MyOkInputGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public InputGuardrailResult validate(InputGuardrailRequest params) { + spy.incrementAndGet(); + return success(); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyKoInputGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public InputGuardrailResult validate(InputGuardrailRequest params) { + spy.incrementAndGet(); + return failure("boom", new ValidationException("boom")); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyOkOutputGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + spy.incrementAndGet(); + return OutputGuardrailResult.success(); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyKoOutputGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + spy.incrementAndGet(); + return failure("boom", new ValidationException("boom")); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyKoWithRetryOutputGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + if (spy.incrementAndGet() == 1) { + return retry("KO"); + } + return success(); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyKoWithRepromprOutputGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + if (spy.incrementAndGet() == 1) { + return reprompt("KO", "retry"); + } + return success(); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class MyChatModel implements ChatModel { + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + return ChatResponse.builder().aiMessage(AiMessage.from("Hi!")).build(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java new file mode 100644 index 0000000000..f72c44472b --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java @@ -0,0 +1,136 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailException; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; +import org.junit.jupiter.api.Test; + +class InputGuardrailChainTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void guardrailChainsAreInvoked() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + aiService.firstOneTwo("1", "foo"); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(firstGuardrail.lastAccess()).isLessThan(secondGuardrail.lastAccess()); + } + + @Test + void guardrailOrderIsCorrect() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + aiService.twoAndFirst("1", "foo"); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.lastAccess()).isLessThan(firstGuardrail.lastAccess()); + } + + @Test + void failTheChain() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + var failingGuardrail = SingletonClassInstanceFactory.getInstance(FailingGuardrail.class); + + assertThatThrownBy(() -> aiService.failingFirstTwo("1", "foo")) + .isInstanceOf(InputGuardrailException.class) + .hasCauseInstanceOf(ValidationException.class) + .hasRootCauseMessage("boom"); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(0); + assertThat(failingGuardrail.spy()).isEqualTo(1); + } + + public interface MyAiService { + @InputGuardrails({FirstGuardrail.class, SecondGuardrail.class}) + String firstOneTwo(@MemoryId String mem, @UserMessage String message); + + @InputGuardrails({SecondGuardrail.class, FirstGuardrail.class}) + String twoAndFirst(@MemoryId String mem, @UserMessage String message); + + @InputGuardrails({FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class}) + String failingFirstTwo(@MemoryId String mem, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class FirstGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicLong lastAccess = new AtomicLong(); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + return success(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + } + + public static class SecondGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private volatile AtomicLong lastAccess = new AtomicLong(); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + return success(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + } + + public static class FailingGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + if (spy.incrementAndGet() == 1) { + return fatal("boom", new ValidationException("boom")); + } + return success(); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassAndMethodTest.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassAndMethodTest.java new file mode 100644 index 0000000000..eaf1bca41b --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassAndMethodTest.java @@ -0,0 +1,89 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class InputGuardrailOnClassAndMethodTest extends BaseGuardrailTests { + @ParameterizedTest + @MethodSource("services") + void guardrailsFromTheClassAreInvoked(String testDescription, MyAiService aiService) { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + + assertThat(koGuardrail.spy()).isEqualTo(0); + } + + static Stream services() { + return Stream.of( + Arguments.of("Using AiServices builder", MyAiServiceWithoutClassAnnotations.create()), + Arguments.of("Using annotation at class level", MyAiServiceUsingClassAnnotations.create())); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @InputGuardrails(OKGuardrail.class) + String hi(@MemoryId String mem); + } + + public interface MyAiServiceWithoutClassAnnotations extends MyAiService { + static MyAiService create() { + return createAiService( + MyAiServiceWithoutClassAnnotations.class, + List.of(OKGuardrail.class), + List.of(), + builder -> builder.chatModel(new MyChatModel())); + } + } + + @InputGuardrails(KOGuardrail.class) + public interface MyAiServiceUsingClassAnnotations extends MyAiService { + static MyAiService create() { + return createAiService( + MyAiServiceUsingClassAnnotations.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassTest.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassTest.java new file mode 100644 index 0000000000..e0e67d9625 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailOnClassTest.java @@ -0,0 +1,69 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class InputGuardrailOnClassTest extends BaseGuardrailTests { + @ParameterizedTest + @MethodSource("services") + void guardrailsFromTheClassAreInvoked(String testDescription, MyAiService aiService) { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + } + + static Stream services() { + return Stream.of( + Arguments.of("Using AiServices builder", MyAiService.create()), + Arguments.of("Using annotation at class level", MyAiServiceWithClassAnnotation.create())); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + String hi(@MemoryId String mem); + + static MyAiService create() { + return createAiService( + MyAiService.class, + List.of(OKGuardrail.class), + List.of(), + builder -> builder.chatModel(new MyChatModel())); + } + } + + @InputGuardrails(OKGuardrail.class) + public interface MyAiServiceWithClassAnnotation extends MyAiService { + static MyAiService create() { + return createAiService( + MyAiServiceWithClassAnnotation.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage ignored) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailPromptTemplateTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailPromptTemplateTests.java new file mode 100644 index 0000000000..4d8d14a5f9 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailPromptTemplateTests.java @@ -0,0 +1,207 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.InputGuardrailValidation; +import dev.langchain4j.memory.chat.MessageWindowChatMemory; +import dev.langchain4j.service.AiServices; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import dev.langchain4j.service.V; +import java.util.List; +import java.util.Map; +import java.util.stream.Stream; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class InputGuardrailPromptTemplateTests { + @BeforeEach + void beforeEach() { + InputGuardrailValidation.getInstance().reset(); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkNoParameters(String testDescription, Assistant aiService) { + assertThat(aiService.getJoke()).isEqualTo("Request: Tell me a joke; Response: Hi!"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me a joke"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()).isEmpty(); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithMemoryId(String testDescription, Assistant aiService) { + aiService.getAnotherJoke("memory-id-001"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me another joke"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of("arg0", "memory-id-001")); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithNoMemoryIdAndOneParameter(String testDescription, Assistant aiService) { + aiService.sayHiToMyFriendNoMemory("Rambo"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Say hi to my friend {{it}}!"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "arg0", "Rambo", + "it", "Rambo")); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithMemoryIdAndOneParameter(String testDescription, Assistant aiService) { + aiService.sayHiToMyFriend("1", "Chuck Norris"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Say hi to my friend {{friend}}!"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "friend", "Chuck Norris", + "arg0", "1")); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithNoMemoryIdAndThreeParameters(String testDescription, Assistant aiService) { + aiService.sayHiToMyFriends("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "topic1", "Chuck Norris", + "topic2", "Jean-Claude Van Damme", + "topic3", "Silvester Stallone")); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithNoMemoryIdAndList(String testDescription, Assistant aiService) { + aiService.sayHiToMyFriends(List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone")); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageText()) + .isEqualTo("Tell me something about [Chuck Norris, Jean-Claude Van Damme, Silvester Stallone]!"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{it}}!"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "arg0", + List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"), + "it", + "[Chuck Norris, Jean-Claude Van Damme, Silvester Stallone]")); + } + + @ParameterizedTest + @MethodSource("assistants") + void shouldWorkWithMemoryIdAndList(String testDescription, Assistant aiService) { + aiService.sayHiToMyFriends( + "memory-id-007", List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone")); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageText()) + .isEqualTo( + "Tell me something about [Chuck Norris, Jean-Claude Van Damme, Silvester Stallone]! This is my memory id: memory-id-007"); + assertThat(InputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{topics}}! This is my memory id: {{memoryId}}"); + assertThat(InputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "topics", + List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"), + "memoryId", + "memory-id-007")); + } + + static Stream assistants() { + return Stream.of( + Arguments.of("Assistant with class-level annotation", ClassLevelAssistant.create()), + Arguments.of("Assistant with method-level annotations", MethodLevelAssistant.create()), + Arguments.of("Assistant with builder-style guardrail", Assistant.create())); + } + + @InputGuardrails(InputGuardrailValidation.class) + interface ClassLevelAssistant extends Assistant { + static Assistant create() { + return Assistant.create(ClassLevelAssistant.class); + } + } + + interface MethodLevelAssistant extends Assistant { + @InputGuardrails(InputGuardrailValidation.class) + @UserMessage("Tell me a joke") + @Override + String getJoke(); + + @UserMessage("Tell me another joke") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String getAnotherJoke(@MemoryId String memoryId); + + @UserMessage("Say hi to my friend {{it}}!") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String sayHiToMyFriendNoMemory(String friend); + + @UserMessage("Say hi to my friend {{friend}}!") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String sayHiToMyFriend(@MemoryId String mem, @V("friend") String friend); + + @UserMessage("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String sayHiToMyFriends(@V("topic1") String topic1, @V("topic2") String topic2, @V("topic3") String topic3); + + @UserMessage("Tell me something about {{it}}!") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String sayHiToMyFriends(List topics); + + @UserMessage("Tell me something about {{topics}}! This is my memory id: {{memoryId}}") + @InputGuardrails(InputGuardrailValidation.class) + @Override + String sayHiToMyFriends(@V("memoryId") @MemoryId String memoryId, @V("topics") List topics); + + static Assistant create() { + return Assistant.create(MethodLevelAssistant.class); + } + } + + interface Assistant { + @UserMessage("Tell me a joke") + String getJoke(); + + @UserMessage("Tell me another joke") + String getAnotherJoke(@MemoryId String memoryId); + + @UserMessage("Say hi to my friend {{it}}!") + String sayHiToMyFriendNoMemory(String friend); + + @UserMessage("Say hi to my friend {{friend}}!") + String sayHiToMyFriend(@MemoryId String mem, @V("friend") String friend); + + @UserMessage("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!") + String sayHiToMyFriends(@V("topic1") String topic1, @V("topic2") String topic2, @V("topic3") String topic3); + + @UserMessage("Tell me something about {{it}}!") + String sayHiToMyFriends(List topics); + + @UserMessage("Tell me something about {{topics}}! This is my memory id: {{memoryId}}") + String sayHiToMyFriends(@V("memoryId") @MemoryId String memoryId, @V("topics") List topics); + + static T create(Class clazz) { + return AiServices.builder(clazz) + .chatModel(new MyChatModel()) + .chatMemoryProvider(memoryId -> MessageWindowChatMemory.withMaxMessages(10)) + .build(); + } + + static Assistant create() { + return AiServices.builder(Assistant.class) + .chatModel(new MyChatModel()) + .chatMemoryProvider(memoryId -> MessageWindowChatMemory.withMaxMessages(10)) + .inputGuardrails(List.of(InputGuardrailValidation.getInstance())) + .build(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailRewritingTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailRewritingTests.java new file mode 100644 index 0000000000..6b02805537 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailRewritingTests.java @@ -0,0 +1,37 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.service.UserMessage; +import dev.langchain4j.service.V; +import org.junit.jupiter.api.Test; + +class InputGuardrailRewritingTests extends BaseGuardrailTests { + @Test + void rewriting() { + assertThat(MyAiService.create().test("first prompt", "second prompt")) + .hasSize(MessageTruncatingGuardrail.MAX_LENGTH); + } + + public interface MyAiService { + @UserMessage("Given {{first}} and {{second}} do something") + @InputGuardrails(MessageTruncatingGuardrail.class) + String test(@V("first") String first, @V("second") String second); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new EchoChatModel())); + } + } + + public static class MessageTruncatingGuardrail implements InputGuardrail { + static final int MAX_LENGTH = 20; + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + String text = um.singleText(); + return successWith(text.substring(0, MAX_LENGTH)); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailTests.java new file mode 100644 index 0000000000..ce42c9834d --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailTests.java @@ -0,0 +1,78 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class InputGuardrailTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void guardrailsAreInvoked() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + } + + @Test + void guardrailCanThrowValidationException() { + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(koGuardrail.spy()).isEqualTo(0); + assertThatThrownBy(() -> aiService.ko("1")).hasCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(1); + assertThatThrownBy(() -> aiService.ko("1")).hasCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(2); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @InputGuardrails(OKGuardrail.class) + String hi(@MemoryId String mem); + + @UserMessage("Say Hi!") + @InputGuardrails(KOGuardrail.class) + String ko(@MemoryId String mem); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return failure("KO", new ValidationException("KO")); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailValidationTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailValidationTests.java new file mode 100644 index 0000000000..69a40a08b2 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailValidationTests.java @@ -0,0 +1,234 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.data.Index.atIndex; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailException; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.StreamingChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.TokenStream; +import dev.langchain4j.service.UserMessage; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class InputGuardrailValidationTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void ok() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + aiService.ok("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.ok("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + } + + @Test + void ko() { + assertThatThrownBy(() -> aiService.ko("2")) + .isInstanceOf(InputGuardrailException.class) + .hasMessageContaining("KO"); + } + + @Test + void okStreaming() throws InterruptedException { + var latch = new CountDownLatch(1); + var text = new AtomicReference(); + var partialResponses = new ArrayList(); + + aiService + .okStream("1") + .onError(t -> latch.countDown()) + .onPartialResponse(partialResponses::add) + .onCompleteResponse(response -> { + text.set(response.aiMessage().text()); + latch.countDown(); + }) + .start(); + + latch.await(10, TimeUnit.SECONDS); + + assertThat(String.join(" ", text.get())).isEqualTo("Streaming hi !"); + assertThat(String.join(" ", partialResponses)).isEqualTo(text.get()); + + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(1); + } + + @Test + void koStreaming() { + assertThatExceptionOfType(InputGuardrailException.class) + .isThrownBy(() -> aiService.koStream("2")) + .withMessageContaining("KO"); + + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(koGuardrail.spy()).isEqualTo(1); + } + + @Test + void fatalException() { + assertThatExceptionOfType(InputGuardrailException.class) + .isThrownBy(() -> aiService.fatal("5")) + .withMessageContaining("Fatal"); + + var fatal = SingletonClassInstanceFactory.getInstance(KOFatalGuardrail.class); + assertThat(fatal.spy()).isEqualTo(1); + } + + @Test + void memoryCheck() { + var memoryCheck = SingletonClassInstanceFactory.getInstance(MemoryCheck.class); + aiService.test("1", "foo"); + assertThat(memoryCheck.spy()).isEqualTo(1); + + aiService.test("1", "bar"); + assertThat(memoryCheck.spy()).isEqualTo(2); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @InputGuardrails(OKGuardrail.class) + String ok(@MemoryId String mem); + + @UserMessage("Say Hi!") + @InputGuardrails(KOGuardrail.class) + String ko(@MemoryId String mem); + + @UserMessage("Say Hi!") + @InputGuardrails(OKGuardrail.class) + TokenStream okStream(@MemoryId String mem); + + @UserMessage("Say Hi!") + @InputGuardrails(KOGuardrail.class) + TokenStream koStream(@MemoryId String mem); + + @UserMessage("Say Hi!") + @InputGuardrails(KOFatalGuardrail.class) + String fatal(@MemoryId String mem); + + @InputGuardrails(MemoryCheck.class) + String test(@MemoryId String name, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel()) + .streamingChatModel(new MyStreamingChatModel())); + } + } + + public static class OKGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOFatalGuardrail implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(dev.langchain4j.data.message.UserMessage um) { + spy.incrementAndGet(); + throw new IllegalArgumentException("Fatal"); + } + + public int spy() { + return spy.get(); + } + } + + public static class MemoryCheck implements InputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(InputGuardrailRequest params) { + spy.incrementAndGet(); + var messages = Optional.ofNullable(params.requestParams().chatMemory()) + .map(ChatMemory::messages) + .orElseGet(List::of); + + if (messages.isEmpty()) { + assertThat(params.userMessage().singleText()).isEqualTo("foo"); + } + + if (messages.size() == 2) { + assertThat(messages) + .satisfies( + message -> assertThat(message) + .isInstanceOf(dev.langchain4j.data.message.UserMessage.class) + .extracting(m -> ((dev.langchain4j.data.message.UserMessage) m).singleText()) + .isEqualTo("foo"), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isInstanceOf(AiMessage.class) + .extracting(m -> ((AiMessage) m).text()) + .isEqualTo("Hi!"), + atIndex(1)); + + assertThat(params.userMessage().singleText()).isEqualTo("bar"); + } + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class MyChatModel implements ChatModel { + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + return ChatResponse.builder().aiMessage(AiMessage.from("Hi!")).build(); + } + } + + public static class MyStreamingChatModel implements StreamingChatModel { + @Override + public void doChat(ChatRequest chatRequest, StreamingChatResponseHandler handler) { + handler.onPartialResponse("Streaming hi"); + handler.onPartialResponse("!"); + handler.onCompleteResponse(ChatResponse.builder() + .aiMessage(AiMessage.from("Streaming hi !")) + .build()); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/MyChatModel.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/MyChatModel.java new file mode 100644 index 0000000000..75459b9979 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/MyChatModel.java @@ -0,0 +1,24 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.ChatMessageType; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; + +public class MyChatModel implements ChatModel { + private static String getUserMessage(ChatRequest chatRequest) { + return chatRequest.messages().stream() + .filter(message -> message.type() == ChatMessageType.USER) + .findFirst() + .map(chatMessage -> ((dev.langchain4j.data.message.UserMessage) chatMessage).singleText()) + .orElseThrow(() -> new IllegalArgumentException("No user message found")); + } + + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + return ChatResponse.builder() + .aiMessage(AiMessage.from("Request: %s; Response: Hi!".formatted(getUserMessage(chatRequest)))) + .build(); + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainOnStreamedResponseTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainOnStreamedResponseTests.java new file mode 100644 index 0000000000..c9d2dfb9dc --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainOnStreamedResponseTests.java @@ -0,0 +1,283 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.atIndex; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.data.message.ChatMessageType; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.StreamingChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.TokenStream; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class OutputGuardrailChainOnStreamedResponseTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void guardrailChainsAreInvoked() throws InterruptedException { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + + var value = execute(() -> aiService.firstOneTwo("1", "foo")); + assertThat(value).isEqualTo("Hi! World! "); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(firstGuardrail.lastAccess()).isLessThan(secondGuardrail.lastAccess()); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void guardrailOrderIsCorrect() throws InterruptedException { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + + var value = execute(() -> aiService.twoAndFirst("1", "foo")); + assertThat(value).isEqualTo("Hi! World! "); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.lastAccess()).isLessThan(firstGuardrail.lastAccess()); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void retryRestartsTheChain() throws InterruptedException { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + var failingGuardrail = SingletonClassInstanceFactory.getInstance(FailingGuardrail.class); + var value = execute(() -> aiService.failingFirstTwo("1", "foo")); + assertThat(value).isEqualTo("Hi! World! "); + assertThat(firstGuardrail.spy()).isEqualTo(2); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(firstGuardrail.lastAccess()).isLessThan(secondGuardrail.lastAccess()); + assertThat(failingGuardrail.spy()).isEqualTo(2); + assertThat(failingGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + public interface MyAiService { + @OutputGuardrails({FirstGuardrail.class, SecondGuardrail.class}) + TokenStream firstOneTwo(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({SecondGuardrail.class, FirstGuardrail.class}) + TokenStream twoAndFirst(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class}) + TokenStream failingFirstTwo(@MemoryId String mem, @UserMessage String message); + + static MyAiService create() { + return createAiService( + MyAiService.class, builder -> builder.streamingChatModel(new MyStreamingChatModel())); + } + } + + public static class FirstGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private volatile AtomicLong lastAccess = new AtomicLong(); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + return success(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class SecondGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private volatile AtomicLong lastAccess = new AtomicLong(); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + return success(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class FailingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + if (spy.incrementAndGet() == 1) { + return reprompt("Retry", "Retry"); + } + return success(); + } + + public int spy() { + return spy.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class MyStreamingChatModel implements StreamingChatModel { + @Override + public void doChat(ChatRequest chatRequest, StreamingChatResponseHandler handler) { + handler.onPartialResponse("Hi!"); + handler.onPartialResponse(" "); + handler.onPartialResponse("World!"); + handler.onCompleteResponse(ChatResponse.builder() + .aiMessage(AiMessage.from("Hi! World! ")) + .build()); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainTests.java new file mode 100644 index 0000000000..5697b9989b --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailChainTests.java @@ -0,0 +1,535 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import static org.assertj.core.api.Assertions.atIndex; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.data.message.ChatMessageType; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class OutputGuardrailChainTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void guardrailChainsAreInvoked() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + aiService.firstOneTwo("1", "foo"); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(firstGuardrail.lastAccess()).isLessThan(secondGuardrail.lastAccess()); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void guardrailOrderIsCorrect() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + aiService.twoAndFirst("1", "foo"); + assertThat(firstGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(secondGuardrail.lastAccess()).isLessThan(firstGuardrail.lastAccess()); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void retryRepromptRestartsTheChain() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondGuardrail.class); + var failingGuardrail = SingletonClassInstanceFactory.getInstance(FailingGuardrail.class); + aiService.failingFirstTwo("1", "foo"); + assertThat(firstGuardrail.spy()).isEqualTo(3); + assertThat(secondGuardrail.spy()).isEqualTo(1); + assertThat(firstGuardrail.lastAccess()).isLessThan(secondGuardrail.lastAccess()); + assertThat(failingGuardrail.spy()).isEqualTo(3); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(failingGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void rewritesTheOutputTwiceInTheChain() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstRewritingGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(SecondRewritingGuardrail.class); + assertThat(aiService.rewritingSuccess("1", "foo")).isEqualTo("Request: foo; Response: Hi!,1,2"); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void repromptAfterRewriteIsNotAllowed() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstRewritingGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(RepromptingGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> aiService.repromptAfterRewrite("1", "foo")) + .withMessageContaining("Retry or reprompt is not allowed after a rewritten output"); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void rewritesTheOutputWithAResult() { + var firstGuardrail = SingletonClassInstanceFactory.getInstance(FirstRewritingGuardrail.class); + var secondGuardrail = SingletonClassInstanceFactory.getInstance(RewritingGuardrailWithResult.class); + assertThat(aiService.rewritingSuccessWithResult("1", "foo")).isSameAs(RewritingGuardrailWithResult.RESULT); + assertThat(firstGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + assertThat(secondGuardrail.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + public interface MyAiService { + @OutputGuardrails({FirstGuardrail.class, SecondGuardrail.class}) + String firstOneTwo(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({SecondGuardrail.class, FirstGuardrail.class}) + String twoAndFirst(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails( + value = {FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class}, + maxRetries = 3) + String failingFirstTwo(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({FirstRewritingGuardrail.class, SecondRewritingGuardrail.class}) + String rewritingSuccess(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({FirstRewritingGuardrail.class, RepromptingGuardrail.class}) + String repromptAfterRewrite(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails({FirstRewritingGuardrail.class, RewritingGuardrailWithResult.class}) + String rewritingSuccessWithResult(@MemoryId String mem, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class FirstGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private volatile AtomicLong lastAccess = new AtomicLong(); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + return success(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + } + + public static class SecondGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private volatile AtomicLong lastAccess = new AtomicLong(); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + lastAccess.set(System.nanoTime()); + + try { + TimeUnit.MILLISECONDS.sleep(1); + } catch (InterruptedException e) { + // Ignore me + } + + return success(); + } + + public int spy() { + return spy.get(); + } + + public long lastAccess() { + return lastAccess.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class FailingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + int v = spy.incrementAndGet(); + + if (v == 1) { + return reprompt("Reprompt", "Reprompt"); + } else if (v == 2) { + return retry("Retry"); + } + + return success(); + } + + public int spy() { + return spy.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class FirstRewritingGuardrail implements OutputGuardrail { + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + String text = responseFromLLM.text(); + return successWith(text + ",1"); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class SecondRewritingGuardrail implements OutputGuardrail { + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + var text = responseFromLLM.text(); + return successWith(text + ",2"); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class RewritingGuardrailWithResult implements OutputGuardrail { + static final String RESULT = String.valueOf(1_000); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + var text = responseFromLLM.text(); + return successWith(text + ",2", RESULT); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class RepromptingGuardrail implements OutputGuardrail { + private boolean firstCall = true; + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + if (firstCall) { + firstCall = false; + String text = responseFromLLM.text(); + return reprompt("Wrong message", text + ", " + text); + } + + return success(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassAndMethodTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassAndMethodTests.java new file mode 100644 index 0000000000..023b4a3ae5 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassAndMethodTests.java @@ -0,0 +1,69 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class OutputGuardrailOnClassAndMethodTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void properGuardrailsAreInvoked() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + assertThat(koGuardrail.spy()).isEqualTo(0); + } + + @OutputGuardrails(KOGuardrail.class) + public interface MyAiService { + + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + String hi(@MemoryId String mem); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassTests.java new file mode 100644 index 0000000000..ede9aea46d --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnClassTests.java @@ -0,0 +1,102 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class OutputGuardrailOnClassTests extends BaseGuardrailTests { + @ParameterizedTest + @MethodSource("services") + void correctGuardrailsAreInvoked(String testDescription, MyAiService aiService) { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + assertThat(koGuardrail.spy()).isEqualTo(0); + } + + static Stream services() { + return Stream.of( + Arguments.of("AiService using builder", MyAiService.create()), + Arguments.of("AiService using annotation at class level", MyAiServiceWithClassAnnotation.create()), + Arguments.of( + "AiService using annotation at class and method level", + MyAiServiceWithClassAndMethodAnnotation.create())); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + String hi(@MemoryId String mem); + + static MyAiService create() { + return createAiService( + MyAiService.class, + List.of(), + List.of(OKGuardrail.class), + builder -> builder.chatModel(new MyChatModel())); + } + } + + @OutputGuardrails(OKGuardrail.class) + public interface MyAiServiceWithClassAnnotation extends MyAiService { + static MyAiService create() { + return createAiService( + MyAiServiceWithClassAnnotation.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + @OutputGuardrails(KOGuardrail.class) + public interface MyAiServiceWithClassAndMethodAnnotation extends MyAiService { + @Override + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + String hi(@MemoryId String mem); + + static MyAiService create() { + return createAiService( + MyAiServiceWithClassAndMethodAnnotation.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseTests.java new file mode 100644 index 0000000000..f91c55ec49 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseTests.java @@ -0,0 +1,159 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import static org.assertj.core.api.Assertions.fail; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.model.chat.StreamingChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.TokenStream; +import dev.langchain4j.service.UserMessage; +import java.util.ArrayList; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class OutputGuardrailOnStreamedResponseTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void outputGuardrailsAreInvoked() throws InterruptedException { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + execute(() -> aiService.ok("1")); + assertThat(okGuardrail.spy()).isEqualTo(1); + execute(() -> aiService.ok("2")); + assertThat(okGuardrail.spy()).isEqualTo(2); + } + + @Test + void guardrailCanThrowValidationException() { + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(koGuardrail.spy()).isEqualTo(0); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> execute(() -> aiService.ko("1"))) + .withCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(1); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> execute(() -> aiService.ko("1"))) + .withCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(2); + } + + @Test + void streamingPartialResponsesBuffered() { + var okPartials = new ArrayList(); + var okComplete = new AtomicReference(); + + this.aiService + .ok("1") + .onPartialResponse(okPartials::add) + .onError(t -> fail(t.getMessage())) + .onCompleteResponse(okComplete::set) + .start(); + + assertThat(okPartials).hasSize(3).containsExactly("Hi!", " ", "World!"); + assertThat(okComplete.get()) + .isNotNull() + .extracting(m -> m.aiMessage().text()) + .isEqualTo("Hi! World! "); + } + + @Test + void streamingPartialResponsesNotShownOnError() { + var repromptPartials = new ArrayList(); + var repromptComplete = new AtomicReference(); + + assertThatExceptionOfType(OutputGuardrailException.class).isThrownBy(() -> this.aiService + .reprompt("1") + .onPartialResponse(repromptPartials::add) + .onError(t -> fail(t.getMessage())) + .onCompleteResponse(repromptComplete::set) + .start()); + + assertThat(repromptPartials).isEmpty(); + assertThat(repromptComplete.get()).isNull(); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + TokenStream ok(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOGuardrail.class) + TokenStream ko(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(RepromptGuardrail.class) + TokenStream reprompt(@MemoryId String mem); + + static MyAiService create() { + return createAiService( + MyAiService.class, builder -> builder.streamingChatModel(new MyStreamingChatModel())); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class RepromptGuardrail implements OutputGuardrail { + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + return reprompt("reorompt", "reprompt"); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + if (responseFromLLM.text().length() > 3) { // Accumulated response. + return failure("KO", new ValidationException("KO")); + } else { // Chunk, do not fail on the first chunk + if (responseFromLLM.text().contains("Hi!")) { + return success(); + } else { + return failure("KO", new ValidationException("KO")); + } + } + } + + public int spy() { + return spy.get(); + } + } + + public static class MyStreamingChatModel implements StreamingChatModel { + @Override + public void doChat(ChatRequest chatRequest, StreamingChatResponseHandler handler) { + handler.onPartialResponse("Hi!"); + handler.onPartialResponse(" "); + handler.onPartialResponse("World!"); + handler.onCompleteResponse(ChatResponse.builder() + .aiMessage(AiMessage.from("Hi! World! ")) + .build()); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseValidationTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseValidationTests.java new file mode 100644 index 0000000000..387cda4c6b --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailOnStreamedResponseValidationTests.java @@ -0,0 +1,199 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.model.chat.StreamingChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.TokenStream; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class OutputGuardrailOnStreamedResponseValidationTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void ok() throws InterruptedException { + assertThat(execute(() -> aiService.ok("1"))).isEqualTo("Hi! World! "); + } + + @Test + void ko() { + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> execute(() -> aiService.ko("2"))) + .withMessageContaining("KO"); + } + + @Test + void retryOk() throws InterruptedException { + var retry = SingletonClassInstanceFactory.getInstance(RetryingGuardrail.class); + + assertThat(execute(() -> aiService.retry("3"))).isEqualTo("Hi! World! "); + assertThat(retry.spy()).isEqualTo(2); + } + + @Test + void fetryFail() { + var retryFail = SingletonClassInstanceFactory.getInstance(RetryingButFailGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> execute(() -> aiService.retryButFail("4"))) + .withMessageContaining("maximum number of retries"); + assertThat(retryFail.spy()).isEqualTo(3); + } + + @Test + void fatalException() { + var fatal = SingletonClassInstanceFactory.getInstance(KOFatalGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> execute(() -> aiService.fatal("5"))) + .withMessageContaining("Fatal"); + assertThat(fatal.spy()).isEqualTo(1); + } + + @Test + void rewritingWhileStreaming() throws InterruptedException { + var rewriting = SingletonClassInstanceFactory.getInstance(RewritingGuardrail.class); + assertThat(execute(() -> aiService.rewriting("1"))).isEqualTo("Hi! World! ,1"); + assertThat(rewriting.spy()).isEqualTo(1); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + TokenStream ok(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOGuardrail.class) + TokenStream ko(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(value = RetryingGuardrail.class, maxRetries = 3) + TokenStream retry(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(value = RetryingButFailGuardrail.class, maxRetries = 3) + TokenStream retryButFail(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOFatalGuardrail.class) + TokenStream fatal(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails({RewritingGuardrail.class}) + TokenStream rewriting(@MemoryId String mem); + + static MyAiService create() { + return createAiService( + MyAiService.class, builder -> builder.streamingChatModel(new MyStreamingChatModel())); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class RetryingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + int v = spy.incrementAndGet(); + if (v >= 2) { + return OutputGuardrailResult.success(); + } + return retry("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class RetryingButFailGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return retry("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOFatalGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + throw new IllegalArgumentException("Fatal"); + } + + public int spy() { + return spy.get(); + } + } + + public static class RewritingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + String text = responseFromLLM.text(); + return successWith(text + ",1"); + } + + public int spy() { + return spy.get(); + } + } + + public static class MyStreamingChatModel implements StreamingChatModel { + @Override + public void doChat(ChatRequest chatRequest, StreamingChatResponseHandler handler) { + handler.onPartialResponse("Hi!"); + handler.onPartialResponse(" "); + handler.onPartialResponse("World!"); + handler.onCompleteResponse(ChatResponse.builder() + .aiMessage(AiMessage.from("Hi! World! ")) + .build()); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailPromptTemplateTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailPromptTemplateTests.java new file mode 100644 index 0000000000..c0f652be3a --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailPromptTemplateTests.java @@ -0,0 +1,197 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.OutputGuardrailValidation; +import dev.langchain4j.memory.chat.MessageWindowChatMemory; +import dev.langchain4j.service.AiServices; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import dev.langchain4j.service.V; +import java.util.List; +import java.util.Map; +import java.util.stream.Stream; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class OutputGuardrailPromptTemplateTests extends BaseGuardrailTests { + @BeforeEach + void beforeEach() { + OutputGuardrailValidation.getInstance().reset(); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkNoParameters(String testDescription, MyAiService aiService) { + aiService.getJoke(); + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me a joke"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()).isEmpty(); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithMemoryId(String testDescription, MyAiService aiService) { + aiService.getAnotherJoke("memory-id-001"); + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me another joke"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of("arg0", "memory-id-001")); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithNoMemoryIdAndOneParameter(String testDescription, MyAiService aiService) { + aiService.sayHiToMyFriendNoMemory("Rambo"); + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Say hi to my friend {{it}}!"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "arg0", "Rambo", + "it", "Rambo")); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithMemoryIdAndOneParameter(String testDescription, MyAiService aiService) { + aiService.sayHiToMyFriend("1", "Chuck Norris"); + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Say hi to my friend {{friend}}!"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "friend", "Chuck Norris", + "arg0", "1")); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithNoMemoryIdAndThreeParameters(String testDescription, MyAiService aiService) { + aiService.sayHiToMyFriends("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"); + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "topic1", "Chuck Norris", + "topic2", "Jean-Claude Van Damme", + "topic3", "Silvester Stallone")); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithNoMemoryIdAndList(String testDescription, MyAiService aiService) { + aiService.sayHiToMyFriends(List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone")); + + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{it}}!"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "arg0", + List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"), + "it", + "[Chuck Norris, Jean-Claude Van Damme, Silvester Stallone]")); + } + + @ParameterizedTest + @MethodSource("services") + void shouldWorkWithMemoryIdAndList(String testDescription, MyAiService aiService) { + aiService.sayHiToMyFriends( + "memory-id-007", List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone")); + + assertThat(OutputGuardrailValidation.getInstance().spyUserMessageTemplate()) + .isEqualTo("Tell me something about {{topics}}! This is my memory id: {{memoryId}}"); + assertThat(OutputGuardrailValidation.getInstance().spyVariables()) + .containsExactlyInAnyOrderEntriesOf(Map.of( + "topics", + List.of("Chuck Norris", "Jean-Claude Van Damme", "Silvester Stallone"), + "memoryId", + "memory-id-007")); + } + + static Stream services() { + return Stream.of( + Arguments.of("AiService with class-level annotation", ClassLevelAiService.create()), + Arguments.of("AiService with method-level annotations", MethodLevelAiService.create()), + Arguments.of("AiService with builder-style guardrails", MyAiService.create())); + } + + @OutputGuardrails(OutputGuardrailValidation.class) + public interface ClassLevelAiService extends MyAiService { + static MyAiService create() { + return createAiService(ClassLevelAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public interface MethodLevelAiService extends MyAiService { + @Override + @UserMessage("Tell me a joke") + @OutputGuardrails(OutputGuardrailValidation.class) + String getJoke(); + + @Override + @UserMessage("Tell me another joke") + @OutputGuardrails(OutputGuardrailValidation.class) + String getAnotherJoke(@MemoryId String memoryId); + + @Override + @UserMessage("Say hi to my friend {{it}}!") + @OutputGuardrails(OutputGuardrailValidation.class) + String sayHiToMyFriendNoMemory(String friend); + + @Override + @UserMessage("Say hi to my friend {{friend}}!") + @OutputGuardrails(OutputGuardrailValidation.class) + String sayHiToMyFriend(@MemoryId String mem, @V("friend") String friend); + + @Override + @UserMessage("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!") + @OutputGuardrails(OutputGuardrailValidation.class) + String sayHiToMyFriends(@V("topic1") String topic1, @V("topic2") String topic2, @V("topic3") String topic3); + + @Override + @UserMessage("Tell me something about {{it}}!") + @OutputGuardrails(OutputGuardrailValidation.class) + String sayHiToMyFriends(List topics); + + @Override + @UserMessage("Tell me something about {{topics}}! This is my memory id: {{memoryId}}") + @OutputGuardrails(OutputGuardrailValidation.class) + String sayHiToMyFriends(@V("memoryId") @MemoryId String memoryId, @V("topics") List topics); + + static MyAiService create() { + return createAiService(MethodLevelAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public interface MyAiService { + @UserMessage("Tell me a joke") + String getJoke(); + + @UserMessage("Tell me another joke") + String getAnotherJoke(@MemoryId String memoryId); + + @UserMessage("Say hi to my friend {{it}}!") + String sayHiToMyFriendNoMemory(String friend); + + @UserMessage("Say hi to my friend {{friend}}!") + String sayHiToMyFriend(@MemoryId String mem, @V("friend") String friend); + + @UserMessage("Tell me something about {{topic1}}, {{topic2}}, {{topic3}}!") + String sayHiToMyFriends(@V("topic1") String topic1, @V("topic2") String topic2, @V("topic3") String topic3); + + @UserMessage("Tell me something about {{it}}!") + String sayHiToMyFriends(List topics); + + @UserMessage("Tell me something about {{topics}}! This is my memory id: {{memoryId}}") + String sayHiToMyFriends(@V("memoryId") @MemoryId String memoryId, @V("topics") List topics); + + static MyAiService create() { + return AiServices.builder(MyAiService.class) + .chatModel(new MyChatModel()) + .chatMemoryProvider(memoryId -> MessageWindowChatMemory.withMaxMessages(10)) + .outputGuardrails(OutputGuardrailValidation.getInstance()) + .build(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingRetryTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingRetryTests.java new file mode 100644 index 0000000000..f44f318363 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingRetryTests.java @@ -0,0 +1,124 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.SystemMessage; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class OutputGuardrailRepromptingRetryTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void ok() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OkGuardrail.class); + aiService.ok("1", "foo"); + assertThat(okGuardrail.getSpy()).isEqualTo(1); + + aiService.ok("1", "bar"); + assertThat(okGuardrail.getSpy()).isEqualTo(2); + } + + @Test + void retryFailing() { + var retryGuardrail = SingletonClassInstanceFactory.getInstance(RetryGuardrail.class); + assertThatThrownBy(() -> aiService.retry("1", "foo")) + .isInstanceOf(OutputGuardrailException.class) + .hasMessageContaining("maximum number of retries"); + assertThat(retryGuardrail.getSpy()).isEqualTo(5); + } + + @Test + void noRetry() { + var retryGuardrail = SingletonClassInstanceFactory.getInstance(RetryGuardrail.class); + retryGuardrail.reset(); + assertThatThrownBy(() -> aiService.noRetry("2", "foo")) + .isInstanceOf(OutputGuardrailException.class) + .hasMessageContaining("maximum number of retries"); + assertThat(retryGuardrail.getSpy()).isEqualTo(1); + } + + @Test + void repromptingFailing() { + var repromptingGuardrail = SingletonClassInstanceFactory.getInstance(RepromptingGuardrail.class); + assertThatThrownBy(() -> aiService.reprompting("1", "foo")) + .isInstanceOf(OutputGuardrailException.class) + .hasMessageContaining("maximum number of retries"); + assertThat(repromptingGuardrail.getSpy()) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + } + + @SystemMessage("Say Hi!") + public interface MyAiService { + @OutputGuardrails(OkGuardrail.class) + String ok(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails(value = RetryGuardrail.class, maxRetries = 5) + String retry(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails(value = RetryGuardrail.class, maxRetries = 0) + String noRetry(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails(RepromptingGuardrail.class) + String reprompting(@MemoryId String mem, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OkGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int getSpy() { + return spy.get(); + } + } + + public static class RetryGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + spy.incrementAndGet(); + return retry("Retry"); + } + + public int getSpy() { + return spy.get(); + } + + public void reset() { + this.spy.set(0); + } + } + + public static class RepromptingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + int v = spy.incrementAndGet(); + return reprompt("Retry", "reprompt"); + } + + public int getSpy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingTests.java new file mode 100644 index 0000000000..466ef16a86 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailRepromptingTests.java @@ -0,0 +1,264 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import static org.assertj.core.api.Assertions.atIndex; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.data.message.ChatMessageType; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.SystemMessage; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class OutputGuardrailRepromptingTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void repromptingOkAfterOneRetry() { + var repromptingOne = SingletonClassInstanceFactory.getInstance(RepromptingOne.class); + aiService.one("1", "foo"); + assertThat(repromptingOne.getSpy()).isEqualTo(2); + assertThat(repromptingOne.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void repromptingOkAfterTwoRetries() { + var repromptingTwo = SingletonClassInstanceFactory.getInstance(RepromptingTwo.class); + aiService.two("2", "foo"); + assertThat(repromptingTwo.getSpy()).isEqualTo(3); + assertThat(repromptingTwo.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @Test + void repromptingFailing() { + var repromptingFailed = SingletonClassInstanceFactory.getInstance(RepromptingFailed.class); + assertThatExceptionOfType(OutputGuardrailException.class).isThrownBy(() -> aiService.fail("3", "foo")); + assertThat(repromptingFailed.getSpy()).isEqualTo(3); + assertThat(repromptingFailed.chatMemory()) + .isNotNull() + .extracting(ChatMemory::messages) + .satisfies(messages -> assertThat(messages) + .isNotNull() + .hasSize(2) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.USER), + atIndex(0)) + .satisfies( + message -> assertThat(message) + .isNotNull() + .extracting(ChatMessage::type) + .isEqualTo(ChatMessageType.AI), + atIndex(1))); + } + + @SystemMessage("Say Hi!") + public interface MyAiService { + @OutputGuardrails(RepromptingOne.class) + String one(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails(value = RepromptingTwo.class, maxRetries = 3) + String two(@MemoryId String mem, @UserMessage String message); + + @OutputGuardrails(value = RepromptingFailed.class, maxRetries = 3) + String fail(@MemoryId String mem, @UserMessage String message); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class RepromptingOne implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + return OutputGuardrail.super.validate(params); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + if (spy.incrementAndGet() == 1) { + return reprompt("Retry", "Retry"); + } + + return success(); + } + + public int getSpy() { + return spy.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class RepromptingTwo implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + int v = spy.incrementAndGet(); + var messages = params.requestParams().chatMemory().messages(); + + if (v == 1) { + ChatMessage last = messages.get(messages.size() - 1); + assertThat(last).isInstanceOfSatisfying(AiMessage.class, am -> assertThat(am.text()) + .isEqualTo("Nope")); + assertThat(params.responseFromLLM().aiMessage().text()).isEqualTo("Nope"); + return reprompt("Retry", "Retry"); + } + + if (v == 2) { + // Check that it's NOT in memory + ChatMessage last = messages.get(messages.size() - 1); + ChatMessage beforeLast = messages.get(messages.size() - 2); + + assertThat(last).isInstanceOfSatisfying(AiMessage.class, am -> assertThat(am.text()) + .isEqualTo("Nope")); + assertThat(params.responseFromLLM().aiMessage().text()).isEqualTo("Hello"); + assertThat(beforeLast) + .isInstanceOfSatisfying( + dev.langchain4j.data.message.UserMessage.class, + um -> assertThat(um.singleText()).isEqualTo("foo")); + + return reprompt("Retry", "Retry"); + } + + if (v != 3) { + throw new IllegalArgumentException("Unexpected call"); + } + + return success(); + } + + public int getSpy() { + return spy.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class RepromptingFailed implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + private final AtomicReference chatMemory = new AtomicReference<>(); + + @Override + public OutputGuardrailResult validate(OutputGuardrailRequest params) { + this.chatMemory.set(params.requestParams().chatMemory()); + int v = spy.incrementAndGet(); + var messages = params.requestParams().chatMemory().messages(); + + if (v == 1) { + ChatMessage last = messages.get(messages.size() - 1); + assertThat(last).isInstanceOfSatisfying(AiMessage.class, am -> assertThat(am.text()) + .isEqualTo("Nope")); + return reprompt("Retry", "Retry Once"); + } + + if (v == 2) { + // Check that it's NOT in memory + ChatMessage last = messages.get(messages.size() - 1); + ChatMessage beforeLast = messages.get(messages.size() - 2); + + assertThat(last).isInstanceOfSatisfying(AiMessage.class, am -> assertThat(am.text()) + .isEqualTo("Nope")); + assertThat(beforeLast) + .isInstanceOfSatisfying( + dev.langchain4j.data.message.UserMessage.class, + um -> assertThat(um.singleText()).isEqualTo("foo")); + return reprompt("Retry", "Retry Twice"); + } + + return reprompt("Retry", "Retry Again"); + } + + public int getSpy() { + return spy.get(); + } + + public ChatMemory chatMemory() { + return chatMemory.get(); + } + } + + public static class MyChatModel implements ChatModel { + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + var messages = chatRequest.messages(); + var last = messages.get(messages.size() - 1); + + if (last instanceof dev.langchain4j.data.message.UserMessage um) { + if ("foo".equals(um.singleText())) { + return ChatResponse.builder() + .aiMessage(AiMessage.from("Nope")) + .build(); + } + + if (um.singleText().contains("Retry")) { + return ChatResponse.builder() + .aiMessage(AiMessage.from("Hello")) + .build(); + } + } + + throw new IllegalArgumentException("Unexpected message: " + messages); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailTests.java new file mode 100644 index 0000000000..f061832ab4 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailTests.java @@ -0,0 +1,79 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class OutputGuardrailTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void outputGuardrailsAreInvoked() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + assertThat(okGuardrail.spy()).isEqualTo(0); + aiService.hi("1"); + assertThat(okGuardrail.spy()).isEqualTo(1); + aiService.hi("2"); + assertThat(okGuardrail.spy()).isEqualTo(2); + } + + @Test + void guardrailCanThrowValidationException() { + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThat(koGuardrail.spy()).isEqualTo(0); + assertThatThrownBy(() -> aiService.ko("1")).hasCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(1); + assertThatThrownBy(() -> aiService.ko("1")).hasCauseExactlyInstanceOf(ValidationException.class); + assertThat(koGuardrail.spy()).isEqualTo(2); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + String hi(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOGuardrail.class) + String ko(@MemoryId String mem); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return failure("KO", new ValidationException("KO")); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailValidationTests.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailValidationTests.java new file mode 100644 index 0000000000..87dd0de73a --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/OutputGuardrailValidationTests.java @@ -0,0 +1,159 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import com.example.SingletonClassInstanceFactory; +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailException; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.service.MemoryId; +import dev.langchain4j.service.UserMessage; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; + +class OutputGuardrailValidationTests extends BaseGuardrailTests { + MyAiService aiService = MyAiService.create(); + + @Test + void ok() { + var okGuardrail = SingletonClassInstanceFactory.getInstance(OKGuardrail.class); + aiService.ok("1"); + assertThat(okGuardrail.spy()).isOne(); + } + + @Test + void ko() { + var koGuardrail = SingletonClassInstanceFactory.getInstance(KOGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> aiService.ko("2")) + .withMessageContaining("KO"); + + assertThat(koGuardrail.spy()).isOne(); + } + + @Test + void retryOk() { + var retry = SingletonClassInstanceFactory.getInstance(RetryingGuardrail.class); + aiService.retry("3"); + assertThat(retry.spy()).isEqualTo(2); + } + + @Test + void retryFail() { + var retryFail = SingletonClassInstanceFactory.getInstance(RetryingButFailGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> aiService.retryButFail("4")) + .withMessageContaining("maximum number of retries"); + assertThat(retryFail.spy()).isEqualTo(3); + } + + @Test + void fatalException() { + var fatal = SingletonClassInstanceFactory.getInstance(KOFatalGuardrail.class); + assertThatExceptionOfType(OutputGuardrailException.class) + .isThrownBy(() -> aiService.fatal("5")) + .withMessageContaining("Fatal"); + assertThat(fatal.spy()).isOne(); + } + + public interface MyAiService { + @UserMessage("Say Hi!") + @OutputGuardrails(OKGuardrail.class) + String ok(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOGuardrail.class) + String ko(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(RetryingGuardrail.class) + String retry(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(value = RetryingButFailGuardrail.class, maxRetries = 3) + String retryButFail(@MemoryId String mem); + + @UserMessage("Say Hi!") + @OutputGuardrails(KOFatalGuardrail.class) + String fatal(@MemoryId String mem); + + static MyAiService create() { + return createAiService(MyAiService.class, builder -> builder.chatModel(new MyChatModel())); + } + } + + public static class OKGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return success(); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + return failure("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class RetryingGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + int v = spy.incrementAndGet(); + if (v == 2) { + return OutputGuardrailResult.success(); + } + return retry("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class RetryingButFailGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + int v = spy.incrementAndGet(); + return retry("KO"); + } + + public int spy() { + return spy.get(); + } + } + + public static class KOFatalGuardrail implements OutputGuardrail { + private final AtomicInteger spy = new AtomicInteger(0); + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + spy.incrementAndGet(); + throw new IllegalArgumentException("Fatal"); + } + + public int spy() { + return spy.get(); + } + } +} diff --git a/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java new file mode 100644 index 0000000000..9b5ec3d529 --- /dev/null +++ b/integration-tests/integration-tests-guardrails/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java @@ -0,0 +1,7 @@ +package dev.langchain4j.service.guardrail; + +public class ValidationException extends RuntimeException { + public ValidationException(String message) { + super(message); + } +} diff --git a/integration-tests/pom.xml b/integration-tests/pom.xml new file mode 100644 index 0000000000..25b39085ef --- /dev/null +++ b/integration-tests/pom.xml @@ -0,0 +1,43 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-parent + 1.1.0-beta7-SNAPSHOT + ../langchain4j-parent/pom.xml + + + langchain4j-integration-tests-parent + pom + LangChain4j :: Integration Tests + Parent POM for all integration tests + + + + Apache-2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + A business-friendly OSS license + + + + + integration-tests-class-instance-loader + integration-tests-class-instance-loader/integration-tests-class-instance-loader-spring + integration-tests-class-instance-loader/integration-tests-class-instance-loader-quarkus + integration-tests-class-metadata-provider + integration-tests-class-metadata-provider/integration-tests-class-metadata-provider-spring + integration-tests-guardrails + + + + + true + true + true + true + requireUpperBoundDeps + + diff --git a/langchain4j-bom/pom.xml b/langchain4j-bom/pom.xml index c99908fb4e..a8a35552f3 100644 --- a/langchain4j-bom/pom.xml +++ b/langchain4j-bom/pom.xml @@ -31,6 +31,12 @@ 1.1.0-SNAPSHOT + + dev.langchain4j + langchain4j-test + 1.1.0-beta7-SNAPSHOT + + dev.langchain4j langchain4j-http-client @@ -564,10 +570,10 @@ flatten - process-resources flatten + process-resources flatten.clean diff --git a/langchain4j-core/pom.xml b/langchain4j-core/pom.xml index 77c61bfd59..37e38c12ca 100644 --- a/langchain4j-core/pom.xml +++ b/langchain4j-core/pom.xml @@ -42,7 +42,6 @@ ${jspecify.version} - @@ -91,6 +90,30 @@ + + + + + + + org.codehaus.mojo + build-helper-maven-plugin + 3.6.0 + + + + add-test-source + + generate-test-sources + + + ${project.basedir}/../langchain4j-test/src/main/java + + + + + + org.apache.maven.plugins maven-source-plugin @@ -143,9 +166,12 @@ + dev.langchain4j.classinstance dev.langchain4j.data.document dev.langchain4j.data.text dev.langchain4j.exception + dev.langchain4j.guardrail + dev.langchain4j.guardrail.config dev.langchain4j.internal dev.langchain4j.model.chat dev.langchain4j.model.chat.listener diff --git a/langchain4j-core/src/main/java/dev/langchain4j/Experimental.java b/langchain4j-core/src/main/java/dev/langchain4j/Experimental.java index 35fbb892d4..13f44a36d6 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/Experimental.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/Experimental.java @@ -1,14 +1,28 @@ package dev.langchain4j; -import java.lang.annotation.Target; - import static java.lang.annotation.ElementType.CONSTRUCTOR; import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PACKAGE; import static java.lang.annotation.ElementType.TYPE; +import java.lang.annotation.Documented; +import java.lang.annotation.Inherited; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + /** * Indicates that a class/constructor/method is experimental and might change in the future. */ -@Target({TYPE, CONSTRUCTOR, METHOD}) +@Inherited +@Documented +@Retention(RetentionPolicy.RUNTIME) +@Target({TYPE, CONSTRUCTOR, METHOD, PACKAGE}) public @interface Experimental { + /** + * Describes why the annotated element is experimental + * + * @return The experimental description + */ + String value() default "This feature is experimental and the API is subject to change"; } diff --git a/langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java b/langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java new file mode 100644 index 0000000000..b74f30a251 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/classinstance/ClassInstanceLoader.java @@ -0,0 +1,45 @@ +package dev.langchain4j.classinstance; + +import dev.langchain4j.spi.classloading.ClassInstanceFactory; +import java.lang.reflect.InvocationTargetException; +import java.util.ServiceLoader; + +/** + * Utility class for creating and retrieving instances of specified class types. + * This class provides a mechanism to delegate instance creation to a factory, if available, + * or fallback to direct instantiation using the no-argument constructor. + *

+ * This is useful in scenarios where dependency injection frameworks or other managed + * object factories might need to be leveraged for object creation. + *

+ */ +public final class ClassInstanceLoader { + private ClassInstanceLoader() {} + + /** + * Retrieves an instance of the specified class type. This method first attempts to obtain the instance + * through a {@link ClassInstanceFactory}, if available. If no factory is present, it creates a new + * instance of the class using its no-argument constructor. + * + * @param the type of the class + * @param clazz the class object representing the type whose instance is to be created + * @return an instance of the specified class type + */ + public static T getClassInstance(Class clazz) { + return ServiceLoader.load(ClassInstanceFactory.class) + .findFirst() + .map(classInstanceFactory -> classInstanceFactory.getInstanceOfClass(clazz)) + .orElseGet(() -> createNewClassInstance(clazz)); + } + + private static T createNewClassInstance(Class clazz) { + try { + return clazz.getDeclaredConstructor().newInstance(); + } catch (InstantiationException + | IllegalAccessException + | InvocationTargetException + | NoSuchMethodException e) { + throw new RuntimeException(e); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/AbstractGuardrailExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/AbstractGuardrailExecutor.java new file mode 100644 index 0000000000..27f1b96e5d --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/AbstractGuardrailExecutor.java @@ -0,0 +1,262 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.Internal; +import dev.langchain4j.guardrail.GuardrailResult.Failure; +import dev.langchain4j.guardrail.config.GuardrailsConfig; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; + +/** + * Abstract base class for {@link GuardrailExecutor}s. + * @param + * The type of {@link GuardrailsConfig} to use for configuration + * @param

+ * The type of {@link GuardrailRequest} to validate + * @param + * The type of {@link GuardrailResult} to return + * @param + * The type of {@link Guardrail}s being executed + * @param + * The type of {@link Failure} to return + */ +@Internal +public abstract sealed class AbstractGuardrailExecutor< + C extends GuardrailsConfig, + P extends GuardrailRequest

, + R extends GuardrailResult, + G extends Guardrail, + F extends Failure> + implements GuardrailExecutor permits InputGuardrailExecutor, OutputGuardrailExecutor { + + private final C config; + private final List guardrails; + + protected AbstractGuardrailExecutor(C config, List guardrails) { + ensureNotNull(config, "config"); + this.config = config; + this.guardrails = Optional.ofNullable(guardrails).orElseGet(List::of); + } + + /** + * Creates a failure result from some {@link Failure}s. + * @param failures The failures + * @return A {@link GuardrailResult} containing the failures + */ + protected abstract R createFailure(List failures); + + /** + * Creates a success result. + * @return A {@link GuardrailResult} representing success + */ + protected abstract R createSuccess(); + + /** + * Creates a {@link GuardrailException} using the provided message and optional cause. + * + * @param message The detailed message for the exception. + * @param cause The underlying cause of the exception, or null if no cause is available. + * @return A new instance of {@link GuardrailException} constructed with the provided message and cause. + */ + protected abstract GuardrailException createGuardrailException(String message, Throwable cause); + + @Override + public C config() { + return this.config; + } + + @Override + public List guardrails() { + return this.guardrails; + } + + /** + * Validates a guardrail against a set of params. + *

+ * If any kind of {@link Exception} is thrown during validation, it will be wrapped in a {@link GuardrailException}. + *

+ * @param params The {@link GuardrailRequest} to validate + * @param guardrail The {@link Guardrail} to evaluate against + * @throws GuardrailException If any kind of {@link Exception} is thrown during validation + * @return The {@link GuardrailResult} of the validation + */ + protected R validate(P params, G guardrail) { + ensureNotNull(params, "params"); + ensureNotNull(guardrail, "guardrail"); + + try { + return guardrail.validate(params).validatedBy(guardrail.getClass()); + } catch (Exception e) { + throw createGuardrailException(e.getMessage(), e); + } + } + + /** + * Handles a fatal result. + * @param accumulatedResult The accumulated result + * @param result The fatal result + * @return The fatal result, possibly wrapped/modified in some way + */ + protected R handleFatalResult(R accumulatedResult, R result) { + return result; + } + + protected R executeGuardrails(P params) { + ensureNotNull(params, "params"); + + var accumulatedResult = createSuccess(); + var accumulatedParams = params; + + for (var guardrail : this.guardrails) { + if (guardrail != null) { + var result = validate(accumulatedParams, guardrail); + + if (result.isFatal()) { + // Fatal result, so stop right here and don't do any more processing + return handleFatalResult(accumulatedResult, result); + } + + if (result.hasRewrittenResult()) { + accumulatedParams = accumulatedParams.withText(result.successfulText()); + } + + accumulatedResult = composeResult(accumulatedResult, result); + } + } + + return accumulatedResult; + } + + protected R composeResult(R oldResult, R newResult) { + if (oldResult.isSuccess()) { + return newResult; + } + + if (newResult.isSuccess()) { + return oldResult; + } + + var failures = new ArrayList(oldResult.failures()); + failures.addAll(newResult.failures()); + + return createFailure(failures); + } + + /** + * A generic abstract builder class for creating instances of {@link GuardrailExecutor}. + * + * @param + * The type of {@link GuardrailsConfig} to use for configuration + * @param

+ * The type of {@link GuardrailRequest} to validate + * @param + * The type of {@link GuardrailResult} to return + * @param + * The type of {@link Guardrail}s being executed + * + * This class is sealed to restrict subclassing to only specific permitted classes, such as + * {@link InputGuardrailExecutor.InputGuardrailExecutorBuilder} and + * {@link OutputGuardrailExecutor.OutputGuardrailExecutorBuilder}. + * + * It provides methods to configure and manage the guardrails and their associated configurations, + * eventually culminating in the construction of a specific {@link GuardrailExecutor}. + */ + public abstract static sealed class GuardrailExecutorBuilder< + C extends GuardrailsConfig, + R extends GuardrailResult, + P extends GuardrailRequest

, + G extends Guardrail, + B extends GuardrailExecutorBuilder> + permits InputGuardrailExecutor.InputGuardrailExecutorBuilder, + OutputGuardrailExecutor.OutputGuardrailExecutorBuilder { + + private final C defaultConfig; + private C config; + private List guardrails = new ArrayList<>(); + + protected GuardrailExecutorBuilder(C defaultConfig) { + this.defaultConfig = ensureNotNull(defaultConfig, "defaultConfig"); + } + + /** + * Constructs and returns an instance of {@link GuardrailExecutor}. + * + * This method finalizes the building process, using the configuration and guardrails + * provided, to create a fully-formed {@link GuardrailExecutor} instance. The returned + * instance enables execution of guardrails on given parameters. + * + * @return A fully initialized instance of {@link GuardrailExecutor}, ready to validate + * interactions based on the configured guardrails and parameters. + */ + public abstract GuardrailExecutor build(); + + /** + * Retrieves the current configuration instance used by this builder. + * + * @return The configuration set in the builder. + */ + protected C config() { + return (this.config != null) ? this.config : this.defaultConfig; + } + + /** + * Retrieves the list of guardrails configured in the builder. + * Guardrails are validation rules applied to interactions with the model, ensuring that inputs or outputs + * meet required conditions for safety and correctness. + * + * @return A list containing the configured guardrails. + */ + protected List guardrails() { + return this.guardrails; + } + + /** + * Sets the configuration for the guardrail executor builder. + * + * @param config The configuration instance to be set, which implements {@link GuardrailsConfig}. + * This can be null if no specific configuration is required. + * @return The updated instance of the builder, allowing for method chaining. + */ + public B config(C config) { + this.config = config; + return (B) this; + } + + /** + * Updates the list of guardrails for the builder. The provided guardrails will replace + * the current list of guardrails in the builder. If the provided list is null, all + * existing guardrails will be cleared. + * + * @param guardrails A list of guardrails to be set for the builder. It can be null, + * in which case the current list of guardrails will be cleared. + * @return The updated instance of the builder, allowing for method chaining. + */ + public B guardrails(List guardrails) { + this.guardrails.clear(); + + if (guardrails != null) { + this.guardrails.addAll(guardrails); + } + + return (B) this; + } + + /** + * Updates the builder with the specified guardrails. This method accepts + * a variadic array of guardrails, which will be used to replace the current + * set of guardrails in the builder. If the input is null, the existing + * guardrails will remain unchanged. + * + * @param guardrails An optional array of guardrails to be set for the builder. + * Null values are accepted and will not clear existing guardrails. + * @return The updated instance of the builder, allowing for method chaining. + */ + public B guardrails(G... guardrails) { + Optional.ofNullable(guardrails).map(List::of).ifPresent(this::guardrails); + + return (B) this; + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/Guardrail.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/Guardrail.java new file mode 100644 index 0000000000..da6b03e97f --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/Guardrail.java @@ -0,0 +1,22 @@ +package dev.langchain4j.guardrail; + +/** + * A guardrail is a rule that is applied when interacting with an LLM either to the input (the user message) or to the + * output of the model to ensure that they are safe and meet the expectations of the model. + * + * @param

+ * The type of the {@link GuardrailRequest} + * @param + * The type of the {@link GuardrailResult} + */ +public interface Guardrail

> { + /** + * Validate the interaction between the model and the user in one of the two directions. + * + * @param params + * The parameters of the request or the response to be validated + * + * @return The result of the validation + */ + R validate(P params); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailException.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailException.java new file mode 100644 index 0000000000..75c153b778 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailException.java @@ -0,0 +1,22 @@ +package dev.langchain4j.guardrail; + +import dev.langchain4j.exception.LangChain4jException; + +/** + * Exception thrown when an input or output guardrail validation fails. + *

+ * This class is not intended to be used within guardrail implementations. It is for the framework only. + *

+ * @see InputGuardrailException + * @see OutputGuardrailException + */ +public sealed class GuardrailException extends LangChain4jException + permits InputGuardrailException, OutputGuardrailException { + protected GuardrailException(String message) { + super(message); + } + + protected GuardrailException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailExecutor.java new file mode 100644 index 0000000000..60db9ecd40 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailExecutor.java @@ -0,0 +1,45 @@ +package dev.langchain4j.guardrail; + +import dev.langchain4j.guardrail.config.GuardrailsConfig; +import java.util.List; + +/** + * Represents a mechanism to execute a set of guardrails on given parameters. + * This interface defines the contract for validating interactions (input or output) + * using multiple guardrails. + * + * @param + * The type of {@link GuardrailsConfig} to use for configuration + * @param

+ * The type of {@link GuardrailRequest} to validate + * @param + * The type of {@link GuardrailResult} to return + * @param + * The type of {@link Guardrail}s being executed + */ +public sealed interface GuardrailExecutor< + C extends GuardrailsConfig, + P extends GuardrailRequest, + R extends GuardrailResult, + G extends Guardrail> + permits AbstractGuardrailExecutor { + + /** + * The {@link GuardrailsConfig} to use for configuration of the guardrail execution + * @return The {@link GuardrailsConfig} to use for configuration of the guardrail execution + */ + C config(); + + /** + * Retrieves the guardrails associated with the implementation. + * @return The guardrails which can be used for validating inputs or outputs against predefined rules. + */ + List guardrails(); + + /** + * Executes the provided guardrails on the given parameters. + * @param params The {@link GuardrailRequest} to validate + * @return The {@link GuardrailResult} of the validation + */ + R execute(P params); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequest.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequest.java new file mode 100644 index 0000000000..f0d959e673 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequest.java @@ -0,0 +1,27 @@ +package dev.langchain4j.guardrail; + +/** + * Represents the parameter passed to {@link Guardrail#validate(GuardrailRequest)}} in order to validate an interaction + * between a user and the LLM. + */ +public sealed interface GuardrailRequest

> + permits InputGuardrailRequest, OutputGuardrailRequest { + + /** + * Retrieves the common parameters that are shared across guardrail checks. + * + * @return an instance of {@code GuardrailRequestParams} containing shared parameters such as chat memory, + * user message template, and additional variables. + */ + GuardrailRequestParams requestParams(); + + /** + * Recreate this guardrail param with the given input or output text. + * + * @param text + * The text of the rewritten param. + * + * @return A clone of this guardrail params with the given input or output text. + */ + P withText(String text); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequestParams.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequestParams.java new file mode 100644 index 0000000000..bc301852d9 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailRequestParams.java @@ -0,0 +1,134 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.rag.AugmentationResult; +import java.util.Map; + +/** + * Represents the common parameters shared across guardrail checks when validating interactions + * between a user and a language model. This class encapsulates the chat memory, user message + * template, and additional variables required for guardrail processing. + */ +public final class GuardrailRequestParams { + private final ChatMemory chatMemory; + private final AugmentationResult augmentationResult; + private final String userMessageTemplate; + private final Map variables; + + private GuardrailRequestParams(Builder builder) { + this.chatMemory = builder.chatMemory; + this.augmentationResult = builder.augmentationResult; + this.userMessageTemplate = ensureNotNull(builder.userMessageTemplate, "userMessageTemplate"); + this.variables = ensureNotNull(builder.variables, "variables"); + } + + /** + * Returns the chat memory. + * + * @return the chat memory, may be null + */ + public ChatMemory chatMemory() { + return chatMemory; + } + + /** + * Returns the augmentation result. + * + * @return the augmentation result, may be null + */ + public AugmentationResult augmentationResult() { + return augmentationResult; + } + + /** + * Returns the user message template. + * + * @return the user message template, never null + */ + public String userMessageTemplate() { + return userMessageTemplate; + } + + /** + * Returns the variables. + * + * @return the variables, never null + */ + public Map variables() { + return variables; + } + + /** + * Creates a new builder for {@link GuardrailRequestParams}. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Builder for {@link GuardrailRequestParams}. + */ + public static class Builder { + private ChatMemory chatMemory; + private AugmentationResult augmentationResult; + private String userMessageTemplate; + private Map variables; + + /** + * Sets the chat memory. + * + * @param chatMemory the chat memory + * @return this builder + */ + public Builder chatMemory(ChatMemory chatMemory) { + this.chatMemory = chatMemory; + return this; + } + + /** + * Sets the augmentation result. + * + * @param augmentationResult the augmentation result + * @return this builder + */ + public Builder augmentationResult(AugmentationResult augmentationResult) { + this.augmentationResult = augmentationResult; + return this; + } + + /** + * Sets the user message template. + * + * @param userMessageTemplate the user message template + * @return this builder + */ + public Builder userMessageTemplate(String userMessageTemplate) { + this.userMessageTemplate = userMessageTemplate; + return this; + } + + /** + * Sets the variables. + * + * @param variables the variables + * @return this builder + */ + public Builder variables(Map variables) { + this.variables = variables; + return this; + } + + /** + * Builds a new {@link GuardrailRequestParams}. + * + * @return a new {@link GuardrailRequestParams} + */ + public GuardrailRequestParams build() { + return new GuardrailRequestParams(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailResult.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailResult.java new file mode 100644 index 0000000000..3f5126568c --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/GuardrailResult.java @@ -0,0 +1,155 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import java.util.stream.Collectors; + +/** + * The result of the validation of an interaction between a user and the LLM. + * + * @param + * The type of guardrail result to expect + * + * @see InputGuardrailResult + * @see OutputGuardrailResult + */ +public sealed interface GuardrailResult> + permits InputGuardrailResult, OutputGuardrailResult { + /** + * The possible results of a guardrails validation. + */ + enum Result { + /** + * A successful validation. + */ + SUCCESS, + /** + * A successful validation with a specific result. + */ + SUCCESS_WITH_RESULT, + /** + * A failed validation not preventing the subsequent validations eventually registered to be evaluated. + */ + FAILURE, + /** + * A fatal failed validation, blocking the evaluation of any other validations eventually registered. + */ + FATAL + } + + /** + * The message and the cause of the failure of a single validation. + */ + sealed interface Failure permits InputGuardrailResult.Failure, OutputGuardrailResult.Failure { + /** + * Build a failure from a specific {@link Guardrail} class + */ + Failure withGuardrailClass(Class guardrailClass); + + /** + * The failure message + */ + String message(); + + /** + * The cause of the failure + */ + Throwable cause(); + + /** + * The {@link Guardrail} class + */ + Class guardrailClass(); + + /** + * The string representation of the failure + * @return A string representation of the failure + */ + default String asString() { + var guardrailName = + Optional.ofNullable(guardrailClass()).map(Class::getName).orElse(""); + + return "The guardrail %s failed with this message: %s".formatted(guardrailName, message()); + } + } + + /** + * The result of the guardrail + */ + Result result(); + + /** + * @return The list of failures eventually resulting from a set of validations. + */ + List failures(); + + /** + * The message of the successful result + */ + String successfulText(); + + /** + * Whether or not the result is successful, but the result was re-written, potentially due to re-prompting + */ + default boolean hasRewrittenResult() { + return result() == Result.SUCCESS_WITH_RESULT; + } + + /** + * Whether or not the result is considered fatal + */ + default boolean isFatal() { + return result() == Result.FATAL; + } + + /** + * Whether or not the result is considered successful + */ + default boolean isSuccess() { + var result = result(); + return (result == Result.SUCCESS) || (result == Result.SUCCESS_WITH_RESULT); + } + + /** + * Gets the exception from the first failure + */ + default Throwable getFirstFailureException() { + return !isSuccess() + ? failures().stream() + .map(Failure::cause) + .filter(Objects::nonNull) + .findFirst() + .orElse(null) + : null; + } + + /** + * The {@link Guardrail} class which performed this validation + */ + default GR validatedBy(Class guardrailClass) { + ensureNotNull(guardrailClass, "guardrailClass"); + + if (!isSuccess()) { + var failures = failures(); + + if (failures.size() != 1) { + throw new IllegalArgumentException(); + } + + failures.set(0, failures.get(0).withGuardrailClass(guardrailClass)); + } + + return (GR) this; + } + + default String asString() { + if (isSuccess()) { + return hasRewrittenResult() ? "Success with '%s'".formatted(successfulText()) : "Success"; + } + + return failures().stream().map(Failure::toString).collect(Collectors.joining(", ")); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java new file mode 100644 index 0000000000..2cf8d98eb2 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrail.java @@ -0,0 +1,120 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.InputGuardrailResult.Failure; + +/** + * An input guardrail is a rule that is applied to the input of the model to ensure that the input (i.e. the user + * message and parameters) is safe and meets the expectations of the model. + *

+ * Input guardrails are either successful or failed. A successful guardrail means that the input is valid and can be sent to + * the model. A failed guardrail means that the input is invalid and cannot be sent to the model. + *

+ *

+ * A failed guardrail will stop further processing of any other input guardrails. + *

+ */ +public interface InputGuardrail extends Guardrail { + /** + * Validates the {@code user message} that will be sent to the LLM. + *

+ * + * @param userMessage + * the response from the LLM + */ + default InputGuardrailResult validate(UserMessage userMessage) { + return failure("Validation not implemented"); + } + + /** + * Validates the input that will be sent to the LLM. + *

+ * Unlike {@link #validate(UserMessage)}, this method allows to access the memory and the augmentation result (in + * the case of a RAG). + *

+ * Implementation must not attempt to write to the memory or the augmentation result. + * + * @param params + * the parameters, including the user message, the memory, and the augmentation result. + */ + @Override + default InputGuardrailResult validate(InputGuardrailRequest params) { + ensureNotNull(params, "params"); + return validate(params.userMessage()); + } + + /** + * Produces a successful result without any successful text + * + * @return The result of a successful input guardrail validation. + */ + default InputGuardrailResult success() { + return InputGuardrailResult.success(); + } + + /** + * Produces a successful result with specific success text + * + * @return The result of a successful input guardrail validation with a specific text. + * + * @param successfulText + * The text of the successful result. + */ + default InputGuardrailResult successWith(String successfulText) { + return InputGuardrailResult.successWith(successfulText); + } + + /** + * Produces a non-fatal failure + * + * @param message + * A message describing the failure. + * + * @return The result of a failed input guardrail validation. + */ + default InputGuardrailResult failure(String message) { + return new InputGuardrailResult(new Failure(message), false); + } + + /** + * Produces a non-fatal failure + * + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * + * @return The result of a failed input guardrail validation. + */ + default InputGuardrailResult failure(String message, Throwable cause) { + return new InputGuardrailResult(new Failure(message, cause), false); + } + + /** + * Produces a fatal failure + * + * @param message + * A message describing the failure. + * + * @return The result of a failed input guardrail validation. + */ + default InputGuardrailResult fatal(String message) { + return new InputGuardrailResult(new Failure(message), true); + } + + /** + * Produces a non-fatal failure + * + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * + * @return The result of a failed input guardrail validation. + */ + default InputGuardrailResult fatal(String message, Throwable cause) { + return new InputGuardrailResult(new Failure(message, cause), true); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailException.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailException.java new file mode 100644 index 0000000000..0df2282241 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailException.java @@ -0,0 +1,17 @@ +package dev.langchain4j.guardrail; + +/** + * Exception thrown when an input guardrail validation fails. + *

+ * This class is not intended to be thrown within guardrail implementations. It is for the framework only. It is ok to catch it. + *

+ */ +public final class InputGuardrailException extends GuardrailException { + public InputGuardrailException(String message) { + super(message); + } + + public InputGuardrailException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailExecutor.java new file mode 100644 index 0000000000..1717b31042 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailExecutor.java @@ -0,0 +1,101 @@ +package dev.langchain4j.guardrail; + +import dev.langchain4j.guardrail.InputGuardrailResult.Failure; +import dev.langchain4j.guardrail.config.InputGuardrailsConfig; +import java.util.List; + +/** + * The {@link GuardrailExecutor} for {@link InputGuardrail}s. + */ +public non-sealed class InputGuardrailExecutor + extends AbstractGuardrailExecutor< + InputGuardrailsConfig, InputGuardrailRequest, InputGuardrailResult, InputGuardrail, Failure> { + + protected InputGuardrailExecutor(InputGuardrailsConfig config, List guardrails) { + super(config, guardrails); + } + + /** + * Creates a failure result from some {@link Failure}s. + * @param failures The failures + * @return A {@link InputGuardrailResult} containing the failures + */ + @Override + protected InputGuardrailResult createFailure(List failures) { + return new InputGuardrailResult(failures, false); + } + + /** + * Creates a success result. + * @return A {@link InputGuardrailResult} representing success + */ + @Override + protected InputGuardrailResult createSuccess() { + return InputGuardrailResult.success(); + } + + @Override + protected InputGuardrailException createGuardrailException(String message, Throwable cause) { + return new InputGuardrailException(message, cause); + } + + /** + * Execeutes the {@link InputGuardrail}s on the given {@link InputGuardrailRequest}. + * + * @param params The {@link InputGuardrailRequest} to validate + * @return The {@link InputGuardrailResult} of the validation + */ + @Override + public InputGuardrailResult execute(InputGuardrailRequest params) { + var result = executeGuardrails(params); + + if (!result.isSuccess()) { + throw new InputGuardrailException(result.toString(), result.getFirstFailureException()); + } + + return result; + } + + /** + * Creates and returns a new builder for {@link InputGuardrailExecutor}. + * + * This builder allows for constructing and configuring an {@link InputGuardrailExecutor} + * instance, enabling customization of parameters such as the configuration and input guardrails. + * + * @return An {@link InputGuardrailExecutorBuilder} used to create {@link InputGuardrailExecutor} instances + */ + public static InputGuardrailExecutorBuilder builder() { + return new InputGuardrailExecutorBuilder(); + } + + /** + * Builder class for constructing instances of {@link InputGuardrailExecutor}. + * + * This builder allows configuration of an {@link InputGuardrailExecutor} by specifying the associated configuration + * type ({@link InputGuardrailsConfig}) and the input guardrails to be executed. + * + * Extends {@link GuardrailExecutorBuilder} for the specific types: + * - Configuration type: {@link InputGuardrailsConfig} + * - Result type: {@link InputGuardrailResult} + * - Parameter type: {@link InputGuardrailRequest} + * - Guardrail type: {@link InputGuardrail} + * + * Provides the {@code build()} method to create an {@link InputGuardrailExecutor} instance. + */ + public static non-sealed class InputGuardrailExecutorBuilder + extends GuardrailExecutorBuilder< + InputGuardrailsConfig, + InputGuardrailResult, + InputGuardrailRequest, + InputGuardrail, + InputGuardrailExecutorBuilder> { + public InputGuardrailExecutorBuilder() { + super(InputGuardrailsConfig.builder().build()); + } + + @Override + public InputGuardrailExecutor build() { + return new InputGuardrailExecutor(config(), guardrails()); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailRequest.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailRequest.java new file mode 100644 index 0000000000..ccb839b7e4 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailRequest.java @@ -0,0 +1,110 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.ContentType; +import dev.langchain4j.data.message.TextContent; +import dev.langchain4j.data.message.UserMessage; +import java.util.Objects; + +/** + * Represents the parameter passed to {@link InputGuardrail#validate(InputGuardrailRequest)}. + */ +public final class InputGuardrailRequest implements GuardrailRequest { + private final UserMessage userMessage; + private final GuardrailRequestParams commonParams; + + private InputGuardrailRequest(Builder builder) { + this.userMessage = ensureNotNull(builder.userMessage, "userMessage"); + this.commonParams = ensureNotNull(builder.commonParams, "requestParams"); + } + + /** + * Returns the user message. + * + * @return the user message + */ + public UserMessage userMessage() { + return userMessage; + } + + /** + * Returns the common parameters shared between types of guardrails. + * + * @return the common parameters + */ + @Override + public GuardrailRequestParams requestParams() { + return commonParams; + } + + @Override + public InputGuardrailRequest withText(String text) { + return new Builder() + .userMessage(rewriteUserMessage(text)) + .commonParams(this.commonParams) + .build(); + } + + public UserMessage rewriteUserMessage(String text) { + if (Objects.isNull(this.userMessage) || Objects.isNull(text)) { + return this.userMessage; + } + + var rewrittenContent = this.userMessage.contents().stream() + .map(c -> (c.type() == ContentType.TEXT) ? new TextContent(text) : c) + .toList(); + + return Objects.nonNull(this.userMessage.name()) + ? UserMessage.from(this.userMessage.name(), rewrittenContent) + : UserMessage.from(rewrittenContent); + } + + /** + * Creates a new builder for {@link InputGuardrailRequest}. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Builder for {@link InputGuardrailRequest}. + */ + public static class Builder { + private UserMessage userMessage; + private GuardrailRequestParams commonParams; + + /** + * Sets the user message. + * + * @param userMessage the user message + * @return this builder + */ + public Builder userMessage(UserMessage userMessage) { + this.userMessage = userMessage; + return this; + } + + /** + * Sets the common parameters. + * + * @param commonParams the common parameters + * @return this builder + */ + public Builder commonParams(GuardrailRequestParams commonParams) { + this.commonParams = commonParams; + return this; + } + + /** + * Builds a new {@link InputGuardrailRequest}. + * + * @return a new {@link InputGuardrailRequest} + */ + public InputGuardrailRequest build() { + return new InputGuardrailRequest(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailResult.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailResult.java new file mode 100644 index 0000000000..c40877897e --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/InputGuardrailResult.java @@ -0,0 +1,164 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.UserMessage; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Objects; +import java.util.Optional; + +/** + * The result of the validation of an {@link InputGuardrail} + */ +public final class InputGuardrailResult implements GuardrailResult { + private static final InputGuardrailResult SUCCESS = new InputGuardrailResult(); + + private final Result result; + private final String successfulText; + private final List failures; + + private InputGuardrailResult(Result result, String successfulText, List failures) { + this.result = ensureNotNull(result, "result"); + this.successfulText = successfulText; + this.failures = Optional.ofNullable(failures).orElseGet(List::of); + } + + private InputGuardrailResult() { + this(Result.SUCCESS, null, Collections.emptyList()); + } + + InputGuardrailResult(List failures, boolean fatal) { + this(fatal ? Result.FATAL : Result.FAILURE, null, failures); + } + + InputGuardrailResult(Failure failure, boolean fatal) { + this(new ArrayList<>(List.of(failure)), fatal); + } + + private InputGuardrailResult(String successfulText) { + this(Result.SUCCESS_WITH_RESULT, successfulText, Collections.emptyList()); + } + + /** + * Gets a successful input guardrail result + */ + public static InputGuardrailResult success() { + return SUCCESS; + } + + /** + * Produces a successful result with specific success text + * + * @return The result of a successful input guardrail validation with a specific text. + * + * @param successfulText + * The text of the successful result. + */ + public static InputGuardrailResult successWith(String successfulText) { + return (successfulText == null) ? success() : new InputGuardrailResult(successfulText); + } + + @Override + public Result result() { + return result; + } + + @Override + public String successfulText() { + return successfulText; + } + + @Override + @SuppressWarnings("unchecked") + public List failures() { + return (List) failures; + } + + @Override + public String toString() { + return asString(); + } + + /** + * Gets the {@link UserMessage} computed from the combination of the original {@link UserMessage} in the {@link InputGuardrailRequest} + * and this result + * @param params The input guardrail params + * @return A {@link UserMessage} computed from the combination of the original {@link UserMessage} in the {@link InputGuardrailRequest} + * * and this result + */ + public UserMessage userMessage(InputGuardrailRequest params) { + return hasRewrittenResult() ? params.rewriteUserMessage(successfulText()) : params.userMessage(); + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (o == null || getClass() != o.getClass()) return false; + InputGuardrailResult that = (InputGuardrailResult) o; + return result == that.result + && Objects.equals(successfulText, that.successfulText) + && Objects.equals(failures, that.failures); + } + + @Override + public int hashCode() { + return Objects.hash(result, successfulText, failures); + } + + /** + * Represents an input guardrail failure + */ + public static final class Failure implements GuardrailResult.Failure { + private final String message; + private final Throwable cause; + private final Class guardrailClass; + + Failure(String message, Throwable cause, Class guardrailClass) { + this.message = ensureNotNull(message, "message"); + this.cause = cause; + this.guardrailClass = guardrailClass; + } + + Failure(String message) { + this(message, null, null); + } + + Failure(String message, Throwable cause) { + this(message, cause, null); + } + + /** + * Adds a guardrail class name to a failure + * + * @param guardrailClass + * The guardrail class + */ + @Override + public Failure withGuardrailClass(Class guardrailClass) { + ensureNotNull(guardrailClass, "guardrailClass"); + return new Failure(this.message, this.cause, guardrailClass); + } + + @Override + public String message() { + return message; + } + + @Override + public Throwable cause() { + return cause; + } + + @Override + public Class guardrailClass() { + return guardrailClass; + } + + @Override + public String toString() { + return asString(); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrail.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrail.java new file mode 100644 index 0000000000..852bd278ce --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrail.java @@ -0,0 +1,139 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.core.type.TypeReference; +import com.fasterxml.jackson.databind.ObjectMapper; +import dev.langchain4j.data.message.AiMessage; +import java.util.Optional; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * An {@link OutputGuardrail} that will check whether or not a response can be successfully deserialized to an object + * of type {@code T} from JSON + *

+ * If deserialization fails, the LLM will be reprompted with {@link #getInvalidJsonReprompt(AiMessage, String)}, which + * defaults to {@link #DEFAULT_REPROMPT_PROMPT}. + *

+ * + * @param The type of object that the class should deserialize from JSON + */ +public class JsonExtractorOutputGuardrail implements OutputGuardrail { + /** + * The default message to use when reprompting + */ + public static final String DEFAULT_REPROMPT_MESSAGE = "Invalid JSON"; + + /** + * The default prompt to append to the LLM during a reprompt + */ + public static final String DEFAULT_REPROMPT_PROMPT = + "Make sure you return a valid JSON object following the specified format"; + + private static final Logger LOGGER = LoggerFactory.getLogger(JsonExtractorOutputGuardrail.class); + private final ObjectMapper objectMapper; + private Class outputClass; + private TypeReference outputType; + + public JsonExtractorOutputGuardrail(ObjectMapper objectMapper, Class outputClass) { + this.objectMapper = ensureNotNull(objectMapper, "objectMapper"); + this.outputClass = ensureNotNull(outputClass, "outputClass"); + } + + public JsonExtractorOutputGuardrail(ObjectMapper objectMapper, TypeReference outputType) { + this.objectMapper = ensureNotNull(objectMapper, "objectMapper"); + this.outputType = ensureNotNull(outputType, "outputType"); + } + + public JsonExtractorOutputGuardrail(Class outputClass) { + this(new ObjectMapper(), outputClass); + } + + public JsonExtractorOutputGuardrail(TypeReference outputType) { + this(new ObjectMapper(), outputType); + } + + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + var llmResponse = ensureNotNull(responseFromLLM, "responseFromLLM").text(); + LOGGER.debug("LLM output: {}", llmResponse); + + return deserialize(llmResponse).map(r -> successWith(llmResponse, r)).orElseGet(() -> { + LOGGER.debug("LLM output contained invalid JSON. Attempting to trim non-JSON"); + var json = trimNonJson(llmResponse); + + LOGGER.debug("Attempting to deserialize trimmed JSON: {}", json); + return deserialize(json) + .map(r -> successWith(json, r)) + .orElseGet(() -> invokeInvalidJson(responseFromLLM, json)); + }); + } + + protected String trimNonJson(String llmResponse) { + var jsonMapStart = llmResponse.indexOf('{'); + var jsonListStart = llmResponse.indexOf('['); + + if ((jsonMapStart < 0) && (jsonListStart < 0)) { + return ""; + } + + var isJsonMap = (jsonMapStart >= 0) && ((jsonMapStart < jsonListStart) || (jsonListStart < 0)); + var jsonStart = isJsonMap ? jsonMapStart : jsonListStart; + var jsonEnd = isJsonMap ? llmResponse.lastIndexOf('}') : llmResponse.lastIndexOf(']'); + + return (jsonEnd >= 0) && (jsonStart < jsonEnd) ? llmResponse.substring(jsonStart, jsonEnd + 1) : ""; + } + + protected OutputGuardrailResult invokeInvalidJson(AiMessage aiMessage, String json) { + LOGGER.debug("Found invalid JSON for aiMessage = {} and json = {}", aiMessage, json); + return reprompt(getInvalidJsonMessage(aiMessage, json), getInvalidJsonReprompt(aiMessage, json)); + } + + /** + * Generates a message indicating that the provided JSON is invalid. + * + * @param aiMessage the AI message associated with the invalid JSON. This parameter is not used. + * @param json the JSON that failed validation. This parameter is not used. + * @return a default message indicating that the JSON is invalid. + */ + protected String getInvalidJsonMessage( + @SuppressWarnings("unused") AiMessage aiMessage, @SuppressWarnings("unused") String json) { + return DEFAULT_REPROMPT_MESSAGE; + } + + /** + * Generates a reprompt message indicating that the provided JSON is invalid. + *

+ * This message is appended to the user message from the previous request. + *

+ * + * @param aiMessage the AI message associated with the invalid JSON. This parameter is not used. + * @param json the JSON input that failed validation. This parameter is not used. + * @return a reprompt message indicating that the JSON is invalid. + */ + protected String getInvalidJsonReprompt( + @SuppressWarnings("unused") AiMessage aiMessage, @SuppressWarnings("unused") String json) { + return DEFAULT_REPROMPT_PROMPT; + } + + /** + * Tries to deserialize the provided LLM response string into an object of type T using the configured {@link ObjectMapper}. + * If deserialization fails, an empty Optional is returned. + * + * @param llmResponse the JSON-formatted response string to be deserialized + * @return an Optional containing the deserialized object if successful, or an empty Optional if deserialization fails + */ + protected Optional deserialize(String llmResponse) { + try { + var obj = (this.outputClass != null) + ? this.objectMapper.readValue(llmResponse, this.outputClass) + : this.objectMapper.readValue(llmResponse, this.outputType); + + return Optional.ofNullable(obj); + } catch (JsonProcessingException e) { + return Optional.empty(); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java new file mode 100644 index 0000000000..07848f4f64 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrail.java @@ -0,0 +1,183 @@ +package dev.langchain4j.guardrail; + +import dev.langchain4j.data.message.AiMessage; +import java.util.Arrays; + +/** + * An output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * the expectations. + *

+ * In the case of reprompting, the reprompt message is added to the LLM context and the request is retried. + *

+ * The maximum number of retries is configurable, defaulting to {@link dev.langchain4j.guardrail.config.OutputGuardrailsConfig#MAX_RETRIES_DEFAULT}. + */ +public interface OutputGuardrail extends Guardrail { + /** + * Validates the response from the LLM. + * + * @param responseFromLLM + * the response from the LLM + */ + default OutputGuardrailResult validate(AiMessage responseFromLLM) { + return failure("Validation not implemented"); + } + + /** + * Validates the response from the LLM. + *

+ * Unlike {@link #validate(AiMessage)}, this method allows to access the memory and the augmentation result (in the + * case of a RAG). + *

+ * Implementation must not attempt to write to the memory or the augmentation result. + * + * @param params + * the parameters, including the response from the LLM, the memory, and the augmentation result. + */ + @Override + default OutputGuardrailResult validate(OutputGuardrailRequest params) { + return validate(params.responseFromLLM().aiMessage()); + } + + /** + * Produces a successful result without any successful text + * + * @return The result of a successful output guardrail validation. + */ + default OutputGuardrailResult success() { + return OutputGuardrailResult.success(); + } + + /** + * Produces a successful result with specific success text + * + * @return The result of a successful output guardrail validation with a specific text. + * + * @param successfulText + * The text of the successful result. + */ + default OutputGuardrailResult successWith(String successfulText) { + return OutputGuardrailResult.successWith(successfulText); + } + + /** + * Produces a non-fatal failure + * + * @return The result of a successful output guardrail validation with a specific text. + * + * @param successfulText + * The text of the successful result. + * @param successfulResult + * The object generated by this successful result. + */ + default OutputGuardrailResult successWith(String successfulText, Object successfulResult) { + return OutputGuardrailResult.successWith(successfulText, successfulResult); + } + + /** + * Produces a non-fatal failure + * + * @param message + * A message describing the failure. + * + * @return The result of a failed output guardrail validation. + */ + default OutputGuardrailResult failure(String message) { + return new OutputGuardrailResult(new OutputGuardrailResult.Failure(message), false); + } + + /** + * Produces a non-fatal failure + * + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * + * @return The result of a failed output guardrail validation. + */ + default OutputGuardrailResult failure(String message, Throwable cause) { + return new OutputGuardrailResult(new OutputGuardrailResult.Failure(message, cause), false); + } + + /** + * Produces a fatal failure + * + * @param message + * A message describing the failure. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation. + */ + default OutputGuardrailResult fatal(String message) { + return new OutputGuardrailResult(Arrays.asList(new OutputGuardrailResult.Failure(message)), true); + } + + /** + * Produces a fatal failure + * + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation. + */ + default OutputGuardrailResult fatal(String message, Throwable cause) { + return new OutputGuardrailResult(Arrays.asList(new OutputGuardrailResult.Failure(message, cause)), true); + } + + /** + * @param message + * A message describing the failure. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation and triggering a retry with the same user prompt. + */ + default OutputGuardrailResult retry(String message) { + return new OutputGuardrailResult(Arrays.asList(new OutputGuardrailResult.Failure(message, null, true)), true); + } + + /** + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation and triggering a retry with the same user prompt. + */ + default OutputGuardrailResult retry(String message, Throwable cause) { + return new OutputGuardrailResult(Arrays.asList(new OutputGuardrailResult.Failure(message, cause, true)), true); + } + + /** + * @param message + * A message describing the failure. + * @param reprompt + * The new prompt to be used for the retry. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation and triggering a retry with a new user prompt. + */ + default OutputGuardrailResult reprompt(String message, String reprompt) { + return new OutputGuardrailResult( + Arrays.asList(new OutputGuardrailResult.Failure(message, null, true, reprompt)), true); + } + + /** + * @param message + * A message describing the failure. + * @param cause + * The exception that caused this failure. + * @param reprompt + * The new prompt to be used for the retry. + * + * @return The result of a fatally failed output guardrail validation, blocking the evaluation of any other + * subsequent validation and triggering a retry with a new user prompt. + */ + default OutputGuardrailResult reprompt(String message, Throwable cause, String reprompt) { + return new OutputGuardrailResult( + Arrays.asList(new OutputGuardrailResult.Failure(message, cause, true, reprompt)), true); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailException.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailException.java new file mode 100644 index 0000000000..9ebba65159 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailException.java @@ -0,0 +1,17 @@ +package dev.langchain4j.guardrail; + +/** + * Exception thrown when an output guardrail validation fails. + *

+ * This class is not intended to be thrown within guardrail implementations. It is for the framework only. It is ok to catch it. + *

+ */ +public final class OutputGuardrailException extends GuardrailException { + public OutputGuardrailException(String message) { + super(message); + } + + public OutputGuardrailException(String message, Throwable cause) { + super(message, cause); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailExecutor.java new file mode 100644 index 0000000000..5586cad99b --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailExecutor.java @@ -0,0 +1,169 @@ +package dev.langchain4j.guardrail; + +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.OutputGuardrailResult.Failure; +import dev.langchain4j.guardrail.config.OutputGuardrailsConfig; +import dev.langchain4j.memory.ChatMemory; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.stream.Collectors; + +/** + * The {@link GuardrailExecutor} for {@link OutputGuardrail}s. + *

+ * When executing output guardrails, if any {@link OutputGuardrail} triggers a reprompt or retry, + * the new response has to go back through the entire chain of output guardrails to ensure the new response + * passes all the output guardrails. + *

+ */ +public non-sealed class OutputGuardrailExecutor + extends AbstractGuardrailExecutor< + OutputGuardrailsConfig, OutputGuardrailRequest, OutputGuardrailResult, OutputGuardrail, Failure> { + + public static final String MAX_RETRIES_MESSAGE_TEMPLATE = + """ + Output validation failed. The guardrails have reached the maximum number of retries. + Guardrail messages: + + %s + """; + + protected OutputGuardrailExecutor(OutputGuardrailsConfig config, List guardrails) { + super(config, guardrails); + } + + /** + * Executes the {@link OutputGuardrail}s on the given {@link OutputGuardrailRequest}. + * + * @param params The {@link OutputGuardrailRequest} to validate + * @return The {@link OutputGuardrailResult} of the validation + */ + @Override + public OutputGuardrailResult execute(OutputGuardrailRequest params) { + OutputGuardrailResult result = null; + var accumulatedParams = params; + var attempt = 0; + var maxAttempts = config().maxRetries(); + + if (maxAttempts == 0) { + maxAttempts = 1; + } else if (maxAttempts < 0) { + maxAttempts = OutputGuardrailsConfig.MAX_RETRIES_DEFAULT; + } + + while (attempt < maxAttempts) { + result = executeGuardrails(accumulatedParams); + + if (result.isSuccess()) { + return result; + } + + // Not successful + if (!result.isRetry()) { + // Not any kind of retry, so just stop here + throw new OutputGuardrailException(result.toString(), result.getFirstFailureException()); + } + + // If we get here we know it is some kind of retry + // We don't want to add intermediary UserMessages to the memory + var chatMessages = Optional.ofNullable( + accumulatedParams.requestParams().chatMemory()) + .map(ChatMemory::messages) + .orElseGet(ArrayList::new); + result.getReprompt().map(UserMessage::from).ifPresent(chatMessages::add); + + // Re-execute the request with the appended message + // But don't add it or the resulting message to the memory + var response = accumulatedParams.chatExecutor().execute(chatMessages); + + attempt++; + accumulatedParams = OutputGuardrailRequest.builder() + .responseFromLLM(response) + .chatExecutor(accumulatedParams.chatExecutor()) + .requestParams(accumulatedParams.requestParams()) + .build(); + } + + if (attempt == maxAttempts) { + var failureMessages = result.failures().stream() + .map(GuardrailResult.Failure::message) + .collect(Collectors.joining(System.lineSeparator())); + + throw new OutputGuardrailException(MAX_RETRIES_MESSAGE_TEMPLATE.formatted(failureMessages)); + } + + return result; + } + + /** + * Creates a failure result from some {@link Failure}s. + * @param failures The failures + * @return A {@link OutputGuardrailResult} containing the failures + */ + @Override + protected OutputGuardrailResult createFailure(List failures) { + return OutputGuardrailResult.failure(failures); + } + + /** + * Creates a success result. + * @return A {@link OutputGuardrailResult} representing success + */ + @Override + protected OutputGuardrailResult createSuccess() { + return OutputGuardrailResult.success(); + } + + @Override + protected OutputGuardrailException createGuardrailException(String message, Throwable cause) { + return new OutputGuardrailException(message, cause); + } + + @Override + protected OutputGuardrailResult handleFatalResult( + OutputGuardrailResult accumulatedResult, OutputGuardrailResult result) { + return accumulatedResult.hasRewrittenResult() ? result.blockRetry() : result; + } + + /** + * Creates a new instance of {@link OutputGuardrailExecutorBuilder}. + * The builder is used to construct and configure instances of {@link OutputGuardrailExecutorBuilder}. + * @return A new {@link OutputGuardrailExecutorBuilder} instance. + */ + public static OutputGuardrailExecutorBuilder builder() { + return new OutputGuardrailExecutorBuilder(); + } + + /** + * Builder class for constructing instances of {@link OutputGuardrailExecutor}. + * + * This builder allows configuration of an {@link OutputGuardrailExecutor} by specifying the associated configuration + * type ({@link OutputGuardrailsConfig}) and the output guardrails to be executed. + * + * Extends {@link GuardrailExecutorBuilder} for the specific types: + * - Configuration type: {@link OutputGuardrailsConfig} + * - Result type: {@link OutputGuardrailResult} + * - Parameter type: {@link OutputGuardrailRequest} + * - Guardrail type: {@link OutputGuardrail} + * + * Provides the {@code build()} method to create an {@link OutputGuardrailExecutor} instance. + */ + public static non-sealed class OutputGuardrailExecutorBuilder + extends GuardrailExecutorBuilder< + OutputGuardrailsConfig, + OutputGuardrailResult, + OutputGuardrailRequest, + OutputGuardrail, + OutputGuardrailExecutorBuilder> { + + protected OutputGuardrailExecutorBuilder() { + super(OutputGuardrailsConfig.builder().build()); + } + + @Override + public OutputGuardrailExecutor build() { + return new OutputGuardrailExecutor(config(), guardrails()); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailRequest.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailRequest.java new file mode 100644 index 0000000000..723efcd107 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailRequest.java @@ -0,0 +1,134 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.model.chat.ChatExecutor; +import dev.langchain4j.model.chat.response.ChatResponse; +import java.util.Optional; + +/** + * Represents the parameter passed to {@link OutputGuardrail#validate(OutputGuardrailRequest)}. + */ +public final class OutputGuardrailRequest implements GuardrailRequest { + private final ChatResponse responseFromLLM; + private final ChatExecutor chatExecutor; + private final GuardrailRequestParams requestParams; + + private OutputGuardrailRequest(Builder builder) { + this.responseFromLLM = ensureNotNull(builder.responseFromLLM, "responseFromLLM"); + this.requestParams = ensureNotNull(builder.requestParams, "requestParams"); + this.chatExecutor = ensureNotNull(builder.chatExecutor, "chatExecutor"); + } + + /** + * Returns the response from the LLM. + * + * @return the response from the LLM + */ + public ChatResponse responseFromLLM() { + return responseFromLLM; + } + + /** + * Returns the chat executor. + * + * @return the chat executor + */ + public ChatExecutor chatExecutor() { + return chatExecutor; + } + + /** + * Returns the common parameters that are shared across guardrail checks. + * + * @return an instance of {@code GuardrailRequestParams} containing shared parameters + */ + @Override + public GuardrailRequestParams requestParams() { + return requestParams; + } + + @Override + public OutputGuardrailRequest withText(String text) { + ensureNotNull(text, "text"); + + var aiMessage = Optional.ofNullable(this.responseFromLLM.aiMessage().toolExecutionRequests()) + .filter(t -> !t.isEmpty()) + .map(t -> new AiMessage(text, t)) + .orElseGet(() -> new AiMessage(text)); + + var chatResponse = ChatResponse.builder() + .aiMessage(aiMessage) + .metadata(this.responseFromLLM.metadata()) + .build(); + + return builder() + .responseFromLLM(chatResponse) + .chatExecutor(this.chatExecutor) + .requestParams(this.requestParams) + .build(); + } + + /** + * Creates a new builder for {@link OutputGuardrailRequest}. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Builder for {@link OutputGuardrailRequest}. + */ + public static class Builder { + private ChatResponse responseFromLLM; + private ChatExecutor chatExecutor; + private GuardrailRequestParams requestParams; + + private Builder() {} + + /** + * Sets the response from the LLM. + * + * @param responseFromLLM the response from the LLM + * @return this builder + */ + public Builder responseFromLLM(ChatResponse responseFromLLM) { + this.responseFromLLM = responseFromLLM; + return this; + } + + /** + * Sets the chat executor. + * + * @param chatExecutor the chat executor + * @return this builder + */ + public Builder chatExecutor(ChatExecutor chatExecutor) { + this.chatExecutor = chatExecutor; + return this; + } + + /** + * Sets the common parameters. + * + * @param requestParams the common parameters + * @return this builder + */ + public Builder requestParams(GuardrailRequestParams requestParams) { + this.requestParams = requestParams; + return this; + } + + /** + * Builds a new {@link OutputGuardrailRequest}. + * + * @return a new {@link OutputGuardrailRequest} + */ + public OutputGuardrailRequest build() { + return new OutputGuardrailRequest(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailResult.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailResult.java new file mode 100644 index 0000000000..a46b38542f --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/OutputGuardrailResult.java @@ -0,0 +1,290 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.model.chat.response.ChatResponse; +import java.util.Collections; +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import java.util.stream.Collectors; +import java.util.stream.Stream; + +/** + * The result of the validation of an {@link OutputGuardrail} + */ +public final class OutputGuardrailResult implements GuardrailResult { + private static final OutputGuardrailResult SUCCESS = new OutputGuardrailResult(); + + private final Result result; + private final String successfulText; + private final Object successfulResult; + private final List failures; + + private OutputGuardrailResult( + Result result, String successfulText, Object successfulResult, List failures) { + this.result = ensureNotNull(result, "result"); + this.successfulText = successfulText; + this.successfulResult = successfulResult; + this.failures = Optional.ofNullable(failures).orElseGet(List::of); + } + + private OutputGuardrailResult() { + this(Result.SUCCESS, null, null, Collections.emptyList()); + } + + private OutputGuardrailResult(String successfulText) { + this(Result.SUCCESS_WITH_RESULT, successfulText, null, Collections.emptyList()); + } + + private OutputGuardrailResult(String successfulText, Object successfulResult) { + this(Result.SUCCESS_WITH_RESULT, successfulText, successfulResult, Collections.emptyList()); + } + + OutputGuardrailResult(List failures, boolean fatal) { + this(fatal ? Result.FATAL : Result.FAILURE, null, null, failures); + } + + OutputGuardrailResult(Failure failure, boolean fatal) { + // Using Stream.of().collect() here because we need a mutable list + this(Stream.of(failure).collect(Collectors.toList()), fatal); + } + + /** + * Gets a successful output guardrail result + */ + public static OutputGuardrailResult success() { + return SUCCESS; + } + + /** + * Produces a successful result with specific success text + * + * @return The result of a successful output guardrail validation with a specific text. + * + * @param successfulText + * The text of the successful result. + */ + public static OutputGuardrailResult successWith(String successfulText) { + return (successfulText == null) ? success() : new OutputGuardrailResult(successfulText); + } + + /** + * Produces a non-fatal failure + * + * @param successfulText + * The text of the successful result. + * @param successfulResult + * The object generated by this successful result. + * @return The result of a successful output guardrail validation with a specific text. + */ + public static OutputGuardrailResult successWith(String successfulText, Object successfulResult) { + return new OutputGuardrailResult(successfulText, successfulResult); + } + + /** + * Produces a non-fatal failure + * + * @param failures A list of {@link Failure}s + * + * @return The result of a failed output guardrail validation. + */ + public static OutputGuardrailResult failure(List failures) { + return new OutputGuardrailResult(failures, false); + } + + /** + * Whether or not the guardrail is forcing a retry + */ + public boolean isRetry() { + return !isSuccess() && this.failures.stream().anyMatch(Failure::retry); + } + + /** + * Whether or not the guardrail is forcing a reprompt + */ + public boolean isReprompt() { + return !isSuccess() + && this.failures.stream() + .map(Failure::reprompt) + .filter(Objects::nonNull) + .count() + > 0; + } + + /** + * Block all retries for this result + */ + public OutputGuardrailResult blockRetry() { + this.failures.set(0, this.failures.get(0).blockRetry()); + return this; + } + + /** + * Gets the reprompt message + */ + public Optional getReprompt() { + return !isSuccess() + ? this.failures.stream() + .map(Failure::reprompt) + .filter(Objects::nonNull) + .findFirst() + : Optional.empty(); + } + + @Override + public String toString() { + return asString(); + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (o == null || getClass() != o.getClass()) return false; + OutputGuardrailResult that = (OutputGuardrailResult) o; + return result == that.result + && Objects.equals(successfulText, that.successfulText) + && Objects.equals(successfulResult, that.successfulResult) + && Objects.equals(failures, that.failures); + } + + @Override + public int hashCode() { + return Objects.hash(result, successfulText, successfulResult, failures); + } + + /** + * Gets the response computed from the combination of the original {@link ChatResponse} in the {@link OutputGuardrailRequest} + * and this result + * @param request The output guardrail request + * @param The type of response + * @return A response computed from the combination of the original {@link ChatResponse} in the {@link OutputGuardrailRequest} + * and this result + */ + public T response(OutputGuardrailRequest request) { + return (T) Optional.ofNullable(successfulResult).orElseGet(() -> createResponse(request)); + } + + private ChatResponse createResponse(OutputGuardrailRequest params) { + var response = params.responseFromLLM(); + var aiMessage = response.aiMessage(); + var newAiMessage = aiMessage; + + if (hasRewrittenResult()) { + newAiMessage = aiMessage.hasToolExecutionRequests() + ? AiMessage.from(successfulText(), aiMessage.toolExecutionRequests()) + : AiMessage.from(successfulText()); + } + + return response.toBuilder().aiMessage(newAiMessage).build(); + } + + @Override + public Result result() { + return result; + } + + @Override + @SuppressWarnings("unchecked") + public List failures() { + return (List) failures; + } + + @Override + public String successfulText() { + return successfulText; + } + + public Object successfulResult() { + return successfulResult; + } + + /** + * Represents an output guardrail failure + */ + public static final class Failure implements GuardrailResult.Failure { + private final String message; + private final Throwable cause; + private final Class guardrailClass; + private final boolean retry; + private final String reprompt; + + Failure( + String message, + Throwable cause, + Class guardrailClass, + boolean retry, + String reprompt) { + this.message = ensureNotNull(message, "message"); + this.cause = cause; + this.guardrailClass = guardrailClass; + this.retry = retry; + this.reprompt = reprompt; + } + + Failure(String message) { + this(message, null); + } + + Failure(String message, Throwable cause) { + this(message, cause, false); + } + + Failure(String message, Throwable cause, boolean retry) { + this(message, cause, null, retry, null); + } + + Failure(String message, Throwable cause, boolean retry, String reprompt) { + this(message, cause, null, retry, reprompt); + } + + @Override + public Failure withGuardrailClass(Class guardrailClass) { + ensureNotNull(guardrailClass, "guardrailClass"); + return new Failure(message(), cause(), guardrailClass, this.retry, this.reprompt); + } + + @Override + public String message() { + return message; + } + + @Override + public Throwable cause() { + return cause; + } + + @Override + public Class guardrailClass() { + return guardrailClass; + } + + /** + * Create a failure from this failure that blocks retries + */ + public Failure blockRetry() { + return this.retry + ? new Failure( + "Retry or reprompt is not allowed after a rewritten output", + cause(), + this.guardrailClass, + false, + this.reprompt) + : this; + } + + @Override + public String toString() { + return asString(); + } + + public boolean retry() { + return retry; + } + + public String reprompt() { + return reprompt; + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultInputGuardrailsConfig.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultInputGuardrailsConfig.java new file mode 100644 index 0000000000..819e131524 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultInputGuardrailsConfig.java @@ -0,0 +1,30 @@ +package dev.langchain4j.guardrail.config; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +/** + * The default implementation of {@link InputGuardrailsConfig} for this library if no other libraries provide their own implementations. + */ +final class DefaultInputGuardrailsConfig implements InputGuardrailsConfig { + DefaultInputGuardrailsConfig(Builder builder) { + ensureNotNull(builder, "builder"); + } + + /** + * Gets a builder instance for building {@link DefaultInputGuardrailsConfig} instances. + * @return The builder instance for building {@link DefaultInputGuardrailsConfig} instances. + */ + static Builder builder() { + return new Builder(); + } + + /** + * Builder for {@link DefaultInputGuardrailsConfig} instances. + */ + static class Builder implements InputGuardrailsConfigBuilder { + @Override + public InputGuardrailsConfig build() { + return new DefaultInputGuardrailsConfig(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultOutputGuardrailsConfig.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultOutputGuardrailsConfig.java new file mode 100644 index 0000000000..6f38959239 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/DefaultOutputGuardrailsConfig.java @@ -0,0 +1,46 @@ +package dev.langchain4j.guardrail.config; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +/** + * The default implementation of {@link OutputGuardrailsConfig} for this library if no other libraries provide their own implementations. + */ +final class DefaultOutputGuardrailsConfig implements OutputGuardrailsConfig { + private final int maxRetries; + + DefaultOutputGuardrailsConfig(Builder builder) { + ensureNotNull(builder, "builder"); + this.maxRetries = builder.maxRetries; + } + + /** + * Gets a builder instance for building {@link DefaultOutputGuardrailsConfig} instances. + * @return The builder instance for building {@link DefaultOutputGuardrailsConfig} instances. + */ + static Builder builder() { + return new Builder(); + } + + @Override + public int maxRetries() { + return this.maxRetries; + } + + /** + * Builder for {@link DefaultOutputGuardrailsConfig} instances. + */ + static class Builder implements OutputGuardrailsConfigBuilder { + private int maxRetries = MAX_RETRIES_DEFAULT; + + @Override + public Builder maxRetries(int maxRetries) { + this.maxRetries = maxRetries; + return this; + } + + @Override + public OutputGuardrailsConfig build() { + return new DefaultOutputGuardrailsConfig(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfig.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfig.java new file mode 100644 index 0000000000..9ed4213ffa --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfig.java @@ -0,0 +1,6 @@ +package dev.langchain4j.guardrail.config; + +/** + * Base interface for common configuration across all kinds of guardrails. + */ +public interface GuardrailsConfig {} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfigBuilder.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfigBuilder.java new file mode 100644 index 0000000000..f92ed96294 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/GuardrailsConfigBuilder.java @@ -0,0 +1,13 @@ +package dev.langchain4j.guardrail.config; + +/** + * Builder for {@link GuardrailsConfig} instances. + * @param The type of configuration being build + */ +public interface GuardrailsConfigBuilder { + /** + * Builds the configuration. + * @return The configuration + */ + C build(); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/InputGuardrailsConfig.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/InputGuardrailsConfig.java new file mode 100644 index 0000000000..539c39053b --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/InputGuardrailsConfig.java @@ -0,0 +1,32 @@ +package dev.langchain4j.guardrail.config; + +import dev.langchain4j.spi.guardrail.config.InputGuardrailsConfigBuilderFactory; +import java.util.ServiceLoader; + +/** + * Configuration specifically for input guardrails. + *

+ * Frameworks that extend this library (like Quarkus or Spring) may provide their own implementations of this configuration. + *

+ */ +public interface InputGuardrailsConfig extends GuardrailsConfig { + /** + * Gets a builder instance for building {@link InputGuardrailsConfig} instances. + * @return A {@link InputGuardrailsConfigBuilder} for building {@link InputGuardrailsConfig} instances. + */ + static InputGuardrailsConfigBuilder builder() { + return ServiceLoader.load(InputGuardrailsConfigBuilderFactory.class) + .findFirst() + .map(InputGuardrailsConfigBuilderFactory::get) + .orElseGet(DefaultInputGuardrailsConfig::builder); + } + + /** + * Builder for {@link InputGuardrailsConfig} instances. + *

+ * This is needed so other frameworks (like Quarkus and Spring) can extend the configuration mechanism with their own + * implementations while also adhering to the interfaces and specs defined here. + *

+ */ + interface InputGuardrailsConfigBuilder extends GuardrailsConfigBuilder {} +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/OutputGuardrailsConfig.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/OutputGuardrailsConfig.java new file mode 100644 index 0000000000..736e437dd3 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/OutputGuardrailsConfig.java @@ -0,0 +1,57 @@ +package dev.langchain4j.guardrail.config; + +import dev.langchain4j.spi.guardrail.config.OutputGuardrailsConfigBuilderFactory; +import java.util.ServiceLoader; + +/** + * Configuration specifically for output guardrails. + *

+ * Frameworks that extend this library (like Quarkus or Spring) may provide their own implementations of this configuration. + *

+ */ +public interface OutputGuardrailsConfig extends GuardrailsConfig { + /** + * Default maximum number of retries for the guardrail. + */ + int MAX_RETRIES_DEFAULT = 2; + + /** + * Configures the maximum number of retries for the guardrail. + *

+ * Defaults to {@link #MAX_RETRIES_DEFAULT} if not set. + *

+ * Set to {@code 0} to disable retries. + */ + int maxRetries(); + + /** + * Gets a newBuilder instance for building {@link OutputGuardrailsConfig} instances. + * @return A {@link OutputGuardrailsConfigBuilder} for building {@link OutputGuardrailsConfig} instances. + */ + static OutputGuardrailsConfigBuilder builder() { + return ServiceLoader.load(OutputGuardrailsConfigBuilderFactory.class) + .findFirst() + .map(OutputGuardrailsConfigBuilderFactory::get) + .orElseGet(DefaultOutputGuardrailsConfig::builder); + } + + /** + * Builder for {@link OutputGuardrailsConfig} instances. + *

+ * This is needed so other frameworks (like Quarkus and Spring) can extend the configuration mechanism with their own + * implementations while also adhering to the interfaces and specs defined here. + *

+ */ + interface OutputGuardrailsConfigBuilder extends GuardrailsConfigBuilder { + /** + * Sets the maximum number of retries for output guardrails. + *

+ * Defaults to {@link OutputGuardrailsConfig#maxRetries()} if not set. + *

+ * @param maxRetries The maximum number of retries for output guardrails + * @return The maximum number of retries for output guardrails + * @see OutputGuardrailsConfig#maxRetries() + */ + OutputGuardrailsConfigBuilder maxRetries(int maxRetries); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/package-info.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/package-info.java new file mode 100644 index 0000000000..73de7d61bf --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/config/package-info.java @@ -0,0 +1,4 @@ +@NullMarked +package dev.langchain4j.guardrail.config; + +import org.jspecify.annotations.NullMarked; diff --git a/langchain4j-core/src/main/java/dev/langchain4j/guardrail/package-info.java b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/package-info.java new file mode 100644 index 0000000000..7d5daa537a --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/guardrail/package-info.java @@ -0,0 +1,4 @@ +@Experimental +package dev.langchain4j.guardrail; + +import dev.langchain4j.Experimental; diff --git a/langchain4j-core/src/main/java/dev/langchain4j/internal/Exceptions.java b/langchain4j-core/src/main/java/dev/langchain4j/internal/Exceptions.java index 9f5127c1e9..5c29fe9c41 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/internal/Exceptions.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/internal/Exceptions.java @@ -20,7 +20,7 @@ public class Exceptions { * @return the constructed exception. */ public static IllegalArgumentException illegalArgument(String format, Object... args) { - return new IllegalArgumentException(String.format(format, args)); + return new IllegalArgumentException(format.formatted(args)); } /** @@ -33,6 +33,6 @@ public class Exceptions { * @return the constructed exception. */ public static RuntimeException runtime(String format, Object... args) { - return new RuntimeException(String.format(format, args)); + return new RuntimeException(format.formatted(args)); } } diff --git a/langchain4j-core/src/main/java/dev/langchain4j/memory/ChatMemory.java b/langchain4j-core/src/main/java/dev/langchain4j/memory/ChatMemory.java index f54f54a8a7..08247bf05a 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/memory/ChatMemory.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/memory/ChatMemory.java @@ -1,7 +1,7 @@ package dev.langchain4j.memory; import dev.langchain4j.data.message.ChatMessage; - +import java.util.Arrays; import java.util.List; /** @@ -25,6 +25,26 @@ public interface ChatMemory { */ void add(ChatMessage message); + /** + * Adds messages to the chat memory + * @param messages The {@link ChatMessage}s to add + */ + default void add(ChatMessage... messages) { + if ((messages != null) && (messages.length > 0)) { + add(Arrays.asList(messages)); + } + } + + /** + * Adds messages to the chat memory + * @param messages The {@link ChatMessage}s to add + */ + default void add(Iterable messages) { + if (messages != null) { + messages.forEach(this::add); + } + } + /** * Retrieves messages from the chat memory. * Depending on the implementation, it may not return all previously added messages, diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/AbstractChatExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/AbstractChatExecutor.java new file mode 100644 index 0000000000..233ca6bcae --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/AbstractChatExecutor.java @@ -0,0 +1,55 @@ +package dev.langchain4j.model.chat; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.Internal; +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import java.util.List; + +/** + * Abstract base class for chat executors that provides a common structure and shared functionality + * for implementing the {@link ChatExecutor} interface. + * + * This class encapsulates a {@link ChatRequest} and allows subclasses to define how + * the request should be processed by implementing the {@code execute(ChatRequest)} method. + * + * Subclasses are expected to be immutable and should provide specific implementations for + * executing chat requests, typically using particular chat models or processing strategies. + * + * Responsibilities: + * - Stores a {@link ChatRequest} object which can be used to build specific chat requests. + * - Provides standard implementations for executing a chat request with a list of messages + * or without any additional input. + * - Defines an abstract method {@code execute(ChatRequest)} for subclasses to implement + * specific execution logic. + */ +@Internal +abstract class AbstractChatExecutor implements ChatExecutor { + protected final ChatRequest chatRequest; + + protected AbstractChatExecutor(AbstractBuilder builder) { + this.chatRequest = ensureNotNull(builder.chatRequest, "chatRequest"); + } + + @Override + public ChatResponse execute(List chatMessages) { + var newChatRequest = this.chatRequest.toBuilder().messages(chatMessages).build(); + + return execute(newChatRequest); + } + + @Override + public ChatResponse execute() { + return execute(this.chatRequest); + } + + /** + * Executes a given chat request and returns the corresponding chat response. + * + * @param chatRequest the chat request to process, containing the input messages and any necessary configurations + * @return the chat response generated as a result of processing the given chat request + */ + protected abstract ChatResponse execute(ChatRequest chatRequest); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatExecutor.java new file mode 100644 index 0000000000..a611e60464 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatExecutor.java @@ -0,0 +1,166 @@ +package dev.langchain4j.model.chat; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.function.Consumer; + +/** + * Generic executor interface that defines a chat interaction + */ +public interface ChatExecutor { + /** + * Execute a chat request + * @return The response + */ + ChatResponse execute(); + + /** + * Executes a chat request using the provided chat memory. + * + * @param chatMemory The chat memory containing the context of the conversation. + * It provides the history of messages required for proper interaction with the chat language model. + * @return A response object containing the AI's response and additional metadata. + * @see #execute(List) + */ + default ChatResponse execute(ChatMemory chatMemory) { + var messages = Optional.ofNullable(chatMemory).map(ChatMemory::messages).orElseGet(ArrayList::new); + + return execute(messages); + } + + /** + * Executes a chat request using the provided chat messages + * @param chatMessages The chat messages containing the context of the conversation. + * It provides the history of messages required for proper interaction with the chat model + * @return A response object containing the AI's response and additional metadata. + */ + ChatResponse execute(List chatMessages); + + /** + * Creates a new {@link SynchronousBuilder} instance for constructing {@link ChatExecutor} objects + * that perform synchronous chat requests. + * + * @return A new {@link SynchronousBuilder} instance to configure and build a {@link ChatExecutor}. + */ + static SynchronousBuilder builder(ChatModel chatModel) { + return new SynchronousBuilder(chatModel); + } + + /** + * Creates a new {@link StreamingToSynchronousBuilder} instance for constructing {@link ChatExecutor} objects + * that perform streaming chat requests. + * + * @return A new {@link StreamingToSynchronousBuilder} instance to configure and build a {@link ChatExecutor}. + */ + static StreamingToSynchronousBuilder builder(StreamingChatModel streamingChatModel) { + return new StreamingToSynchronousBuilder(streamingChatModel); + } + + /** + * An abstract base-builder class for constructing instances of {@link ChatExecutor}. + * + * This class provides a fluent API for setting required components, such as + * {@link ChatRequest}, and defines a contract for building {@link ChatExecutor} + * instances. Subclasses should implement the {@code build()} method to ensure + * proper construction of the target chat executor object. + * + * @param the type of the builder subclass for enabling fluent method chaining + */ + abstract class AbstractBuilder> { + protected ChatRequest chatRequest; + + protected AbstractBuilder() {} + + /** + * Sets the {@link ChatRequest} instance for the synchronousBuilder. + * The {@link ChatRequest} encapsulates the input messages and parameters required + * to generate a response from the chat model. + * + * @param chatRequest the {@link ChatRequest} containing the input messages and parameters + * @return the updated SynchronousBuilder instance + */ + public AbstractBuilder chatRequest(ChatRequest chatRequest) { + this.chatRequest = chatRequest; + return this; + } + + /** + * Constructs and returns an instance of {@link ChatExecutor}. + * Ensures that all required parameters have been appropriately set + * before building the {@link ChatExecutor}. + * + * @return a fully constructed {@link ChatExecutor} instance + */ + public abstract ChatExecutor build(); + } + + /** + * SynchronousBuilder for constructing instances of {@link ChatExecutor}. + * + * This synchronousBuilder provides a fluent API for setting required components + * like {@link ChatRequest}, and for building an instance of the {@link ChatExecutor}. + */ + class SynchronousBuilder extends AbstractBuilder { + protected final ChatModel chatModel; + + protected SynchronousBuilder(ChatModel chatModel) { + this.chatModel = ensureNotNull(chatModel, "chatModel"); + } + + /** + * Constructs and returns an instance of {@link ChatExecutor}. + * Ensures that all required parameters have been appropriately set + * before building the {@link ChatExecutor}. + * + * @return a fully constructed {@link ChatExecutor} instance + */ + public ChatExecutor build() { + return new SynchronousChatExecutor(this); + } + } + + /** + * StreamingToSynchronousBuilder for constructing instances of {@link ChatExecutor}. + * + * This streaming build provides a fluent API for setting required components + * like {@link ChatRequest}, and for building an instance of the {@link ChatExecutor} + * that simulates streaming. + */ + class StreamingToSynchronousBuilder extends AbstractBuilder { + protected final StreamingChatModel streamingChatModel; + protected Consumer errorHandler; + + protected StreamingToSynchronousBuilder(StreamingChatModel streamingChatModel) { + this.streamingChatModel = ensureNotNull(streamingChatModel, "streamingChatModel"); + } + + /** + * Sets a custom error handler to manage exceptions or errors that occur during the execution. + * + * @param errorHandler a {@link Consumer} of {@link Throwable} that processes the error + * @return the current {@link StreamingToSynchronousBuilder} instance for method chaining + */ + public StreamingToSynchronousBuilder errorHandler(Consumer errorHandler) { + this.errorHandler = errorHandler; + return this; + } + + /** + * Constructs and returns an instance of {@link ChatExecutor}. + * Ensures that all required parameters have been appropriately set + * before building the {@link ChatExecutor}. + * + * @return a fully constructed {@link ChatExecutor} instance + */ + public ChatExecutor build() { + return new StreamingToSynchronousChatExecutor(this); + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatModel.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatModel.java index de1523f797..6878edefdb 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatModel.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/ChatModel.java @@ -1,5 +1,10 @@ package dev.langchain4j.model.chat; +import static dev.langchain4j.model.ModelProvider.OTHER; +import static dev.langchain4j.model.chat.ChatModelListenerUtils.onError; +import static dev.langchain4j.model.chat.ChatModelListenerUtils.onRequest; +import static dev.langchain4j.model.chat.ChatModelListenerUtils.onResponse; + import dev.langchain4j.data.message.ChatMessage; import dev.langchain4j.data.message.UserMessage; import dev.langchain4j.model.ModelProvider; @@ -8,17 +13,11 @@ import dev.langchain4j.model.chat.request.ChatRequest; import dev.langchain4j.model.chat.request.ChatRequestParameters; import dev.langchain4j.model.chat.request.DefaultChatRequestParameters; import dev.langchain4j.model.chat.response.ChatResponse; - import java.util.List; import java.util.Map; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; -import static dev.langchain4j.model.chat.ChatModelListenerUtils.onError; -import static dev.langchain4j.model.chat.ChatModelListenerUtils.onRequest; -import static dev.langchain4j.model.chat.ChatModelListenerUtils.onResponse; -import static dev.langchain4j.model.ModelProvider.OTHER; - /** * Represents a language model that has a chat API. * @@ -71,9 +70,8 @@ public interface ChatModel { default String chat(String userMessage) { - ChatRequest chatRequest = ChatRequest.builder() - .messages(UserMessage.from(userMessage)) - .build(); + ChatRequest chatRequest = + ChatRequest.builder().messages(UserMessage.from(userMessage)).build(); ChatResponse chatResponse = chat(chatRequest); @@ -82,18 +80,14 @@ public interface ChatModel { default ChatResponse chat(ChatMessage... messages) { - ChatRequest chatRequest = ChatRequest.builder() - .messages(messages) - .build(); + ChatRequest chatRequest = ChatRequest.builder().messages(messages).build(); return chat(chatRequest); } default ChatResponse chat(List messages) { - ChatRequest chatRequest = ChatRequest.builder() - .messages(messages) - .build(); + ChatRequest chatRequest = ChatRequest.builder().messages(messages).build(); return chat(chatRequest); } diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/StreamingToSynchronousChatExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/StreamingToSynchronousChatExecutor.java new file mode 100644 index 0000000000..0aea9da15d --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/StreamingToSynchronousChatExecutor.java @@ -0,0 +1,93 @@ +package dev.langchain4j.model.chat; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.Internal; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; +import java.util.Optional; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * A concrete implementation of the {@link ChatExecutor} interface that executes + * chat requests using a specified {@link StreamingChatModel}. It then executes the requests as if it were + * synchronous, essentially transforming a streaming request to a synchronous request + * + * This class utilizes a {@link ChatRequest} to encapsulate the input messages + * and parameters and delegates the execution of the chat to the provided {@link StreamingChatModel}. + * + * Instances of this class are immutable and are typically instantiated using + * the {@link StreamingToSynchronousBuilder}. + */ +@Internal +final class StreamingToSynchronousChatExecutor extends AbstractChatExecutor { + private final StreamingChatModel streamingChatModel; + private final Consumer errorHandler; + + protected StreamingToSynchronousChatExecutor(StreamingToSynchronousBuilder builder) { + super(builder); + + this.streamingChatModel = ensureNotNull(builder.streamingChatModel, "streamingChatModel"); + this.errorHandler = builder.errorHandler; + } + + @Override + protected ChatResponse execute(ChatRequest chatRequest) { + var responseHandler = new StreamingToSyncResponseHandler(this.errorHandler); + this.streamingChatModel.chat(chatRequest, responseHandler); + + return Optional.ofNullable(responseHandler.getResponse()).orElseGet(ChatResponse.builder()::build); + } + + private static class StreamingToSyncResponseHandler implements StreamingChatResponseHandler { + private static final Logger LOG = LoggerFactory.getLogger(StreamingToSyncResponseHandler.class); + private final Consumer errorHandler; + private final CountDownLatch latch = new CountDownLatch(1); + private AtomicReference response = new AtomicReference<>(); + + StreamingToSyncResponseHandler(Consumer errorHandler) { + this.errorHandler = errorHandler; + } + + @Override + public void onPartialResponse(String partialResponse) {} + + @Override + public void onCompleteResponse(ChatResponse completeResponse) { + response.set(completeResponse); + this.latch.countDown(); + } + + private void waitForCompletion() { + try { + this.latch.await(); + } catch (InterruptedException e) { + throw new RuntimeException(e); + } + } + + ChatResponse getResponse() { + waitForCompletion(); + return this.response.get(); + } + + @Override + public void onError(Throwable error) { + if (errorHandler != null) { + try { + errorHandler.accept(error); + } catch (Exception e) { + LOG.error("While handling the following error...", error); + LOG.error("...the following error happened", e); + } + } else { + LOG.warn("Ignored error", error); + } + } + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/SynchronousChatExecutor.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/SynchronousChatExecutor.java new file mode 100644 index 0000000000..839565ae26 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/SynchronousChatExecutor.java @@ -0,0 +1,33 @@ +package dev.langchain4j.model.chat; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.Internal; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; + +/** + * A concrete implementation of the {@link ChatExecutor} interface that executes + * chat requests using a specified {@link ChatModel}. + * + * This class utilizes a {@link ChatRequest} to encapsulate the input messages + * and parameters and delegates the execution of the chat to the provided + * {@link ChatModel}. + * + * Instances of this class are immutable and are typically instantiated using + * the {@link SynchronousBuilder}. + */ +@Internal +final class SynchronousChatExecutor extends AbstractChatExecutor { + private final ChatModel chatModel; + + protected SynchronousChatExecutor(SynchronousBuilder builder) { + super(builder); + this.chatModel = ensureNotNull(builder.chatModel, "chatModel"); + } + + @Override + protected ChatResponse execute(ChatRequest chatRequest) { + return this.chatModel.chat(chatRequest); + } +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/request/ChatRequest.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/request/ChatRequest.java index 59dbc4d279..252cb1dd63 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/request/ChatRequest.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/request/ChatRequest.java @@ -1,17 +1,15 @@ package dev.langchain4j.model.chat.request; -import dev.langchain4j.agent.tool.ToolSpecification; -import dev.langchain4j.data.message.ChatMessage; - -import java.util.ArrayList; -import java.util.List; -import java.util.Objects; - import static dev.langchain4j.internal.Utils.copy; import static dev.langchain4j.internal.Utils.isNullOrEmpty; import static dev.langchain4j.internal.ValidationUtils.ensureNotEmpty; import static java.util.Arrays.asList; +import dev.langchain4j.agent.tool.ToolSpecification; +import dev.langchain4j.data.message.ChatMessage; +import java.util.List; +import java.util.Objects; + public class ChatRequest { private final List messages; @@ -131,8 +129,7 @@ public class ChatRequest { if (this == o) return true; if (o == null || getClass() != o.getClass()) return false; ChatRequest that = (ChatRequest) o; - return Objects.equals(this.messages, that.messages) - && Objects.equals(this.parameters, that.parameters); + return Objects.equals(this.messages, that.messages) && Objects.equals(this.parameters, that.parameters); } @Override @@ -142,10 +139,14 @@ public class ChatRequest { @Override public String toString() { - return "ChatRequest {" + - " messages = " + messages + - ", parameters = " + parameters + - " }"; + return "ChatRequest {" + " messages = " + messages + ", parameters = " + parameters + " }"; + } + + /** + * Transforms this instance to a {@link Builder} with all of the same field values + */ + public Builder toBuilder() { + return new Builder(this); } public static Builder builder() { @@ -169,6 +170,13 @@ public class ChatRequest { private ToolChoice toolChoice; private ResponseFormat responseFormat; + public Builder() {} + + public Builder(ChatRequest chatRequest) { + this.messages = chatRequest.messages; + this.parameters = chatRequest.parameters; + } + public Builder messages(List messages) { this.messages = messages; return this; diff --git a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/response/ChatResponse.java b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/response/ChatResponse.java index 765afa0468..ef4966040f 100644 --- a/langchain4j-core/src/main/java/dev/langchain4j/model/chat/response/ChatResponse.java +++ b/langchain4j-core/src/main/java/dev/langchain4j/model/chat/response/ChatResponse.java @@ -1,13 +1,12 @@ package dev.langchain4j.model.chat.response; +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + import dev.langchain4j.data.message.AiMessage; import dev.langchain4j.model.output.FinishReason; import dev.langchain4j.model.output.TokenUsage; - import java.util.Objects; -import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; - public class ChatResponse { private final AiMessage aiMessage; @@ -44,6 +43,16 @@ public class ChatResponse { return aiMessage; } + /** + * Converts the current instance of {@code ChatResponse} into a {@link Builder}, + * allowing modifications to the current object's fields. + * + * @return a new {@link Builder} instance initialized with the current state of this {@code ChatResponse}. + */ + public Builder toBuilder() { + return new Builder(this); + } + public ChatResponseMetadata metadata() { return metadata; } @@ -69,8 +78,7 @@ public class ChatResponse { if (this == o) return true; if (o == null || getClass() != o.getClass()) return false; ChatResponse that = (ChatResponse) o; - return Objects.equals(this.aiMessage, that.aiMessage) - && Objects.equals(this.metadata, that.metadata); + return Objects.equals(this.aiMessage, that.aiMessage) && Objects.equals(this.metadata, that.metadata); } @Override @@ -80,10 +88,7 @@ public class ChatResponse { @Override public String toString() { - return "ChatResponse {" + - " aiMessage = " + aiMessage + - ", metadata = " + metadata + - " }"; + return "ChatResponse {" + " aiMessage = " + aiMessage + ", metadata = " + metadata + " }"; } public static Builder builder() { @@ -91,7 +96,6 @@ public class ChatResponse { } public static class Builder { - private AiMessage aiMessage; private ChatResponseMetadata metadata; @@ -100,6 +104,13 @@ public class ChatResponse { private TokenUsage tokenUsage; private FinishReason finishReason; + public Builder() {} + + public Builder(ChatResponse chatResponse) { + this.aiMessage = chatResponse.aiMessage; + this.metadata = chatResponse.metadata; + } + public Builder aiMessage(AiMessage aiMessage) { this.aiMessage = aiMessage; return this; diff --git a/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassInstanceFactory.java b/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassInstanceFactory.java new file mode 100644 index 0000000000..2006f76000 --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassInstanceFactory.java @@ -0,0 +1,19 @@ +package dev.langchain4j.spi.classloading; + +/** + * A factory for providing instances of classes + *

+ * Intended to be implemented by downstream frameworks (like Quarkus and Spring) where rather than creating + * classes on-the-fly, they will most likely be managed by some dependency injection framework. + *

+ */ +public interface ClassInstanceFactory { + /** + * Provides an instance of the specified class type. + * + * @param the type of the class + * @param clazz the class object representing the type + * @return an instance of the specified class type + */ + T getInstanceOfClass(Class clazz); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassMetadataProviderFactory.java b/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassMetadataProviderFactory.java new file mode 100644 index 0000000000..a4d7d239eb --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/spi/classloading/ClassMetadataProviderFactory.java @@ -0,0 +1,46 @@ +package dev.langchain4j.spi.classloading; + +import java.lang.annotation.Annotation; +import java.util.Optional; + +/** + * A factory interface for providing access to class metadata. Intended to be implemented by downstream frameworks. + *

+ * {@code dev.langchain4j.classinstance.ReflectionBasedClassMetadataProviderFactory} + * provides an implementation that uses reflection, which is probably fine in most cases, but this provides hooks for + * other frameworks (like Quarkus) which don't use reflection to provide class metadata. + *

+ * + * @param The type of the method key, representing a unique identifier for methods. Can be whatever it needs to be. + */ +public interface ClassMetadataProviderFactory { + /** + * Retrieves an annotation of the specified type from the given method. + * + * @param The type of the annotation to locate, which must extend {@link Annotation}. + * @param method The method from which the annotation is to be retrieved. + * @param annotationClass The class object corresponding to the annotation type to find. + * @return An {@code Optional} containing the located annotation, or an empty {@code Optional} if the annotation + * is not present on the specified method. + */ + Optional getAnnotation(MethodKey method, Class annotationClass); + + /** + * Retrieves an annotation of the specified type from the given class. + * + * @param The type of the annotation to locate, which must extend {@link Annotation}. + * @param clazz The class from which the annotation is to be retrieved. + * @param annotationClass The class object corresponding to the annotation type to find. + * @return An {@code Optional} containing the located annotation, or an empty {@code Optional} if the annotation + * is not present on the specified class. + */ + Optional getAnnotation(Class clazz, Class annotationClass); + + /** + * Retrieves an iterable containing method keys for all non-static methods defined in the specified class. + * + * @param clazz The class from which to retrieve methods. + * @return An iterable of method keys corresponding to the methods of the specified class. + */ + Iterable getNonStaticMethodsOnClass(Class clazz); +} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/InputGuardrailsConfigBuilderFactory.java b/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/InputGuardrailsConfigBuilderFactory.java new file mode 100644 index 0000000000..356869036f --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/InputGuardrailsConfigBuilderFactory.java @@ -0,0 +1,10 @@ +package dev.langchain4j.spi.guardrail.config; + +import dev.langchain4j.guardrail.config.InputGuardrailsConfig; +import java.util.function.Supplier; + +/** + * SPI for overriding and/or extending the default {@link InputGuardrailsConfig.InputGuardrailsConfigBuilder} implementation. + */ +public interface InputGuardrailsConfigBuilderFactory + extends Supplier {} diff --git a/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/OutputGuardrailsConfigBuilderFactory.java b/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/OutputGuardrailsConfigBuilderFactory.java new file mode 100644 index 0000000000..b3cc1e987a --- /dev/null +++ b/langchain4j-core/src/main/java/dev/langchain4j/spi/guardrail/config/OutputGuardrailsConfigBuilderFactory.java @@ -0,0 +1,10 @@ +package dev.langchain4j.spi.guardrail.config; + +import dev.langchain4j.guardrail.config.OutputGuardrailsConfig; +import java.util.function.Supplier; + +/** + * SPI for overriding and/or extending the default {@link OutputGuardrailsConfig.OutputGuardrailsConfigBuilder} implementation. + */ +public interface OutputGuardrailsConfigBuilderFactory + extends Supplier {} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/classinstance/ClassInstanceLoaderTests.java b/langchain4j-core/src/test/java/dev/langchain4j/classinstance/ClassInstanceLoaderTests.java new file mode 100644 index 0000000000..fd326574e6 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/classinstance/ClassInstanceLoaderTests.java @@ -0,0 +1,18 @@ +package dev.langchain4j.classinstance; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; + +class ClassInstanceLoaderTests { + @Test + void loadsClassInstance() { + var instance1 = ClassInstanceLoader.getClassInstance(SomeClass.class); + var instance2 = ClassInstanceLoader.getClassInstance(SomeClass.class); + + assertThat(instance1).isNotNull().isExactlyInstanceOf(SomeClass.class); + assertThat(instance2).isNotNull().isExactlyInstanceOf(SomeClass.class).isNotEqualTo(instance1); + } + + public static class SomeClass {} +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/guardrail/InputGuardrailExecutorTests.java b/langchain4j-core/src/test/java/dev/langchain4j/guardrail/InputGuardrailExecutorTests.java new file mode 100644 index 0000000000..76acf2ceac --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/guardrail/InputGuardrailExecutorTests.java @@ -0,0 +1,264 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.test.guardrail.GuardrailAssertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.config.InputGuardrailsConfig; +import java.util.Map; +import java.util.stream.IntStream; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ParameterContext; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.aggregator.AggregateWith; +import org.junit.jupiter.params.aggregator.ArgumentsAccessor; +import org.junit.jupiter.params.aggregator.ArgumentsAggregationException; +import org.junit.jupiter.params.aggregator.ArgumentsAggregator; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +import org.mockito.Mockito; + +class InputGuardrailExecutorTests { + @ParameterizedTest(name = "{0}") + @MethodSource("successGuardrails") + void allSuccessfulGuardrails( + @SuppressWarnings("unused") String testDesc, + int howManyShouldExecute, + @AggregateWith(InputGuardrailAggregator.class) InputGuardrail... guardrails) { + + var spiedGuardrails = Stream.of(guardrails).map(Mockito::spy).toArray(InputGuardrail[]::new); + var params = from(UserMessage.from("test")); + var executor = + InputGuardrailExecutor.builder().guardrails(spiedGuardrails).build(); + var result = executor.execute(params); + + assertThat(result).isSuccessful(); + + IntStream.range(0, howManyShouldExecute) + .mapToObj(i -> (SuccessInputGuardrail) spiedGuardrails[i]) + .forEach(guardrail -> { + assertThat(guardrail.shouldBeExecuted).isTrue(); + verify(guardrail).validate(params); + }); + + IntStream.range(howManyShouldExecute, spiedGuardrails.length) + .mapToObj(i -> (SuccessInputGuardrail) spiedGuardrails[i]) + .forEach(guardrail -> { + assertThat(guardrail.shouldBeExecuted).isFalse(); + verify(guardrail, never()).validate(params); + }); + } + + @Test + void noGuardrails() { + var params = from(UserMessage.from("test")); + var executor = InputGuardrailExecutor.builder().build(); + var result = executor.execute(params); + + assertThat(result).isSuccessful(); + } + + @ParameterizedTest(name = "{0}") + @MethodSource("failedFatalGuardrails") + void failedFatal( + @SuppressWarnings("unused") String testDesc, + int howManyShouldExecute, + int howManyFailures, + @AggregateWith(InputGuardrailAggregator.class) InputGuardrail... guardrails) { + + var spiedGuardrails = Stream.of(guardrails).map(Mockito::spy).toArray(InputGuardrail[]::new); + var params = from(UserMessage.from("test")); + var executor = InputGuardrailExecutor.builder() + .guardrails(spiedGuardrails) + .config(InputGuardrailsConfig.builder().build()) + .build(); + + assertThatExceptionOfType(InputGuardrailException.class) + .isThrownBy(() -> executor.execute(params)) + .withMessageMatching("The guardrail " + getClass().getName() + + "\\$.+Guardrail failed with this message: failure \\d"); + + IntStream.range(0, howManyShouldExecute) + .mapToObj(i -> spiedGuardrails[i]) + .forEach(guardrail -> { + var shouldBeExecuted = (guardrail instanceof SuccessInputGuardrail s) + ? s.shouldBeExecuted + : ((FailureInputGuardrail) guardrail).shouldBeExecuted; + + assertThat(shouldBeExecuted).isTrue(); + verify(guardrail).validate(params); + }); + + IntStream.range(howManyShouldExecute, spiedGuardrails.length) + .mapToObj(i -> spiedGuardrails[i]) + .forEach(guardrail -> { + var shouldBeExecuted = (guardrail instanceof SuccessInputGuardrail s) + ? s.shouldBeExecuted + : ((FailureInputGuardrail) guardrail).shouldBeExecuted; + + assertThat(shouldBeExecuted).isFalse(); + verify(guardrail, never()).validate(params); + }); + + var numFailedGuardrails = Stream.of(spiedGuardrails) + .filter(FailureInputGuardrail.class::isInstance) + .map(FailureInputGuardrail.class::cast) + .filter(guardrail -> guardrail.shouldBeExecuted) + .count(); + + assertThat(numFailedGuardrails).isEqualTo(howManyFailures); + } + + static Stream successGuardrails() { + return Stream.of( + Arguments.of("No guardrails", 0), + Arguments.of("One successful guardrail", 1, new SuccessInputGuardrail()), + Arguments.of("Two successful guardrails", 2, new SuccessInputGuardrail(), new SuccessInputGuardrail()), + Arguments.of( + "Three successful guardrails", + 3, + new SuccessInputGuardrail(), + new SuccessInputGuardrail(), + new SuccessInputGuardrail())); + } + + static Stream failedFatalGuardrails() { + return Stream.of( + Arguments.of( + "One successful one fatal guardrail", + 2, + 1, + new SuccessInputGuardrail(), + new FatalInputGuardrail(1)), + Arguments.of( + "One fatal one successful guardrail", + 1, + 1, + new FatalInputGuardrail(1), + new SuccessInputGuardrail(false)), + Arguments.of( + "One successful one fatal one successful guardrails", + 2, + 1, + new SuccessInputGuardrail(), + new FatalInputGuardrail(1), + new SuccessInputGuardrail(false)), + Arguments.of( + "One successful one fatal one failed guardrails", + 2, + 1, + new SuccessInputGuardrail(), + new FatalInputGuardrail(1), + new FailureInputGuardrail<>(2).shouldNotBeExecuted()), + Arguments.of( + "One failure one successful guardrail", + 2, + 1, + new FailureInputGuardrail<>(1), + new SuccessInputGuardrail()), + Arguments.of( + "One successful one failure one successful guardrails", + 3, + 1, + new SuccessInputGuardrail(), + new FailureInputGuardrail<>(1), + new SuccessInputGuardrail()), + Arguments.of( + "One successful one fatal one failure guardrails", + 2, + 1, + new SuccessInputGuardrail(), + new FatalInputGuardrail(1), + new FailureInputGuardrail<>(2).shouldNotBeExecuted()), + Arguments.of( + "Two failure guardrails", 2, 2, new FailureInputGuardrail<>(1), new FailureInputGuardrail<>(2)), + Arguments.of( + "One successful one failure one fatal one failure guardrails", + 3, + 2, + new SuccessInputGuardrail(), + new FailureInputGuardrail<>(2), + new FatalInputGuardrail(1), + new FailureInputGuardrail<>(3).shouldNotBeExecuted())); + } + + public static InputGuardrailRequest from(UserMessage userMessage) { + var newCommonParams = GuardrailRequestParams.builder() + .chatMemory(null) + .augmentationResult(null) + .userMessageTemplate("") + .variables(Map.of()) + .build(); + + return InputGuardrailRequest.builder() + .userMessage(userMessage) + .commonParams(newCommonParams) + .build(); + } + + private static class FatalInputGuardrail extends FailureInputGuardrail { + private FatalInputGuardrail(int failureNumber) { + super(failureNumber); + } + + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return fatal(this.failureMessage); + } + } + + private static class FailureInputGuardrail implements InputGuardrail { + protected final String failureMessage; + private boolean shouldBeExecuted = true; + + private FailureInputGuardrail(int failureNumber) { + this("failure " + failureNumber); + } + + private FailureInputGuardrail(String failureMessage) { + this.failureMessage = failureMessage; + } + + G shouldNotBeExecuted() { + this.shouldBeExecuted = false; + return (G) this; + } + + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return failure(this.failureMessage); + } + } + + private static class SuccessInputGuardrail implements InputGuardrail { + private boolean shouldBeExecuted = true; + + SuccessInputGuardrail(boolean shouldBeExecuted) { + this.shouldBeExecuted = shouldBeExecuted; + } + + SuccessInputGuardrail() { + this(true); + } + + @Override + public InputGuardrailResult validate(final UserMessage userMessage) { + return InputGuardrailResult.success(); + } + } + + static class InputGuardrailAggregator implements ArgumentsAggregator { + @Override + public Object aggregateArguments(ArgumentsAccessor accessor, ParameterContext context) + throws ArgumentsAggregationException { + + return accessor.toList().stream() + .skip(context.getIndex()) + .map(InputGuardrail.class::cast) + .toArray(InputGuardrail[]::new); + } + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrailTests.java b/langchain4j-core/src/test/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrailTests.java new file mode 100644 index 0000000000..21fdb836c8 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/guardrail/JsonExtractorOutputGuardrailTests.java @@ -0,0 +1,96 @@ +package dev.langchain4j.guardrail; + +import static dev.langchain4j.test.guardrail.GuardrailAssertions.assertThat; +import static org.mockito.Mockito.any; +import static org.mockito.Mockito.anyString; +import static org.mockito.Mockito.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.spy; +import static org.mockito.Mockito.verify; + +import com.fasterxml.jackson.core.type.TypeReference; +import dev.langchain4j.data.message.AiMessage; +import java.util.Map; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class JsonExtractorOutputGuardrailTests { + private static final String JSON = + """ + { + "name": "MyObject", + "description": "Description of MyObject" + }"""; + + private static final JsonExtractorOutputGuardrail MY_OBJECT_JSON_OUTPUT_GUARDRAIL = + new JsonExtractorOutputGuardrail<>(MyObject.class); + private static final JsonExtractorOutputGuardrail> MAP_OF_MY_OBJECT_JSON_OUTPUT_GUARDRAIL = + new JsonExtractorOutputGuardrail<>(new TypeReference<>() {}); + + @ParameterizedTest + @MethodSource("guardrails") + void successfulValidation(String json, JsonExtractorOutputGuardrail guardrail, Object expectedResult) { + var guardrailSpy = spy(guardrail); + var result = guardrailSpy.validate(AiMessage.from(json)); + + assertThat(result) + .isNotNull() + .extracting( + OutputGuardrailResult::result, + OutputGuardrailResult::successfulText, + OutputGuardrailResult::successfulResult) + .containsExactly(GuardrailResult.Result.SUCCESS_WITH_RESULT, json, expectedResult); + + verify(guardrailSpy, never()).trimNonJson(anyString()); + } + + @ParameterizedTest + @MethodSource("guardrails") + void successfulValidationAfterTrimming( + String json, JsonExtractorOutputGuardrail guardrail, Object expectedResult) { + var guardrailSpy = spy(guardrail); + var input = "abc" + json; + var result = guardrailSpy.validate(AiMessage.from(input)); + + assertThat(result) + .isNotNull() + .extracting( + OutputGuardrailResult::result, + OutputGuardrailResult::successfulText, + OutputGuardrailResult::successfulResult) + .containsExactly(GuardrailResult.Result.SUCCESS_WITH_RESULT, json, expectedResult); + + verify(guardrailSpy).trimNonJson(input); + } + + @Test + void invalidJson() { + var guardrail = spy(MY_OBJECT_JSON_OUTPUT_GUARDRAIL); + var input = "{{" + JSON; + var result = guardrail.validate(AiMessage.from(input)); + + assertThat(result) + .hasSingleFailureWithMessageAndReprompt( + JsonExtractorOutputGuardrail.DEFAULT_REPROMPT_MESSAGE, + JsonExtractorOutputGuardrail.DEFAULT_REPROMPT_PROMPT); + + verify(guardrail).trimNonJson(input); + verify(guardrail).invokeInvalidJson(any(AiMessage.class), eq(input)); + } + + static Stream guardrails() { + var result = new MyObject("MyObject", "Description of MyObject"); + + return Stream.of( + Arguments.of(JSON, MY_OBJECT_JSON_OUTPUT_GUARDRAIL, result), + Arguments.of( + "{ \"myObject\": %s}".formatted(JSON), + MAP_OF_MY_OBJECT_JSON_OUTPUT_GUARDRAIL, + Map.of("myObject", result))); + } + + record MyObject(String name, String description) {} +} diff --git a/langchain4j-kotlin/src/test/kotlin/dev/langchain4j/kotlin/service/ServiceWithFlowTest.kt b/langchain4j-kotlin/src/test/kotlin/dev/langchain4j/kotlin/service/ServiceWithFlowTest.kt index 5cd48c9bd1..94275c222c 100644 --- a/langchain4j-kotlin/src/test/kotlin/dev/langchain4j/kotlin/service/ServiceWithFlowTest.kt +++ b/langchain4j-kotlin/src/test/kotlin/dev/langchain4j/kotlin/service/ServiceWithFlowTest.kt @@ -10,6 +10,7 @@ import dev.langchain4j.data.message.AiMessage import dev.langchain4j.kotlin.model.chat.StreamingChatModelReply import dev.langchain4j.kotlin.model.chat.StreamingChatModelReply.CompleteResponse import dev.langchain4j.kotlin.model.chat.StreamingChatModelReply.PartialResponse +import dev.langchain4j.model.chat.ChatModel import dev.langchain4j.model.chat.StreamingChatModel import dev.langchain4j.model.chat.request.ChatRequest import dev.langchain4j.model.chat.response.ChatResponse @@ -34,7 +35,10 @@ import org.mockito.kotlin.whenever @ExtendWith(MockitoExtension::class) internal class ServiceWithFlowTest { @Mock - private lateinit var model: StreamingChatModel + private lateinit var streamingModel: StreamingChatModel + + @Mock + private lateinit var model: ChatModel @Test fun `Should use TokenStreamToStringFlowAdapter`() = runTest { @@ -47,14 +51,9 @@ internal class ServiceWithFlowTest { handler.onPartialResponse(partialToken1) handler.onPartialResponse(partialToken2) handler.onCompleteResponse(completeResponse) - }.whenever(model).chat(any(), any()) - - val assistant = - AiServices - .builder(Assistant::class.java) - .streamingChatModel(model) - .build() + }.whenever(streamingModel).chat(any(), any()) + val assistant = AiServices.create(Assistant::class.java, streamingModel) val result = assistant.askQuestion(userName = "My friend", question = "How are you?") .toList() @@ -72,16 +71,10 @@ internal class ServiceWithFlowTest { handler.onPartialResponse(partialToken1) handler.onPartialResponse(partialToken2) handler.onError(error) - }.whenever(model) + }.whenever(streamingModel) .chat(any(), any()) - val assistant = - AiServices - .builder(Assistant::class.java) - .streamingChatModel(model) - .build() - - + val assistant = AiServices.create(Assistant::class.java, streamingModel) val response = assistant.askQuestion(userName = "My friend", question = "How are you?") .catch { val message = @@ -103,14 +96,9 @@ internal class ServiceWithFlowTest { handler.onPartialResponse(partialToken1) handler.onPartialResponse(partialToken2) handler.onCompleteResponse(completeResponse) - }.whenever(model).chat(any(), any()) - - val assistant = - AiServices - .builder(Assistant::class.java) - .streamingChatModel(model) - .build() + }.whenever(streamingModel).chat(any(), any()) + val assistant = AiServices.create(Assistant::class.java, streamingModel) val result = assistant.askQuestion2(userName = "My friend", question = "How are you?") .toList() @@ -132,14 +120,9 @@ internal class ServiceWithFlowTest { handler.onPartialResponse(partialToken1) handler.onPartialResponse(partialToken2) handler.onError(error) - }.whenever(model).chat(any(), any()) - - val assistant = - AiServices - .builder(Assistant::class.java) - .streamingChatModel(model) - .build() + }.whenever(streamingModel).chat(any(), any()) + val assistant = AiServices.create(Assistant::class.java, streamingModel) val response = assistant.askQuestion2(userName = "My friend", question = "How are you?") .catch { emit(StreamingChatModelReply.Error(it)) } .toList() diff --git a/langchain4j-test/pom.xml b/langchain4j-test/pom.xml new file mode 100644 index 0000000000..8a63e56e92 --- /dev/null +++ b/langchain4j-test/pom.xml @@ -0,0 +1,106 @@ + + + 4.0.0 + + + dev.langchain4j + langchain4j-parent + 1.1.0-beta7-SNAPSHOT + ../langchain4j-parent/pom.xml + + + langchain4j-test + LangChain4j :: Test + Testing utility classes and interfaces of LangChain4j + + + + Apache License, Version 2.0 + https://www.apache.org/licenses/LICENSE-2.0.txt + repo + + + + + + -options + + requireUpperBoundDeps + + + + + org.assertj + assertj-core + + + + dev.langchain4j + langchain4j-core + 1.1.0-SNAPSHOT + + + + + + org.junit.jupiter + junit-jupiter-api + test + + + + org.junit.jupiter + junit-jupiter-params + test + + + + + + + org.apache.maven.plugins + maven-jar-plugin + + + + test-jar + + + + + + + org.apache.maven.plugins + maven-source-plugin + + + attach-sources + + jar-no-fork + + test-jar-no-fork + + + + + + + + + + + org.jacoco + jacoco-maven-plugin + + + + + report + + + + + + + + diff --git a/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailAssertions.java b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailAssertions.java new file mode 100644 index 0000000000..46247ebaea --- /dev/null +++ b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailAssertions.java @@ -0,0 +1,31 @@ +package dev.langchain4j.test.guardrail; + +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import org.assertj.core.api.Assertions; + +/** + * Custom assertions for working with Guardrails + *

+ * This follows the pattern described in https://assertj.github.io/doc/#assertj-core-custom-assertions-entry-point + *

+ */ +public class GuardrailAssertions extends Assertions { + /** + * Returns an {@link OutputGuardrailResultAssert} for assertions on an {@link OutputGuardrailResult} + * @param actual The actual {@link OutputGuardrailResult} + * @return The {@link OutputGuardrailResultAssert} + */ + public static OutputGuardrailResultAssert assertThat(OutputGuardrailResult actual) { + return OutputGuardrailResultAssert.assertThat(actual); + } + + /** + * Returns an {@link InputGuardrailResultAssert} for assertions on an {@link InputGuardrailResult} + * @param actual The actual {@link InputGuardrailResult} + * @return The {@link InputGuardrailResultAssert} + */ + public static InputGuardrailResultAssert assertThat(InputGuardrailResult actual) { + return InputGuardrailResultAssert.assertThat(actual); + } +} diff --git a/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailResultAssert.java b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailResultAssert.java new file mode 100644 index 0000000000..d501c94340 --- /dev/null +++ b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/GuardrailResultAssert.java @@ -0,0 +1,159 @@ +package dev.langchain4j.test.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; +import static org.assertj.core.api.Assertions.assertThat; + +import dev.langchain4j.guardrail.GuardrailResult; +import dev.langchain4j.guardrail.GuardrailResult.Failure; +import dev.langchain4j.guardrail.GuardrailResult.Result; +import dev.langchain4j.guardrail.InputGuardrailResult; +import java.util.Objects; +import java.util.function.Consumer; +import org.assertj.core.api.AbstractObjectAssert; +import org.assertj.core.api.InstanceOfAssertFactories; +import org.assertj.core.api.ListAssert; + +/** + * Custom assertions for {@link GuardrailResult}s + *

+ * This follows the pattern described in https://assertj.github.io/doc/#assertj-core-custom-assertions-creation + *

+ * @param The type of {@link GuardrailResultAssert} + * @param The type of {@link GuardrailResult} + * @param The type of {@link Failure} + */ +public abstract sealed class GuardrailResultAssert< + A extends GuardrailResultAssert, R extends GuardrailResult, F extends Failure> + extends AbstractObjectAssert permits InputGuardrailResultAssert, OutputGuardrailResultAssert { + + private final Class failureClass; + + protected GuardrailResultAssert(R r, Class resultType, Class failureClass) { + super(r, resultType); + this.failureClass = failureClass; + } + + /** + * Asserts that the actual object's {@link Result} matches the given expected result. + * If the result does not match, an assertion error is thrown with the actual and expected values. + * + * @param result the expected result to compare against the actual object's result + * @return this assertion object for method chaining + * @throws AssertionError if the actual result does not match the expected result + */ + public A hasResult(Result result) { + isNotNull(); + + if (!Objects.equals(actual.result(), result)) { + throw failureWithActualExpected( + actual.result(), result, "Expected result to be <%s> but was <%s>", result, actual.result()); + } + + return (A) this; + } + + /** + * Asserts that the actual {@link Result} contains the specified successful text. + * This method verifies that the actual instance is in a successful state and that + * the successful text matches the expected value. If the assertion fails, an error + * is thrown, detailing the mismatch. + * + * @param successfulText the expected text for the successful state + * @return this assertion object for method chaining + * @throws AssertionError if the actual object is not successful or if the successful text does not match the expected text + */ + public A hasSuccessfulText(String successfulText) { + isSuccessful(); + + if (!Objects.equals(actual.successfulText(), successfulText)) { + throw failureWithActualExpected( + actual.successfulText(), + successfulText, + "Expected successful text to be <%s> but was <%s>", + successfulText, + actual.successfulText()); + } + + return (A) this; + } + + /** + * Asserts that the actual {@code InputGuardrailResult} represents a successful state. + * A successful state is determined by having {@link InputGuardrailResult#isSuccess()}. + * + * @return this assertion object for method chaining + * @throws AssertionError if the actual result is not successful as per the aforementioned criteria + */ + public A isSuccessful() { + isNotNull(); + + if (!actual.isSuccess()) { + throw failure("Expected result to be successful but was <%s>", actual.result()); + } + + return (A) this; + } + + /** + * Asserts that the actual {@code InputGuardrailResult} contains failures. + * The method validates that the object being asserted is not null and + * that there are failures present within the result. + * + * @return this assertion object for method chaining + * @throws AssertionError if the actual object is null or if the failures are empty + */ + public A hasFailures() { + isNotNull(); + withFailures().isNotEmpty(); + + return (A) this; + } + + /** + * Asserts that the actual {@code InputGuardrailResult} contains exactly one failure with the specified message. + * If the assertion fails, an error is thrown detailing the problem. + * + * @param expectedFailureMessage the expected message of the single failure + * @return this assertion object for method chaining + * @throws AssertionError if the actual object is null, if there are no failures, + * if there is more than one failure, or if the single failure + * does not match the specified message + */ + public A hasSingleFailureWithMessage(String expectedFailureMessage) { + isNotNull(); + + withFailures().singleElement().extracting(Failure::message).isEqualTo(expectedFailureMessage); + + return (A) this; + } + + /** + * Asserts that the {@code InputGuardrailResult} contains exactly one {@link GuardrailResult.Failure} and verifies + * that this failure meets the specified requirements. The requirements are defined by the provided {@link Consumer}. + * + * @param requirements a {@link Consumer} that defines the assertions to be applied to the single failure. + * Must not be {@code null}. + * @return this assertion object for method chaining. + * @throws NullPointerException if the {@code requirements} is {@code null}. + * @throws AssertionError if the actual object is {@code null}, if there are no failures, if there is more than + * one failure, or if the single failure does not satisfy the specified requirements. + * @see #satisfies(Consumer[]) + */ + public A assertSingleFailureSatisfies(Consumer requirements) { + isNotNull(); + ensureNotNull(requirements, "requirements"); + + withFailures().singleElement().satisfies(requirements); + + return (A) this; + } + + /** + * Returns a {@link ListAssert} for the failures of the actual {@link GuardrailResult}. + */ + public ListAssert withFailures() { + return assertThat(actual.failures()) + .isNotNull() + .asInstanceOf(InstanceOfAssertFactories.list(this.failureClass)); + } +} diff --git a/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/InputGuardrailResultAssert.java b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/InputGuardrailResultAssert.java new file mode 100644 index 0000000000..4990951789 --- /dev/null +++ b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/InputGuardrailResultAssert.java @@ -0,0 +1,22 @@ +package dev.langchain4j.test.guardrail; + +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.InputGuardrailResult.Failure; + +/** + * Custom assertions for {@link InputGuardrailResult}s + *

+ * This follows the pattern described in https://assertj.github.io/doc/#assertj-core-custom-assertions-creation + *

+ */ +public final class InputGuardrailResultAssert + extends GuardrailResultAssert { + + private InputGuardrailResultAssert(InputGuardrailResult inputGuardrailResult) { + super(inputGuardrailResult, InputGuardrailResultAssert.class, Failure.class); + } + + public static InputGuardrailResultAssert assertThat(InputGuardrailResult actual) { + return new InputGuardrailResultAssert(actual); + } +} diff --git a/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/OutputGuardrailResultAssert.java b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/OutputGuardrailResultAssert.java new file mode 100644 index 0000000000..ed6ab7d769 --- /dev/null +++ b/langchain4j-test/src/main/java/dev/langchain4j/test/guardrail/OutputGuardrailResultAssert.java @@ -0,0 +1,53 @@ +package dev.langchain4j.test.guardrail; + +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrailResult.Failure; + +/** + * Custom assertions for {@link OutputGuardrailResult}s + *

+ * This follows the pattern described in https://assertj.github.io/doc/#assertj-core-custom-assertions-creation + *

+ */ +public final class OutputGuardrailResultAssert + extends GuardrailResultAssert { + + private OutputGuardrailResultAssert(OutputGuardrailResult outputGuardrailResult) { + super(outputGuardrailResult, OutputGuardrailResultAssert.class, Failure.class); + } + + /** + * Creates a new {@code OutputGuardrailResultAssert} for the provided {@code OutputGuardrailResult}. + * + * @param actual the {@code OutputGuardrailResult} to be asserted + * @return an {@code OutputGuardrailResultAssert} instance for chaining further assertions + */ + public static OutputGuardrailResultAssert assertThat(OutputGuardrailResult actual) { + return new OutputGuardrailResultAssert(actual); + } + + /** + * Asserts that the actual {@code OutputGuardrailResult} contains exactly one failure with the specified message and + * reprompt. + * If the assertion fails, an error is thrown detailing the problem. + * + * @param expectedFailureMessage the expected message of the single failure + * @param expectedReprompt the expected reprompt + * @return this assertion object for method chaining + * @throws AssertionError if the actual object is null, if there are no failures, + * if there is more than one failure, or if the single failure + * does not match the specified message + */ + public OutputGuardrailResultAssert hasSingleFailureWithMessageAndReprompt( + String expectedFailureMessage, String expectedReprompt) { + + isNotNull(); + + withFailures() + .singleElement() + .extracting(Failure::message, Failure::retry, Failure::reprompt) + .containsExactly(expectedFailureMessage, true, expectedReprompt); + + return this; + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/classloading/ClassMetadataProvider.java b/langchain4j/src/main/java/dev/langchain4j/classloading/ClassMetadataProvider.java new file mode 100644 index 0000000000..7ed75b1624 --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/classloading/ClassMetadataProvider.java @@ -0,0 +1,35 @@ +package dev.langchain4j.classloading; + +import dev.langchain4j.spi.classloading.ClassMetadataProviderFactory; +import java.util.ServiceLoader; +import java.util.ServiceLoader.Provider; + +/** + * Utility class for returning metadata about a class and its methods. Intended to allow downstream frameworks (like Quarkus + * or Spring) to use their own mechanisms for providing this information. + */ +public final class ClassMetadataProvider { + private static final ReflectionBasedClassMetadataProviderFactory DEFAULT_CLASS_METADATA_PROVIDER_FACTORY = + new ReflectionBasedClassMetadataProviderFactory(); + + private ClassMetadataProvider() {} + + /** + * Retrieves an implementation of a {@link ClassMetadataProviderFactory}. This method first looks for + * implementations of the factory via the {@link ServiceLoader}. It filters out the default factory implementation + * ({@link ReflectionBasedClassMetadataProviderFactory}) to allow for custom implementations provided by external frameworks. + * If no custom implementations are available, the method returns the default factory. + * + * @param The type of the method key, representing a unique identifier for methods. + * @return An instance of {@link ClassMetadataProviderFactory} either provided by an external framework or falling back + * to the default implementation. + */ + public static ClassMetadataProviderFactory getClassMetadataProviderFactory() { + return ServiceLoader.load(ClassMetadataProviderFactory.class).stream() + .filter(provider -> + !DEFAULT_CLASS_METADATA_PROVIDER_FACTORY.getClass().equals(provider.type())) + .map(Provider::get) + .findFirst() + .orElse(DEFAULT_CLASS_METADATA_PROVIDER_FACTORY); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/classloading/ReflectionBasedClassMetadataProviderFactory.java b/langchain4j/src/main/java/dev/langchain4j/classloading/ReflectionBasedClassMetadataProviderFactory.java new file mode 100644 index 0000000000..6bf639e7f4 --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/classloading/ReflectionBasedClassMetadataProviderFactory.java @@ -0,0 +1,32 @@ +package dev.langchain4j.classloading; + +import dev.langchain4j.spi.classloading.ClassMetadataProviderFactory; +import java.lang.annotation.Annotation; +import java.lang.reflect.Method; +import java.lang.reflect.Modifier; +import java.util.Optional; +import java.util.stream.Stream; + +/** + * Implementation of the {@link ClassMetadataProviderFactory} interface using Java Reflection. + * This class provides methods to retrieve annotations and method metadata from classes + * via reflection-based mechanisms. + */ +public final class ReflectionBasedClassMetadataProviderFactory implements ClassMetadataProviderFactory { + @Override + public Optional getAnnotation(Method method, Class annotationClass) { + return Optional.ofNullable(method.getAnnotation(annotationClass)); + } + + @Override + public Optional getAnnotation(Class clazz, Class annotationClass) { + return Optional.ofNullable(clazz.getAnnotation(annotationClass)); + } + + @Override + public Iterable getNonStaticMethodsOnClass(Class clazz) { + return Stream.of(clazz.getMethods()) + .filter(method -> !Modifier.isStatic(method.getModifiers())) + .toList(); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceContext.java b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceContext.java index 1b0ab3e727..f2bf891125 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceContext.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceContext.java @@ -2,14 +2,16 @@ package dev.langchain4j.service; import dev.langchain4j.Internal; import dev.langchain4j.memory.ChatMemory; -import dev.langchain4j.service.memory.ChatMemoryService; import dev.langchain4j.memory.chat.ChatMemoryProvider; import dev.langchain4j.model.chat.ChatModel; import dev.langchain4j.model.chat.StreamingChatModel; import dev.langchain4j.model.moderation.ModerationModel; import dev.langchain4j.rag.RetrievalAugmentor; +import dev.langchain4j.service.guardrail.GuardrailService; +import dev.langchain4j.service.memory.ChatMemoryService; import dev.langchain4j.service.tool.ToolService; import java.util.Optional; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.Function; @Internal @@ -26,6 +28,9 @@ public class AiServiceContext { public ToolService toolService = new ToolService(); + public final GuardrailService.Builder guardrailServiceBuilder; + private final AtomicReference guardrailService = new AtomicReference<>(); + public ModerationModel moderationModel; public RetrievalAugmentor retrievalAugmentor; @@ -34,6 +39,7 @@ public class AiServiceContext { public AiServiceContext(Class aiServiceClass) { this.aiServiceClass = aiServiceClass; + this.guardrailServiceBuilder = GuardrailService.builder(aiServiceClass); } public boolean hasChatMemory() { @@ -47,4 +53,9 @@ public class AiServiceContext { public void initChatMemories(ChatMemoryProvider chatMemoryProvider) { chatMemoryService = new ChatMemoryService(chatMemoryProvider); } + + public GuardrailService guardrailService() { + return this.guardrailService.updateAndGet( + service -> (service != null) ? service : guardrailServiceBuilder.build()); + } } diff --git a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceStreamingResponseHandler.java b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceStreamingResponseHandler.java index f443a7e793..28652982a0 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceStreamingResponseHandler.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceStreamingResponseHandler.java @@ -1,39 +1,44 @@ package dev.langchain4j.service; +import static dev.langchain4j.internal.Utils.copy; +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + import dev.langchain4j.Internal; import dev.langchain4j.agent.tool.ToolExecutionRequest; import dev.langchain4j.agent.tool.ToolSpecification; import dev.langchain4j.data.message.AiMessage; import dev.langchain4j.data.message.ChatMessage; import dev.langchain4j.data.message.ToolExecutionResultMessage; +import dev.langchain4j.guardrail.GuardrailRequestParams; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.ChatExecutor; import dev.langchain4j.model.chat.request.ChatRequest; import dev.langchain4j.model.chat.response.ChatResponse; import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; import dev.langchain4j.model.output.TokenUsage; import dev.langchain4j.service.tool.ToolExecution; import dev.langchain4j.service.tool.ToolExecutor; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - import java.util.ArrayList; import java.util.List; import java.util.Map; import java.util.function.Consumer; - -import static dev.langchain4j.internal.Utils.copy; -import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** - * Handles response from a language model for AI Service that is streamed token-by-token. - * Handles both regular (text) responses and responses with the request to execute one or multiple tools. + * Handles response from a language model for AI Service that is streamed token-by-token. Handles both regular (text) + * responses and responses with the request to execute one or multiple tools. */ @Internal class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler { + private static final Logger LOG = LoggerFactory.getLogger(AiServiceStreamingResponseHandler.class); - private final Logger log = LoggerFactory.getLogger(AiServiceStreamingResponseHandler.class); - + private final ChatExecutor chatExecutor; private final AiServiceContext context; private final Object memoryId; + private final GuardrailRequestParams commonGuardrailParams; + private final Object methodKey; private final Consumer partialResponseHandler; private final Consumer toolExecutionHandler; @@ -41,45 +46,59 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler private final Consumer errorHandler; - private final List temporaryMemory; + private final ChatMemory temporaryMemory; private final TokenUsage tokenUsage; private final List toolSpecifications; private final Map toolExecutors; + private final List responseBuffer = new ArrayList<>(); + private final boolean hasOutputGuardrails; - AiServiceStreamingResponseHandler(AiServiceContext context, - Object memoryId, - Consumer partialResponseHandler, - Consumer toolExecutionHandler, - Consumer completeResponseHandler, - Consumer errorHandler, - List temporaryMemory, - TokenUsage tokenUsage, - List toolSpecifications, - Map toolExecutors) { + AiServiceStreamingResponseHandler( + ChatExecutor chatExecutor, + AiServiceContext context, + Object memoryId, + Consumer partialResponseHandler, + Consumer toolExecutionHandler, + Consumer completeResponseHandler, + Consumer errorHandler, + ChatMemory temporaryMemory, + TokenUsage tokenUsage, + List toolSpecifications, + Map toolExecutors, + GuardrailRequestParams commonGuardrailParams, + Object methodKey) { + this.chatExecutor = ensureNotNull(chatExecutor, "chatExecutor"); this.context = ensureNotNull(context, "context"); this.memoryId = ensureNotNull(memoryId, "memoryId"); + this.methodKey = methodKey; this.partialResponseHandler = ensureNotNull(partialResponseHandler, "partialResponseHandler"); this.completeResponseHandler = completeResponseHandler; this.toolExecutionHandler = toolExecutionHandler; this.errorHandler = errorHandler; - this.temporaryMemory = new ArrayList<>(temporaryMemory); + this.temporaryMemory = temporaryMemory; this.tokenUsage = ensureNotNull(tokenUsage, "tokenUsage"); + this.commonGuardrailParams = commonGuardrailParams; this.toolSpecifications = copy(toolSpecifications); this.toolExecutors = copy(toolExecutors); + this.hasOutputGuardrails = context.guardrailService().hasOutputGuardrails(methodKey); } @Override public void onPartialResponse(String partialResponse) { - partialResponseHandler.accept(partialResponse); + // If we're using output guardrails, then buffer the partial response until the guardrails have completed + if (hasOutputGuardrails) { + responseBuffer.add(partialResponse); + } else { + partialResponseHandler.accept(partialResponse); + } } @Override public void onCompleteResponse(ChatResponse completeResponse) { - AiMessage aiMessage = completeResponse.aiMessage(); addToMemory(aiMessage); @@ -88,10 +107,8 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler String toolName = toolExecutionRequest.name(); ToolExecutor toolExecutor = toolExecutors.get(toolName); String toolExecutionResult = toolExecutor.execute(toolExecutionRequest, memoryId); - ToolExecutionResultMessage toolExecutionResultMessage = ToolExecutionResultMessage.from( - toolExecutionRequest, - toolExecutionResult - ); + ToolExecutionResultMessage toolExecutionResultMessage = + ToolExecutionResultMessage.from(toolExecutionRequest, toolExecutionResult); addToMemory(toolExecutionResultMessage); if (toolExecutionHandler != null) { @@ -108,7 +125,8 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler .toolSpecifications(toolSpecifications) .build(); - StreamingChatResponseHandler handler = new AiServiceStreamingResponseHandler( + var handler = new AiServiceStreamingResponseHandler( + chatExecutor, context, memoryId, partialResponseHandler, @@ -118,8 +136,9 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler temporaryMemory, TokenUsage.sum(tokenUsage, completeResponse.metadata().tokenUsage()), toolSpecifications, - toolExecutors - ); + toolExecutors, + commonGuardrailParams, + methodKey); context.streamingChatModel.chat(chatRequest, handler); } else { @@ -127,26 +146,57 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler ChatResponse finalChatResponse = ChatResponse.builder() .aiMessage(aiMessage) .metadata(completeResponse.metadata().toBuilder() - .tokenUsage(tokenUsage.add(completeResponse.metadata().tokenUsage())) + .tokenUsage(tokenUsage.add( + completeResponse.metadata().tokenUsage())) .build()) .build(); + + // Invoke output guardrails + if (hasOutputGuardrails) { + if (commonGuardrailParams != null) { + var newCommonParams = GuardrailRequestParams.builder() + .chatMemory(getMemory()) + .augmentationResult(commonGuardrailParams.augmentationResult()) + .userMessageTemplate(commonGuardrailParams.userMessageTemplate()) + .variables(commonGuardrailParams.variables()) + .build(); + + var outputGuardrailParams = OutputGuardrailRequest.builder() + .responseFromLLM(finalChatResponse) + .chatExecutor(chatExecutor) + .requestParams(newCommonParams) + .build(); + + finalChatResponse = + context.guardrailService().executeGuardrails(methodKey, outputGuardrailParams); + } + + // If we have output guardrails, we should process all of the partial responses first before + // completing + responseBuffer.forEach(partialResponseHandler::accept); + responseBuffer.clear(); + } + + // TODO should completeResponseHandler accept all ChatResponses that happened? completeResponseHandler.accept(finalChatResponse); } } } + private ChatMemory getMemory() { + return getMemory(memoryId); + } + + private ChatMemory getMemory(Object memId) { + return context.hasChatMemory() ? context.chatMemoryService.getOrCreateChatMemory(memoryId) : temporaryMemory; + } + private void addToMemory(ChatMessage chatMessage) { - if (context.hasChatMemory()) { - context.chatMemoryService.getOrCreateChatMemory(memoryId).add(chatMessage); - } else { - temporaryMemory.add(chatMessage); - } + getMemory().add(chatMessage); } private List messagesToSend(Object memoryId) { - return context.hasChatMemory() - ? context.chatMemoryService.getOrCreateChatMemory(memoryId).messages() - : temporaryMemory; + return getMemory(memoryId).messages(); } @Override @@ -155,11 +205,11 @@ class AiServiceStreamingResponseHandler implements StreamingChatResponseHandler try { errorHandler.accept(error); } catch (Exception e) { - log.error("While handling the following error...", error); - log.error("...the following error happened", e); + LOG.error("While handling the following error...", error); + LOG.error("...the following error happened", e); } } else { - log.warn("Ignored error", error); + LOG.warn("Ignored error", error); } } } diff --git a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStream.java b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStream.java index 4fbe64c653..f3d1dadf90 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStream.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStream.java @@ -1,25 +1,25 @@ package dev.langchain4j.service; -import dev.langchain4j.Internal; -import dev.langchain4j.agent.tool.ToolSpecification; -import dev.langchain4j.data.message.ChatMessage; -import dev.langchain4j.model.chat.request.ChatRequest; -import dev.langchain4j.model.chat.response.ChatResponse; -import dev.langchain4j.model.chat.response.StreamingChatResponseHandler; -import dev.langchain4j.model.output.TokenUsage; -import dev.langchain4j.rag.content.Content; -import dev.langchain4j.service.tool.ToolExecution; -import dev.langchain4j.service.tool.ToolExecutor; - -import java.util.ArrayList; -import java.util.List; -import java.util.Map; -import java.util.function.Consumer; - import static dev.langchain4j.internal.Utils.copy; import static dev.langchain4j.internal.ValidationUtils.ensureNotEmpty; import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; -import static java.util.Collections.emptyList; + +import dev.langchain4j.Internal; +import dev.langchain4j.agent.tool.ToolSpecification; +import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.guardrail.GuardrailRequestParams; +import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.memory.chat.MessageWindowChatMemory; +import dev.langchain4j.model.chat.ChatExecutor; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.model.output.TokenUsage; +import dev.langchain4j.rag.content.Content; +import dev.langchain4j.service.tool.ToolExecution; +import dev.langchain4j.service.tool.ToolExecutor; +import java.util.List; +import java.util.Map; +import java.util.function.Consumer; @Internal public class AiServiceTokenStream implements TokenStream { @@ -30,6 +30,8 @@ public class AiServiceTokenStream implements TokenStream { private final List retrievedContents; private final AiServiceContext context; private final Object memoryId; + private final GuardrailRequestParams commonGuardrailParams; + private final Object methodKey; private Consumer partialResponseHandler; private Consumer> contentsHandler; @@ -50,6 +52,7 @@ public class AiServiceTokenStream implements TokenStream { * @param parameters the parameters for creating the token stream */ public AiServiceTokenStream(AiServiceTokenStreamParameters parameters) { + ensureNotNull(parameters, "parameters"); this.messages = copy(ensureNotEmpty(parameters.messages(), "messages")); this.toolSpecifications = copy(parameters.toolSpecifications()); this.toolExecutors = copy(parameters.toolExecutors()); @@ -57,6 +60,8 @@ public class AiServiceTokenStream implements TokenStream { this.context = ensureNotNull(parameters.context(), "context"); ensureNotNull(this.context.streamingChatModel, "streamingChatModel"); this.memoryId = ensureNotNull(parameters.memoryId(), "memoryId"); + this.commonGuardrailParams = parameters.commonGuardrailParams(); + this.methodKey = parameters.methodKey(); } @Override @@ -110,7 +115,13 @@ public class AiServiceTokenStream implements TokenStream { .toolSpecifications(toolSpecifications) .build(); - StreamingChatResponseHandler handler = new AiServiceStreamingResponseHandler( + ChatExecutor chatExecutor = ChatExecutor.builder(context.streamingChatModel) + .errorHandler(errorHandler) + .chatRequest(chatRequest) + .build(); + + var handler = new AiServiceStreamingResponseHandler( + chatExecutor, context, memoryId, partialResponseHandler, @@ -120,7 +131,9 @@ public class AiServiceTokenStream implements TokenStream { initTemporaryMemory(context, messages), new TokenUsage(), toolSpecifications, - toolExecutors); + toolExecutors, + commonGuardrailParams, + methodKey); if (contentsHandler != null && retrievedContents != null) { contentsHandler.accept(retrievedContents); @@ -148,11 +161,13 @@ public class AiServiceTokenStream implements TokenStream { } } - private List initTemporaryMemory(AiServiceContext context, List messagesToSend) { - if (context.hasChatMemory()) { - return emptyList(); - } else { - return new ArrayList<>(messagesToSend); + private ChatMemory initTemporaryMemory(AiServiceContext context, List messagesToSend) { + var chatMemory = MessageWindowChatMemory.withMaxMessages(Integer.MAX_VALUE); + + if (!context.hasChatMemory()) { + chatMemory.add(messagesToSend); } + + return chatMemory; } } diff --git a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStreamParameters.java b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStreamParameters.java index 149815f482..0de530e4fb 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStreamParameters.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/AiServiceTokenStreamParameters.java @@ -1,11 +1,14 @@ package dev.langchain4j.service; +import static dev.langchain4j.internal.Utils.copyIfNotNull; +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + import dev.langchain4j.Internal; import dev.langchain4j.agent.tool.ToolSpecification; import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.guardrail.GuardrailRequestParams; import dev.langchain4j.rag.content.Content; import dev.langchain4j.service.tool.ToolExecutor; - import java.util.List; import java.util.Map; @@ -21,14 +24,19 @@ public class AiServiceTokenStreamParameters { private final List retrievedContents; private final AiServiceContext context; private final Object memoryId; + private final GuardrailRequestParams commonGuardrailParams; + private final Object methodKey; protected AiServiceTokenStreamParameters(Builder builder) { this.messages = builder.messages; - this.toolSpecifications = builder.toolSpecifications; - this.toolExecutors = builder.toolExecutors; + this.toolSpecifications = copyIfNotNull(builder.toolSpecifications); + this.toolExecutors = copyIfNotNull(builder.toolExecutors); this.retrievedContents = builder.retrievedContents; - this.context = builder.context; - this.memoryId = builder.memoryId; + this.context = ensureNotNull(builder.context, "context"); + this.memoryId = ensureNotNull(builder.memoryId, "memoryId"); + this.commonGuardrailParams = builder.commonGuardrailParams; + this.methodKey = builder.methodKey; + ensureNotNull(context.streamingChatModel, "streamingChatModel"); } /** @@ -73,6 +81,26 @@ public class AiServiceTokenStreamParameters { return memoryId; } + /** + * Retrieves the common parameters shared across guardrail checks for validating interactions + * between a user and a language model, if available. + * + * @return the {@link GuardrailRequestParams} containing chat memory, user message template, + * and additional variables required for guardrail processing, or null if not set. + */ + public GuardrailRequestParams commonGuardrailParams() { + return commonGuardrailParams; + } + + /** + * Retrieves the method key associated with this instance. + * + * @return the method key as an Object + */ + public Object methodKey() { + return methodKey; + } + /** * Creates a new builder for {@link AiServiceTokenStreamParameters}. * @@ -93,9 +121,10 @@ public class AiServiceTokenStreamParameters { private List retrievedContents; private AiServiceContext context; private Object memoryId; + private GuardrailRequestParams commonGuardrailParams; + private Object methodKey; - protected Builder() { - } + protected Builder() {} /** * Sets the messages. @@ -163,6 +192,30 @@ public class AiServiceTokenStreamParameters { return this; } + /** + * Sets the common guardrail parameters for validating interactions between a user and a language model. + * + * @param commonGuardrailParams an instance of {@link GuardrailRequestParams} containing the shared parameters + * required for guardrail checks, such as chat memory, user message template, + * and additional variables. + * @return this builder instance. + */ + public Builder commonGuardrailParams(GuardrailRequestParams commonGuardrailParams) { + this.commonGuardrailParams = commonGuardrailParams; + return this; + } + + /** + * Sets the method key. + * + * @param methodKey the method key + * @return this builder + */ + public Builder methodKey(Object methodKey) { + this.methodKey = methodKey; + return this; + } + /** * Builds a new {@link AiServiceTokenStreamParameters}. * diff --git a/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java b/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java index 735c823063..c6dd4966ba 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/AiServices.java @@ -12,6 +12,10 @@ import dev.langchain4j.agent.tool.ToolSpecification; import dev.langchain4j.data.message.AiMessage; import dev.langchain4j.data.message.ChatMessage; import dev.langchain4j.data.message.ToolExecutionResultMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.config.InputGuardrailsConfig; +import dev.langchain4j.guardrail.config.OutputGuardrailsConfig; import dev.langchain4j.memory.ChatMemory; import dev.langchain4j.memory.chat.ChatMemoryProvider; import dev.langchain4j.model.chat.ChatModel; @@ -161,9 +165,7 @@ public abstract class AiServices { * @return An instance of the provided interface, implementing all its defined methods. */ public static T create(Class aiService, StreamingChatModel streamingChatModel) { - return builder(aiService) - .streamingChatModel(streamingChatModel) - .build(); + return builder(aiService).streamingChatModel(streamingChatModel).build(); } /** @@ -398,6 +400,286 @@ public abstract class AiServices { return this; } + /** + * Configures the input guardrails for the AI service context by setting the provided InputGuardrailsConfig. + * + * @param inputGuardrailsConfig the configuration object that defines input guardrails for the AI service + * @return the current instance of {@link AiServices} to allow method chaining + */ + public AiServices inputGuardrailsConfig(InputGuardrailsConfig inputGuardrailsConfig) { + context.guardrailServiceBuilder.inputGuardrailsConfig(inputGuardrailsConfig); + return this; + } + + /** + * Configures the output guardrails for AI services. + * + * @param outputGuardrailsConfig the configuration object specifying the output guardrails + * @return the current instance of {@link AiServices} to allow for method chaining + */ + public AiServices outputGuardrailsConfig(OutputGuardrailsConfig outputGuardrailsConfig) { + context.guardrailServiceBuilder.outputGuardrailsConfig(outputGuardrailsConfig); + return this; + } + + /** + * Configures the input guardrail classes for the AI services. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.InputGuardrails InptputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * An input guardrail is a rule that is applied to the input of the model (essentially the user message) to ensure + * that the input is safe and meets the expectations of the model. It does not replace a moderation model, but it can + * be used to add additional checks (i.e. prompt injection, etc). + *

+ *

+ * Unlike for output guardrails, the input guardrails do not support retry or reprompt. The failure is passed directly + * to the caller, wrapped into a {@link dev.langchain4j.guardrail.GuardrailException GuardrailException}. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in the order + * they are listed. + *

+ * + * @param guardrailClasses A list of {@link InputGuardrailsConfig} classes, which will be used for input validation. + * The list can be {@code null} if no guardrails are to be applied. + * @param The type of {@link InputGuardrail} + * @return The instance of {@link AiServices} to allow method chaining. + */ + public AiServices inputGuardrailClasses(List> guardrailClasses) { + context.guardrailServiceBuilder.inputGuardrailClasses(guardrailClasses); + return this; + } + + /** + * Configures input guardrail classes for the AI service. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.InputGuardrails InptputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * An input guardrail is a rule that is applied to the input of the model (essentially the user message) to ensure + * that the input is safe and meets the expectations of the model. It does not replace a moderation model, but it can + * be used to add additional checks (i.e. prompt injection, etc). + *

+ *

+ * Unlike for output guardrails, the input guardrails do not support retry or reprompt. The failure is passed directly + * to the caller, wrapped into a {@link dev.langchain4j.guardrail.GuardrailException GuardrailException}. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in the order + * they are listed. + *

+ * + * @param guardrailClasses A list of {@link InputGuardrail} classes, which + * can include {@code null} to indicate no guardrails or optional configurations. + * @param The type of {@link InputGuardrail} + * @return the current instance of {@link AiServices} for chaining further configurations. + */ + public AiServices inputGuardrailClasses(Class... guardrailClasses) { + context.guardrailServiceBuilder.inputGuardrailClasses(guardrailClasses); + return this; + } + + /** + * Sets the input guardrails to be used by the guardrail service builder in the current context. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.InputGuardrails InptputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * An input guardrail is a rule that is applied to the input of the model (essentially the user message) to ensure + * that the input is safe and meets the expectations of the model. It does not replace a moderation model, but it can + * be used to add additional checks (i.e. prompt injection, etc). + *

+ *

+ * Unlike for output guardrails, the input guardrails do not support retry or reprompt. The failure is passed directly + * to the caller, wrapped into a {@link dev.langchain4j.guardrail.GuardrailException GuardrailException}. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in the order + * they are listed. + *

+ * + * @param guardrails a list of input guardrails, or null if no guardrails are to be set + * @return the current instance of {@link AiServices} for method chaining + */ + public AiServices inputGuardrails(List guardrails) { + context.guardrailServiceBuilder.inputGuardrails(guardrails); + return this; + } + + /** + * Adds the specified input guardrails to the context's guardrail service builder. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.InputGuardrails InptputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * An input guardrail is a rule that is applied to the input of the model (essentially the user message) to ensure + * that the input is safe and meets the expectations of the model. It does not replace a moderation model, but it can + * be used to add additional checks (i.e. prompt injection, etc). + *

+ *

+ * Unlike for output guardrails, the input guardrails do not support retry or reprompt. The failure is passed directly + * to the caller, wrapped into a {@link dev.langchain4j.guardrail.GuardrailException GuardrailException}. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in the order + * they are listed. + *

+ * + * @param guardrails an array of input guardrails to set, may be null + * @return the current instance of {@link AiServices} for chaining + */ + public AiServices inputGuardrails(I... guardrails) { + context.guardrailServiceBuilder.inputGuardrails(guardrails); + return this; + } + + /** + * Configures the output guardrail classes for the AI services. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.OutputGuardrails OutputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * Am output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * certain expectations. + *

+ *

+ * When a validation fails, the result can indicate whether the request should be retried as-is, or to provide a + * {@code reprompt} message to append to the prompt. + *

+ *

+ * In the case of re-prompting, the reprompt message is added to the LLM context and the request is then retried. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in + * the order they are listed. + *

+ *

+ * When several {@link OutputGuardrail}s are applied, if any guardrail forces a retry or reprompt, then all of the + * guardrails will be re-applied to the new response. + *

+ * + * @param guardrailClasses a list of {@link OutputGuardrail} classes. These classes + * define the output guardrails to be applied. Can be null. + * @param The type of {@link OutputGuardrail} + * @return the current instance of {@link AiServices}. + */ + public AiServices outputGuardrailClasses(List> guardrailClasses) { + context.guardrailServiceBuilder.outputGuardrailClasses(guardrailClasses); + return this; + } + + /** + * Sets the output guardrail classes to be used in the guardrail service. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.OutputGuardrails OutputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * Am output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * certain expectations. + *

+ *

+ * When a validation fails, the result can indicate whether the request should be retried as-is, or to provide a + * {@code reprompt} message to append to the prompt. + *

+ *

+ * In the case of re-prompting, the reprompt message is added to the LLM context and the request is then retried. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in + * the order they are listed. + *

+ *

+ * When several {@link OutputGuardrail}s are applied, if any guardrail forces a retry or reprompt, then all of the + * guardrails will be re-applied to the new response. + *

+ * + * @param guardrailClasses A list of {@link OutputGuardrail} classes. + * These classes define the guardrails for output behavior. + * Nullable, meaning guardrails can be omitted. + * @param The type of {@link OutputGuardrail} + * @return The current instance of {@link AiServices}, enabling method chaining. + */ + public AiServices outputGuardrailClasses(Class... guardrailClasses) { + context.guardrailServiceBuilder.outputGuardrailClasses(guardrailClasses); + return this; + } + + /** + * Configures the output guardrails for the AI service. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.OutputGuardrails OutputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * Am output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * certain expectations. + *

+ *

+ * When a validation fails, the result can indicate whether the request should be retried as-is, or to provide a + * {@code reprompt} message to append to the prompt. + *

+ *

+ * In the case of re-prompting, the reprompt message is added to the LLM context and the request is then retried. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in + * the order they are listed. + *

+ *

+ * When several {@link OutputGuardrail}s are applied, if any guardrail forces a retry or reprompt, then all of the + * guardrails will be re-applied to the new response. + *

+ * + * @param guardrails a list of output guardrails to be applied; can be {@code null} + * @return the current instance of {@link AiServices} for method chaining + */ + public AiServices outputGuardrails(List guardrails) { + context.guardrailServiceBuilder.outputGuardrails(guardrails); + return this; + } + + /** + * Configures output guardrails for the AI services. + *

+ * Configuring this way is exactly the same as using the {@link dev.langchain4j.service.guardrail.OutputGuardrails OutputGuardrails} + * annotation at the class level. Using the annotation takes precedence. + *

+ *

+ * Am output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * certain expectations. + *

+ *

+ * When a validation fails, the result can indicate whether the request should be retried as-is, or to provide a + * {@code reprompt} message to append to the prompt. + *

+ *

+ * In the case of re-prompting, the reprompt message is added to the LLM context and the request is then retried. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in + * the order they are listed. + *

+ *

+ * When several {@link OutputGuardrail}s are applied, if any guardrail forces a retry or reprompt, then all of the + * guardrails will be re-applied to the new response. + *

+ * + * @param guardrails an array of output guardrails to be applied; can be {@code null} + * or contain multiple instances of OutputGuardrail + * @return the current instance of {@link AiServices} with the specified guardrails applied + */ + public AiServices outputGuardrails(O... guardrails) { + context.guardrailServiceBuilder.outputGuardrails(guardrails); + return this; + } + /** * Constructs and returns the AI Service. * diff --git a/langchain4j/src/main/java/dev/langchain4j/service/DefaultAiServices.java b/langchain4j/src/main/java/dev/langchain4j/service/DefaultAiServices.java index c1eaeaf5ae..345581688c 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/DefaultAiServices.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/DefaultAiServices.java @@ -12,7 +12,11 @@ import dev.langchain4j.Internal; import dev.langchain4j.data.message.ChatMessage; import dev.langchain4j.data.message.SystemMessage; import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.GuardrailRequestParams; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailRequest; import dev.langchain4j.memory.ChatMemory; +import dev.langchain4j.model.chat.ChatExecutor; import dev.langchain4j.model.chat.request.ChatRequest; import dev.langchain4j.model.chat.request.ChatRequestParameters; import dev.langchain4j.model.chat.request.ResponseFormat; @@ -21,9 +25,11 @@ import dev.langchain4j.model.chat.response.ChatResponse; import dev.langchain4j.model.input.Prompt; import dev.langchain4j.model.input.PromptTemplate; import dev.langchain4j.model.moderation.Moderation; +import dev.langchain4j.model.output.FinishReason; import dev.langchain4j.rag.AugmentationRequest; import dev.langchain4j.rag.AugmentationResult; import dev.langchain4j.rag.query.Metadata; +import dev.langchain4j.service.guardrail.GuardrailService; import dev.langchain4j.service.memory.ChatMemoryAccess; import dev.langchain4j.service.memory.ChatMemoryService; import dev.langchain4j.service.output.ServiceOutputParser; @@ -148,7 +154,10 @@ class DefaultAiServices extends AiServices { : null; Optional systemMessage = prepareSystemMessage(memoryId, method, args); - UserMessage userMessage = prepareUserMessage(method, args); + var userMessageTemplate = getUserMessageTemplate(method, args); + var variables = InternalReflectionVariableResolver.findTemplateVariables( + userMessageTemplate, method, args); + UserMessage userMessage = prepareUserMessage(method, args, userMessageTemplate, variables); AugmentationResult augmentationResult = null; if (context.retrievalAugmentor != null) { List chatMemoryMessages = chatMemory != null ? chatMemory.messages() : null; @@ -158,9 +167,24 @@ class DefaultAiServices extends AiServices { userMessage = (UserMessage) augmentationResult.chatMessage(); } + var commonGuardrailParam = GuardrailRequestParams.builder() + .chatMemory(chatMemory) + .augmentationResult(augmentationResult) + .userMessageTemplate(userMessageTemplate) + .variables(variables) + .build(); + + // Invoke input guardrails + userMessage = invokeInputGuardrails( + context.guardrailService(), method, userMessage, commonGuardrailParam); + + // TODO give user ability to provide custom OutputParser Type returnType = method.getGenericReturnType(); boolean streaming = returnType == TokenStream.class || canAdaptTokenStreamTo(returnType); - boolean supportsJsonSchema = supportsJsonSchema(); + + boolean supportsJsonSchema = supportsJsonSchema(); // TODO should it be called for + // returnType==String? + Optional jsonSchema = Optional.empty(); if (supportsJsonSchema && !streaming) { jsonSchema = serviceOutputParser.jsonSchema(returnType); @@ -169,13 +193,13 @@ class DefaultAiServices extends AiServices { userMessage = appendOutputFormatInstructions(returnType, userMessage); } - List messages; - if (chatMemory != null) { + List messages = new ArrayList<>(); + + if (context.hasChatMemory()) { systemMessage.ifPresent(chatMemory::add); chatMemory.add(userMessage); - messages = chatMemory.messages(); + messages.addAll(chatMemory.messages()); } else { - messages = new ArrayList<>(); systemMessage.ifPresent(messages::add); messages.add(userMessage); } @@ -186,7 +210,7 @@ class DefaultAiServices extends AiServices { context.toolService.createContext(memoryId, userMessage); if (streaming) { - TokenStream tokenStream = new AiServiceTokenStream(AiServiceTokenStreamParameters.builder() + var tokenStreamParameters = AiServiceTokenStreamParameters.builder() .messages(messages) .toolSpecifications(toolServiceContext.toolSpecifications()) .toolExecutors(toolServiceContext.toolExecutors()) @@ -194,7 +218,11 @@ class DefaultAiServices extends AiServices { augmentationResult != null ? augmentationResult.contents() : null) .context(context) .memoryId(memoryId) - .build()); + .commonGuardrailParams(commonGuardrailParam) + .methodKey(method) + .build(); + + TokenStream tokenStream = new AiServiceTokenStream(tokenStreamParameters); // TODO moderation if (returnType == TokenStream.class) { return tokenStream; @@ -221,7 +249,11 @@ class DefaultAiServices extends AiServices { .parameters(parameters) .build(); - ChatResponse chatResponse = context.chatModel.chat(chatRequest); + ChatExecutor chatExecutor = ChatExecutor.builder(context.chatModel) + .chatRequest(chatRequest) + .build(); + + ChatResponse chatResponse = chatExecutor.execute(); verifyModerationIfNeeded(moderationFuture); @@ -236,13 +268,22 @@ class DefaultAiServices extends AiServices { chatResponse = toolServiceResult.chatResponse(); - Object parsedResponse = serviceOutputParser.parse(chatResponse, returnType); + FinishReason finishReason = chatResponse.metadata().finishReason(); + var response = invokeOutputGuardrails( + context.guardrailService(), method, chatResponse, chatExecutor, commonGuardrailParam); + + if ((response != null) && typeHasRawClass(returnType, response.getClass())) { + return response; + } + + var parsedResponse = serviceOutputParser.parse((ChatResponse) response, returnType); + if (typeHasRawClass(returnType, Result.class)) { return Result.builder() .content(parsedResponse) .tokenUsage(chatResponse.tokenUsage()) .sources(augmentationResult == null ? null : augmentationResult.contents()) - .finishReason(chatResponse.finishReason()) + .finishReason(finishReason) .toolExecutions(toolServiceResult.toolExecutions()) .build(); } else { @@ -300,6 +341,43 @@ class DefaultAiServices extends AiServices { return (T) proxyInstance; } + private UserMessage invokeInputGuardrails( + GuardrailService guardrailService, + Method method, + UserMessage userMessage, + GuardrailRequestParams commonGuardrailParams) { + + // NOTE: This check is cached, so it really only needs to be computed the first time for each method + if (guardrailService.hasInputGuardrails(method)) { + var inputGuardrailRequest = InputGuardrailRequest.builder() + .userMessage(userMessage) + .commonParams(commonGuardrailParams) + .build(); + return guardrailService.executeGuardrails(method, inputGuardrailRequest); + } + + return userMessage; + } + + private T invokeOutputGuardrails( + GuardrailService guardrailService, + Method method, + ChatResponse responseFromLLM, + ChatExecutor chatExecutor, + GuardrailRequestParams commonGuardrailParams) { + + if (guardrailService.hasOutputGuardrails(method)) { + var outputGuardrailRequest = OutputGuardrailRequest.builder() + .responseFromLLM(responseFromLLM) + .chatExecutor(chatExecutor) + .requestParams(commonGuardrailParams) + .build(); + return guardrailService.executeGuardrails(method, outputGuardrailRequest); + } + + return (T) responseFromLLM; + } + private Optional prepareSystemMessage(Object memoryId, Method method, Object[] args) { return findSystemMessageTemplate(memoryId, method).map(systemMessageTemplate -> PromptTemplate.from( systemMessageTemplate) @@ -318,13 +396,9 @@ class DefaultAiServices extends AiServices { return context.systemMessageProvider.apply(memoryId); } - private static UserMessage prepareUserMessage(Method method, Object[] args) { - - String template = getUserMessageTemplate(method, args); - Map variables = - InternalReflectionVariableResolver.findTemplateVariables(template, method, args); - - Prompt prompt = PromptTemplate.from(template).apply(variables); + private static UserMessage prepareUserMessage( + Method method, Object[] args, String userMessageTemplate, Map variables) { + Prompt prompt = PromptTemplate.from(userMessageTemplate).apply(variables); Optional maybeUserName = findUserName(method.getParameters(), args); return maybeUserName diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/AbstractGuardrailService.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/AbstractGuardrailService.java new file mode 100644 index 0000000000..b98235d2dc --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/AbstractGuardrailService.java @@ -0,0 +1,112 @@ +package dev.langchain4j.service.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.Internal; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailExecutor; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailExecutor; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +/** + * Responsible for managing and applying input and output guardrails to methods + * of a specified AI service class. Guardrails are defined through annotations at either the + * class or method level and are used to enforce constraints, validation, or transformation + * logic on inputs and outputs of AI service methods. + * + * This class handles the initialization, configuration, and execution of input and output + * guardrail logic for the methods of the associated AI service class, ensuring both input + * and output constraints are applied automatically when methods are invoked. + */ +@Internal +public abstract class AbstractGuardrailService implements GuardrailService { + private final Class aiServiceClass; + private final Map inputGuardrails = new HashMap<>(); + private final Map outputGuardrails = new HashMap<>(); + + // Caches for whether or not a method has input or output guardrails + private final Map inputGuardrailMethods = new HashMap<>(); + private final Map outputGuardrailMethods = new HashMap<>(); + + protected AbstractGuardrailService( + Class aiServiceClass, + Map inputGuardrails, + Map outputGuardrails) { + this.aiServiceClass = ensureNotNull(aiServiceClass, "aiServiceClass"); + Optional.ofNullable(inputGuardrails).ifPresent(this.inputGuardrails::putAll); + Optional.ofNullable(outputGuardrails).ifPresent(this.outputGuardrails::putAll); + } + + @Override + public Class aiServiceClass() { + return this.aiServiceClass; + } + + @Override + public InputGuardrailResult executeInputGuardrails(MethodKey method, InputGuardrailRequest params) { + return Optional.ofNullable(method) + .map(this.inputGuardrails::get) + .map(executor -> executor.execute(params)) + .orElseGet(InputGuardrailResult::success); + } + + @Override + public OutputGuardrailResult executeOutputGuardrails(MethodKey method, OutputGuardrailRequest params) { + return Optional.ofNullable(method) + .map(this.outputGuardrails::get) + .map(executor -> executor.execute(params)) + .orElseGet(OutputGuardrailResult::success); + } + + @Override + public boolean hasInputGuardrails(MethodKey method) { + return this.inputGuardrailMethods.computeIfAbsent( + method, m -> !getInputGuardrails(m).isEmpty()); + } + + @Override + public boolean hasOutputGuardrails(MethodKey method) { + return this.outputGuardrailMethods.computeIfAbsent( + method, m -> !getOutputGuardrails(m).isEmpty()); + } + + // These methods below really only exist for testing purposes + // That's why they are package-scoped + int getInputGuardrailMethodCount() { + return this.inputGuardrails.size(); + } + + int getOutputGuardrailMethodCount() { + return this.outputGuardrails.size(); + } + + Optional getInputConfig(MethodKey method) { + return Optional.ofNullable(this.inputGuardrails.get(method)).map(InputGuardrailExecutor::config); + } + + Optional getOutputConfig(MethodKey method) { + return Optional.ofNullable(this.outputGuardrails.get(method)).map(OutputGuardrailExecutor::config); + } + + List getInputGuardrails(MethodKey method) { + return Optional.ofNullable(method) + .map(this.inputGuardrails::get) + .map(InputGuardrailExecutor::guardrails) + .orElseGet(List::of); + } + + List getOutputGuardrails(MethodKey method) { + return Optional.ofNullable(method) + .map(this.outputGuardrails::get) + .map(OutputGuardrailExecutor::guardrails) + .orElseGet(List::of); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/DefaultGuardrailService.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/DefaultGuardrailService.java new file mode 100644 index 0000000000..a9ed2613db --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/DefaultGuardrailService.java @@ -0,0 +1,58 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailExecutor; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailExecutor; +import java.lang.reflect.Method; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.stream.Stream; + +/** + * Responsible for managing and applying input and output guardrails + * to methods of a specified AI service class. Guardrails are defined through annotations + * at either the class or method level and are used to enforce constraints or rules for + * processing requests and responses. + * + * This class initializes guardrails for all methods of the specified AI service class, + * allowing input and output validation, transformation, or restriction through the + * specified guardrail implementations. The guardrails can be customized through configurations + * specific to each guardrail type. + *

+ * Obtain instances via {@link GuardrailService#builder(Class)} + *

+ */ +final class DefaultGuardrailService extends AbstractGuardrailService { + DefaultGuardrailService( + Class aiServiceClass, + Map inputGuardrails, + Map outputGuardrails) { + super(aiServiceClass, inputGuardrails, outputGuardrails); + } + + // These methods below really only exist for testing purposes + // Thats why they are package-scoped + Optional getInputConfig(String methodName) { + return findMethod(methodName).flatMap(super::getInputConfig); + } + + Optional getOutputConfig(String methodName) { + return findMethod(methodName).flatMap(super::getOutputConfig); + } + + List getInputGuardrails(String methodName) { + return findMethod(methodName).map(super::getInputGuardrails).orElseGet(List::of); + } + + List getOutputGuardrails(String methodName) { + return findMethod(methodName).map(super::getOutputGuardrails).orElseGet(List::of); + } + + private Optional findMethod(String methodName) { + return Stream.of(aiServiceClass().getMethods()) + .filter(method -> methodName.equals(method.getName())) + .findFirst(); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailService.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailService.java new file mode 100644 index 0000000000..4fc0b15410 --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailService.java @@ -0,0 +1,232 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailRequest; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailRequest; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.model.chat.response.ChatResponse; +import java.lang.reflect.Method; +import java.util.List; +import java.util.Optional; + +/** + * Defines a service for executing guardrails associated with methods in an AI service. + * Guardrails are constraints or validations applied either to input or output of a method. + */ +public interface GuardrailService { + /** + * Retrieves the class representing the AI service to which the guardrails apply. + * + * @return The {@code Class} object representing the AI service. + */ + Class aiServiceClass(); + + /** + * Executes the input guardrails associated with a given {@link Method} + * + * @param method The method whose input guardrails are to be executed. + * @param params The parameters to validate against the input guardrails. Must not be null. + * @return The result of executing the input guardrails, encapsulated in an {@code InputGuardrailResult}. + * If no guardrails are associated with the method, a successful result is returned by default. + * @param > The type of the method key, representing a unique identifier for methods. + */ + InputGuardrailResult executeInputGuardrails(MethodKey method, InputGuardrailRequest params); + + /** + * Executes the input guardrails associated with the given method and parameters, + * and retrieves a modified or validated {@link UserMessage} based on the result. + * + * @param The type of the method key, representing a unique identifier for methods. + * @param method The method whose input guardrails are to be executed. Nullable. + * @param params The parameters to validate against the input guardrails. Must not be null. + * @return A {@link UserMessage} derived from the provided parameters and the result + * of the input guardrails execution. If guardrails are applied successfully, + * a potentially rewritten user message is returned. If no guardrails are + * associated with the method, the original user message is returned. + */ + default UserMessage executeGuardrails(MethodKey method, InputGuardrailRequest params) { + return executeInputGuardrails(method, params).userMessage(params); + } + + /** + * Executes the output guardrails associated with a given {@code Method}. + * + * @param method The method whose output guardrails are to be executed. + * @param params The parameters to validate against the output guardrails. Must not be null. + * @return The result of executing the output guardrails, encapsulated in an {@code OutputGuardrailResult}. + * If no guardrails are associated with the method, a successful result is returned by default. + * @param > The type of the method key, representing a unique identifier for methods. + */ + OutputGuardrailResult executeOutputGuardrails(MethodKey method, OutputGuardrailRequest params); + + /** + * Whether or not a method has any input guardrails associated with it + * @param method The method + * @return {@code true} If {@code method} has input guardrails. {@code false} otherwise + * @param > The type of the method key, representing a unique identifier for methods. + */ + boolean hasInputGuardrails(MethodKey method); + + /** + * Whether or not a method has any output guardrails associated with it + * @param method The method + * @return {@code true} If {@code method} has output guardrails. {@code false} otherwise + * @param > The type of the method key, representing a unique identifier for methods. + */ + boolean hasOutputGuardrails(MethodKey method); + + /** + * Executes the guardrails associated with a given method and parameters, returning the appropriate response. + * + * @param The type of the method key, representing a unique identifier for methods. + * @param The type of response to produce + * @param method The method whose output guardrails are to be executed. Nullable. + * @param params The parameters to validate against the output guardrails. Must not be null. + * @return A {@link ChatResponse} that encapsulates the output of executing the guardrails based on the provided parameters. + */ + default T executeGuardrails(MethodKey method, OutputGuardrailRequest params) { + return executeOutputGuardrails(method, params).response(params); + } + + /** + * Creates a new instance of {@link Builder} for the specified AI service class. + * + * @param aiServiceClass The {@code Class} object representing the AI service for which the builder is being created. + * @return A {@link Builder} instance initialized with the specified AI service class. + */ + static Builder builder(Class aiServiceClass) { + return new GuardrailServiceBuilder(aiServiceClass); + } + + interface Builder { + /** + * Configures the input guardrails for the builder. + * + * @param config The configuration for input guardrails. Must not be null. + * @return The current instance of {@link Builder} for method chaining. + * @throws IllegalArgumentException if {@code config} is null. + */ + Builder inputGuardrailsConfig(dev.langchain4j.guardrail.config.InputGuardrailsConfig config); + + /** + * Configures the output guardrails for the Builder. + * + * @param config The configuration for output guardrails. Must not be null. + * @return The current instance of {@link Builder} for method chaining. + * @throws IllegalArgumentException if {@code config} is null. + */ + Builder outputGuardrailsConfig(dev.langchain4j.guardrail.config.OutputGuardrailsConfig config); + + /** + * Configures the classes of input guardrails for the Builder. Existing input guardrail classes will be cleared. + * + * @param guardrailClasses A list of classes implementing the {@link InputGuardrail} interface to be used + * as input guardrails. May be {@code null}. + * @param The type of {@link InputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + Builder inputGuardrailClasses(List> guardrailClasses); + + /** + * Configures the classes of input guardrails for the Builder. + * Existing input guardrail classes will be cleared. + * + * @param guardrailClasses An array of classes implementing the {@link InputGuardrail} interface to be used + * as input guardrails. May be {@code null}. + * @param The type of {@link InputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + default Builder inputGuardrailClasses(Class... guardrailClasses) { + return Optional.ofNullable(guardrailClasses) + .map(g -> inputGuardrailClasses(List.of(g))) + .orElse(this); + } + + /** + * Configures the classes of output guardrails for the Builder. + * Existing output guardrail classes will be cleared. + * + * @param guardrailClasses A list of classes implementing the {@link OutputGuardrail} interface to be used + * as output guardrails. May be {@code null}. + * @param The type of {@link OutputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + Builder outputGuardrailClasses(List> guardrailClasses); + + /** + * Configures the classes of output guardrails for the Builder. + * Existing output guardrail classes will be cleared. + * + * @param guardrailClasses An array of classes implementing the {@link OutputGuardrail} interface to be used + * as output guardrails. May be {@code null}. + * @param The type of {@link OutputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + default Builder outputGuardrailClasses(Class... guardrailClasses) { + return Optional.ofNullable(guardrailClasses) + .map(g -> outputGuardrailClasses(List.of(g))) + .orElse(this); + } + + /** + * Sets the input guardrails for the Builder. Existing input guardrails + * will be cleared, and the provided input guardrails will be added. + * + * @param guardrails A list of input guardrails implementing the {@link InputGuardrail} interface. + * Can be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + Builder inputGuardrails(List guardrails); + + /** + * Configures the input guardrails for the Builder. + * + * @param guardrails An array of input guardrails implementing the {@link InputGuardrail} interface. + * May be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + default Builder inputGuardrails(I... guardrails) { + return Optional.ofNullable(guardrails) + .map(ig -> inputGuardrails(List.of(ig))) + .orElse(this); + } + + /** + * Sets the output guardrails for the Builder. Existing output guardrails + * will be cleared, and the provided output guardrails will be added. + * + * @param guardrails A list of output guardrails implementing the {@link OutputGuardrail} + * interface. Can be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + Builder outputGuardrails(List guardrails); + + /** + * Configures the output guardrails for the Builder. + * + * @param guardrails An array of output guardrails implementing the {@link OutputGuardrail} interface. + * May be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + default Builder outputGuardrails(O... guardrails) { + return Optional.ofNullable(guardrails) + .map(og -> outputGuardrails(List.of(og))) + .orElse(this); + } + + /** + * Builds and returns an instance of {@link GuardrailService}. + * This method configures input and output guardrails at the service level + * using the provided class-level or method-level annotations. If no + * method-level annotations are present, it defers to class-level annotations, + * and if those are absent, it uses the settings defined in the builder. + * + * @return an instance of {@link GuardrailService} configured with appropriate + * input and output guardrails. + */ + GuardrailService build(); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailServiceBuilder.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailServiceBuilder.java new file mode 100644 index 0000000000..725cd11bc5 --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/GuardrailServiceBuilder.java @@ -0,0 +1,328 @@ +package dev.langchain4j.service.guardrail; + +import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; + +import dev.langchain4j.classinstance.ClassInstanceLoader; +import dev.langchain4j.classloading.ClassMetadataProvider; +import dev.langchain4j.guardrail.Guardrail; +import dev.langchain4j.guardrail.GuardrailRequest; +import dev.langchain4j.guardrail.GuardrailResult; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailExecutor; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailExecutor; +import dev.langchain4j.service.guardrail.GuardrailService.Builder; +import dev.langchain4j.spi.classloading.ClassMetadataProviderFactory; +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.function.Supplier; +import java.util.stream.Collectors; +import java.util.stream.Stream; + +/** + * A builder class for creating and configuring a {@link GuardrailService} instance, which provides input and output guardrail + * mechanisms for an AI service class. This class allows customization of the guardrails through configuration objects, + * guardrail classes, or direct instances of guardrails. + *

+ * This builder supports setting guardrails via both annotations and explicit configurations in code. Annotations at the + * method level take precedence over annotations at the class level, and annotations take precedence over guardrails + * configured programmatically in the builder. + */ +final class GuardrailServiceBuilder implements Builder { + private final Supplier defaultInputGuardrailSupplier = + () -> InputGuardrailExecutor.builder() + .config(this.inputGuardrailsConfig) + .guardrails( + getNonAnnotationBasedClassLevelGuardrails(this.inputGuardrails, this.inputGuardrailClasses)) + .build(); + + private final Supplier defaultOutputGuardrailSupplier = + () -> OutputGuardrailExecutor.builder() + .config(this.outputGuardrailsConfig) + .guardrails(getNonAnnotationBasedClassLevelGuardrails( + this.outputGuardrails, this.outputGuardrailClasses)) + .build(); + + private final Class aiServiceClass; + private dev.langchain4j.guardrail.config.InputGuardrailsConfig inputGuardrailsConfig; + private dev.langchain4j.guardrail.config.OutputGuardrailsConfig outputGuardrailsConfig; + private List> inputGuardrailClasses = new ArrayList<>(); + private List> outputGuardrailClasses = new ArrayList<>(); + private List inputGuardrails = new ArrayList<>(); + private List outputGuardrails = new ArrayList<>(); + + GuardrailServiceBuilder(Class aiServiceClass) { + this.aiServiceClass = ensureNotNull(aiServiceClass, "aiServiceClass"); + } + + /** + * Configures the input guardrails for the Builder. + * + * @param config The configuration for input guardrails. Must not be null. + * @return The current instance of {@link Builder} for method chaining. + * @throws IllegalArgumentException if {@code config} is null. + */ + @Override + public Builder inputGuardrailsConfig(dev.langchain4j.guardrail.config.InputGuardrailsConfig config) { + this.inputGuardrailsConfig = ensureNotNull(config, "config"); + return this; + } + + /** + * Configures the output guardrails for the Builder. + * + * @param config The configuration for output guardrails. Must not be null. + * @return The current instance of {@link Builder} for method chaining. + * @throws IllegalArgumentException if {@code config} is null. + */ + @Override + public Builder outputGuardrailsConfig(dev.langchain4j.guardrail.config.OutputGuardrailsConfig config) { + this.outputGuardrailsConfig = ensureNotNull(config, "config"); + return this; + } + + /** + * Configures the classes of input guardrails for the Builder. Existing input guardrail classes will be cleared. + * + * @param guardrailClasses A list of classes implementing the {@link InputGuardrail} interface to be used + * as input guardrails. May be {@code null}. + * @param The type of {@link InputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + @Override + public Builder inputGuardrailClasses(List> guardrailClasses) { + this.inputGuardrailClasses.clear(); + + if (guardrailClasses != null) { + this.inputGuardrailClasses.addAll(guardrailClasses); + } + + return this; + } + + /** + * Configures the classes of output guardrails for the Builder. + * Existing output guardrail classes will be cleared. + * + * @param guardrailClasses A list of classes implementing the {@link OutputGuardrail} interface to be used + * as output guardrails. May be {@code null}. + * @param The type of {@link OutputGuardrail} + * @return The current instance of {@link Builder} for method chaining. + */ + @Override + public Builder outputGuardrailClasses(List> guardrailClasses) { + this.outputGuardrailClasses.clear(); + + if (guardrailClasses != null) { + this.outputGuardrailClasses.addAll(guardrailClasses); + } + + return this; + } + + /** + * Sets the input guardrails for the Builder. Existing input guardrails + * will be cleared, and the provided input guardrails will be added. + * + * @param guardrails A list of input guardrails implementing the {@link InputGuardrail} interface. + * Can be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + @Override + public Builder inputGuardrails(List guardrails) { + this.inputGuardrails.clear(); + + if (guardrails != null) { + this.inputGuardrails.addAll(guardrails); + } + + return this; + } + + /** + * Sets the output guardrails for the Builder. Existing output guardrails + * will be cleared, and the provided output guardrails will be added. + * + * @param guardrails A list of output guardrails implementing the {@link OutputGuardrail} + * interface. Can be {@code null}, in which case no guardrails will be added. + * @return The current instance of {@link Builder} for method chaining. + */ + @Override + public Builder outputGuardrails(List guardrails) { + this.outputGuardrails.clear(); + + if (guardrails != null) { + this.outputGuardrails.addAll(guardrails); + } + + return this; + } + + /** + * Builds and returns an instance of {@link GuardrailService}. + * This method configures input and output guardrails using the settings defined in the builder. + * If not set it then uses the provided class-level or method-level annotations. If no + * method-level annotations are present, it defers to class-level annotations. + * + * @return an instance of {@link GuardrailService} configured with appropriate + * input and output guardrails. + */ + @Override + public DefaultGuardrailService build() { + // Anything set here in this builder is relevant at the AiService level, NOT at the method level + // Setting guardrails at the method level can only be done via the annotations + + // Next, compute method-level guardrails based on the annotations + // Go method-by-method, if there are annotations on the method, then use them + // Otherwise use the annotations on the class + // If there aren't any annotations on the class, then use the ones set on this builder + var inputGuardrailsByMethod = new HashMap(); + var outputGuardrailsByMethod = new HashMap(); + var factory = ClassMetadataProvider.getClassMetadataProviderFactory(); + + factory.getNonStaticMethodsOnClass(this.aiServiceClass).forEach(method -> { + var inputGuardrailsForMethod = computeInputGuardrailsForAiServiceMethod(method, factory); + var outputGuardrailsForMethod = computeOutputGuardrailsForAiServiceMethod(method, factory); + + if (!inputGuardrailsForMethod.guardrails().isEmpty()) { + inputGuardrailsByMethod.put(method, inputGuardrailsForMethod); + } + + if (!outputGuardrailsForMethod.guardrails().isEmpty()) { + outputGuardrailsByMethod.put(method, outputGuardrailsForMethod); + } + }); + + return new DefaultGuardrailService(this.aiServiceClass, inputGuardrailsByMethod, outputGuardrailsByMethod); + } + + private static

, G extends Guardrail> + List getNonAnnotationBasedClassLevelGuardrails( + List guardrails, List> guardrailClasses) { + ensureNotNull(guardrails, "guardrails"); + ensureNotNull(guardrailClasses, "guardrailClasses"); + + var guardrailsSetByBuilderAtClassLevel = guardrails.stream(); + var guardrailsSetByBuilderAtClassLevelByClassName = + guardrailClasses.stream().map(GuardrailServiceBuilder::getGuardrailClassInstance); + + return Stream.concat(guardrailsSetByBuilderAtClassLevel, guardrailsSetByBuilderAtClassLevelByClassName) + .collect(Collectors.toCollection(ArrayList::new)); + } + + private static

, G extends Guardrail> + G getGuardrailClassInstance(Class guardrailClass) { + ensureNotNull(guardrailClass, "guardrailClass"); + return ClassInstanceLoader.getClassInstance(guardrailClass); + } + + private static List getGuardrails(InputGuardrails inputGuardrails) { + return Stream.of(inputGuardrails.value()) + .map(guardrailClass -> (I) getGuardrailClassInstance(guardrailClass)) + .toList(); + } + + private static List getGuardrails(OutputGuardrails outputGuardrails) { + return Stream.of(outputGuardrails.value()) + .map(guardrailClass -> (O) getGuardrailClassInstance(guardrailClass)) + .toList(); + } + + private static dev.langchain4j.guardrail.config.InputGuardrailsConfig computeConfig(InputGuardrails annotation) { + return dev.langchain4j.guardrail.config.InputGuardrailsConfig.builder().build(); + } + + private static dev.langchain4j.guardrail.config.OutputGuardrailsConfig computeConfig(OutputGuardrails annotation) { + + return dev.langchain4j.guardrail.config.OutputGuardrailsConfig.builder() + .maxRetries(annotation.maxRetries()) + .build(); + } + + private InputGuardrailExecutor computeInputGuardrails(InputGuardrails annotation) { + return InputGuardrailExecutor.builder() + .config(hasInputGuardrailConfigSetOnBuilder() ? this.inputGuardrailsConfig : computeConfig(annotation)) + .guardrails( + hasInputGuardrailsSetOnBuilder() + ? getNonAnnotationBasedClassLevelGuardrails( + this.inputGuardrails, this.inputGuardrailClasses) + : getGuardrails(annotation)) + .build(); + } + + private OutputGuardrailExecutor computeOutputGuardrails(OutputGuardrails annotation) { + return OutputGuardrailExecutor.builder() + .config( + hasOutputGuardrailConfigSetOnBuilder() + ? this.outputGuardrailsConfig + : computeConfig(annotation)) + .guardrails( + hasOutputGuardrailsSetOnBuilder() + ? getNonAnnotationBasedClassLevelGuardrails( + this.outputGuardrails, this.outputGuardrailClasses) + : getGuardrails(annotation)) + .build(); + } + + private InputGuardrailExecutor computeInputGuardrailsForAiServiceMethod( + MethodKey method, ClassMetadataProviderFactory factory) { + // For both input & output guardrails, first check the builder + // If nothing on the builder, then check the method + // If nothing on the method, then fall back to the class + if (inputGuardrailsAndConfigSetOnBuilder()) { + // Don't need to introspect the annotation at all since everything is set on the builder + return this.defaultInputGuardrailSupplier.get(); + } + + // If we get here we know we need to introspect the annotation for one reason or another + return factory.getAnnotation(method, InputGuardrails.class) + .map(this::computeInputGuardrails) + // Didn't exist at the method level, so check the class level + .orElseGet(() -> factory.getAnnotation(this.aiServiceClass, InputGuardrails.class) + .map(this::computeInputGuardrails) + .orElseGet(this.defaultInputGuardrailSupplier::get)); + } + + private OutputGuardrailExecutor computeOutputGuardrailsForAiServiceMethod( + MethodKey method, ClassMetadataProviderFactory factory) { + // For both input & output guardrails, first check the builder + // If nothing on the builder, then check the method + // If nothing on the method, then fall back to the class + if (outputGuardrailsAndConfigSetOnBuilder()) { + return this.defaultOutputGuardrailSupplier.get(); + } + + // If we get here we know we need to introspect the annotation for one reason or another + return factory.getAnnotation(method, OutputGuardrails.class) + .map(this::computeOutputGuardrails) + // Didn't exist at the method level, so check the class level + .orElseGet(() -> factory.getAnnotation(this.aiServiceClass, OutputGuardrails.class) + .map(this::computeOutputGuardrails) + .orElseGet(this.defaultOutputGuardrailSupplier::get)); + } + + private boolean hasInputGuardrailsSetOnBuilder() { + return !this.inputGuardrails.isEmpty() || !this.inputGuardrailClasses.isEmpty(); + } + + private boolean hasInputGuardrailConfigSetOnBuilder() { + return this.inputGuardrailsConfig != null; + } + + private boolean hasOutputGuardrailsSetOnBuilder() { + return !this.outputGuardrails.isEmpty() || !this.outputGuardrailClasses.isEmpty(); + } + + private boolean hasOutputGuardrailConfigSetOnBuilder() { + return this.outputGuardrailsConfig != null; + } + + private boolean inputGuardrailsAndConfigSetOnBuilder() { + return hasInputGuardrailsSetOnBuilder() && hasInputGuardrailConfigSetOnBuilder(); + } + + private boolean outputGuardrailsAndConfigSetOnBuilder() { + return hasOutputGuardrailsSetOnBuilder() && hasOutputGuardrailConfigSetOnBuilder(); + } +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java new file mode 100644 index 0000000000..01998869cb --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/InputGuardrails.java @@ -0,0 +1,41 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.guardrail.InputGuardrail; +import java.lang.annotation.Documented; +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * An annotation to apply input guardrails to the input of the model using the declarative {@link dev.langchain4j.service.AiServices AiServices} approach. + *

+ * An input guardrail is a rule that is applied to the input of the model (essentially the user message) to ensure + * that the input is safe and meets the expectations of the model. It does not replace a moderation model, but it can + * be used to add additional checks (i.e. prompt injection, etc). + *

+ *

+ * Unlike for output guardrails, the input guardrails do not support retry or reprompt. The failure is passed directly + * to the caller, wrapped into a {@link dev.langchain4j.guardrail.GuardrailException GuardrailException}. + *

+ *

+ * If the annotation is present on a class, the guardrails will be applied to all the methods of the class. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in the order + * they are listed. + *

+ */ +@Retention(RetentionPolicy.RUNTIME) +@Documented +@Target({ElementType.TYPE, ElementType.METHOD}) +public @interface InputGuardrails { + /** + * The ordered list of {@link InputGuardrail}s to apply to the input of the model. + *

+ * The order of the classes is important as the guardrails are applied in the order they are listed. + * Guardrails can not be present twice in the list. + *

+ */ + Class[] value(); +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java new file mode 100644 index 0000000000..c7906744e2 --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/OutputGuardrails.java @@ -0,0 +1,57 @@ +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.guardrail.OutputGuardrail; +import java.lang.annotation.Documented; +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * An annotation to apply guardrails to the output of the model using the declarative {@link dev.langchain4j.service.AiServices AiServices} + * approach. + *

+ * Am output guardrail is a rule that is applied to the output of the model to ensure that the output is safe and meets + * certain expectations. + *

+ *

+ * When a validation fails, the result can indicate whether the request should be retried as-is, or to provide a + * {@code reprompt} message to append to the prompt. + *

+ *

+ * In the case of re-prompting, the reprompt message is added to the LLM context and the request is then retried. + *

+ *

+ * If the annotation is present on a class, the guardrails will be applied to all the methods of the class. + *

+ *

+ * When several guardrails are applied, the order of the guardrails is important, as the guardrails are applied in + * the order they are listed. + *

+ *

+ * When several {@link OutputGuardrail}s are applied, if any guardrail forces a retry or reprompt, then all of the + * guardrails will be re-applied to the new response. + *

+ */ +@Documented +@Retention(RetentionPolicy.RUNTIME) +@Target({ElementType.TYPE, ElementType.METHOD}) +public @interface OutputGuardrails { + /** + * The ordered list of guardrails to apply to the output of the model. + *

+ * The order of the classes is important as the guardrails are applied in the order they are listed. + * Guardrails can not be present twice in the list. + *

+ */ + Class[] value(); + + /** + * The maximum number of retries to perform when an output guardrail forces a retry or reprompt. + *

+ * Set to {@code 0} to disable retries + *

+ * @see dev.langchain4j.guardrail.config.OutputGuardrailsConfig#maxRetries() + */ + int maxRetries() default dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT; +} diff --git a/langchain4j/src/main/java/dev/langchain4j/service/guardrail/package-info.java b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/package-info.java new file mode 100644 index 0000000000..2a3aafc6be --- /dev/null +++ b/langchain4j/src/main/java/dev/langchain4j/service/guardrail/package-info.java @@ -0,0 +1,4 @@ +@Experimental +package dev.langchain4j.service.guardrail; + +import dev.langchain4j.Experimental; diff --git a/langchain4j/src/main/java/dev/langchain4j/service/tool/ToolExecutionRequestUtil.java b/langchain4j/src/main/java/dev/langchain4j/service/tool/ToolExecutionRequestUtil.java index 4602d98c6a..918aaeb5de 100644 --- a/langchain4j/src/main/java/dev/langchain4j/service/tool/ToolExecutionRequestUtil.java +++ b/langchain4j/src/main/java/dev/langchain4j/service/tool/ToolExecutionRequestUtil.java @@ -1,17 +1,16 @@ package dev.langchain4j.service.tool; +import static dev.langchain4j.internal.Utils.isNullOrBlank; + import dev.langchain4j.Internal; import dev.langchain4j.agent.tool.ToolExecutionRequest; import dev.langchain4j.internal.Json; - import java.lang.reflect.ParameterizedType; import java.lang.reflect.Type; import java.util.Map; import java.util.regex.Matcher; import java.util.regex.Pattern; -import static dev.langchain4j.internal.Utils.isNullOrBlank; - /** * Utility class for {@link ToolExecutionRequest}. */ @@ -22,14 +21,13 @@ class ToolExecutionRequestUtil { private static final Pattern LEADING_TRAILING_QUOTE_PATTERN = Pattern.compile("^\"|\"$"); private static final Pattern ESCAPED_QUOTE_PATTERN = Pattern.compile("\\\\\""); - private ToolExecutionRequestUtil() { - } + private ToolExecutionRequestUtil() {} private static final Type MAP_TYPE = new ParameterizedType() { @Override public Type[] getActualTypeArguments() { - return new Type[]{String.class, Object.class}; + return new Type[] {String.class, Object.class}; } @Override @@ -93,5 +91,4 @@ class ToolExecutionRequestUtil { Matcher escapedQuoteMatcher = ESCAPED_QUOTE_PATTERN.matcher(normalizedJson); return escapedQuoteMatcher.replaceAll("\""); } - } diff --git a/langchain4j/src/test/java/dev/langchain4j/classinstance/ClassMetadataProviderTests.java b/langchain4j/src/test/java/dev/langchain4j/classinstance/ClassMetadataProviderTests.java new file mode 100644 index 0000000000..a564211e50 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/classinstance/ClassMetadataProviderTests.java @@ -0,0 +1,65 @@ +package dev.langchain4j.classinstance; + +import static org.assertj.core.api.Assertions.assertThat; + +import dev.langchain4j.Experimental; +import dev.langchain4j.classloading.ClassMetadataProvider; +import dev.langchain4j.classloading.ReflectionBasedClassMetadataProviderFactory; +import java.lang.annotation.Target; +import java.lang.reflect.Method; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; +import org.junit.jupiter.api.Test; + +class ClassMetadataProviderTests { + @Test + void loadsThingsCorrectly() { + var factory = ClassMetadataProvider.getClassMetadataProviderFactory(); + + assertThat(factory).isNotNull().isExactlyInstanceOf(ReflectionBasedClassMetadataProviderFactory.class); + + assertThat(factory.getAnnotation(SomeInterface.class, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is plain and boring!"); + + assertThat(factory.getAnnotation(SomeInterface.class, Target.class)).isEmpty(); + + var methods = factory.getNonStaticMethodsOnClass(SomeInterface.class); + + assertThat(methods).hasSize(2).extracting(Method::getName).containsExactlyInAnyOrder("hello", "goodbye"); + + var methodsByName = StreamSupport.stream(methods.spliterator(), false) + .collect(Collectors.toMap(Method::getName, method -> method)); + + var helloMethod = methodsByName.get("hello"); + var goodbyeMethod = methodsByName.get("goodbye"); + + assertThat(helloMethod).isNotNull(); + + assertThat(goodbyeMethod).isNotNull(); + + assertThat(factory.getAnnotation(goodbyeMethod, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("Just trying things out"); + + assertThat(factory.getAnnotation(goodbyeMethod, Target.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Experimental.class)).isEmpty(); + + assertThat(factory.getAnnotation(helloMethod, Target.class)).isEmpty(); + } + + @Experimental("This is plain and boring!") + interface SomeInterface { + String hello(); + + @Experimental("Just trying things out") + String goodbye(); + + static String wave() { + return "wave"; + } + } +} diff --git a/langchain4j/src/test/java/dev/langchain4j/classinstance/ReflectionBasedClassMetadataProviderFactoryTests.java b/langchain4j/src/test/java/dev/langchain4j/classinstance/ReflectionBasedClassMetadataProviderFactoryTests.java new file mode 100644 index 0000000000..71b10b7db8 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/classinstance/ReflectionBasedClassMetadataProviderFactoryTests.java @@ -0,0 +1,77 @@ +package dev.langchain4j.classinstance; + +import static org.assertj.core.api.Assertions.assertThat; + +import dev.langchain4j.Experimental; +import dev.langchain4j.classloading.ReflectionBasedClassMetadataProviderFactory; +import dev.langchain4j.service.TokenStream; +import dev.langchain4j.service.guardrail.InputGuardrails; +import java.lang.reflect.Method; +import java.util.Map; +import java.util.function.Function; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; +import org.junit.jupiter.api.Test; + +class ReflectionBasedClassMetadataProviderFactoryTests { + ReflectionBasedClassMetadataProviderFactory factory = new ReflectionBasedClassMetadataProviderFactory(); + + private Map getMethodsOnClass() { + var nonStaticMethodsOnClass = StreamSupport.stream( + factory.getNonStaticMethodsOnClass(Assistant.class).spliterator(), false) + .collect(Collectors.toMap(Method::getName, Function.identity())); + + assertThat(nonStaticMethodsOnClass).hasSize(2).containsOnlyKeys("hello", "helloStreaming"); + + return nonStaticMethodsOnClass; + } + + @Test + void annotationOnClass() { + assertThat(factory.getAnnotation(Assistant.class, Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is a test"); + } + + @Test + void annotationOnClassNotFound() { + assertThat(factory.getAnnotation(Assistant.class, InputGuardrails.class)) + .isEmpty(); + } + + @Test + void annotationOnMethod() { + var nonStaticMethodsOnClass = getMethodsOnClass(); + + assertThat(factory.getAnnotation(nonStaticMethodsOnClass.get("hello"), Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is just a test"); + + assertThat(factory.getAnnotation(nonStaticMethodsOnClass.get("helloStreaming"), Experimental.class)) + .get() + .extracting(Experimental::value) + .isEqualTo("This is another test"); + } + + @Test + void annotationsOnMethodNotFound() { + var nonStaticMethodsOnClass = getMethodsOnClass(); + + assertThat(factory.getAnnotation(nonStaticMethodsOnClass.get("hello"), InputGuardrails.class)) + .isEmpty(); + + assertThat(factory.getAnnotation(nonStaticMethodsOnClass.get("helloStreaming"), InputGuardrails.class)) + .isEmpty(); + } + + @Experimental("This is a test") + interface Assistant { + @Experimental("This is just a test") + String hello(); + + @Experimental("This is another test") + TokenStream helloStreaming(); + } +} diff --git a/langchain4j/src/test/java/dev/langchain4j/service/AiServiceTokenStreamTest.java b/langchain4j/src/test/java/dev/langchain4j/service/AiServiceTokenStreamTest.java index a5ac03ce8a..c90386c8cd 100644 --- a/langchain4j/src/test/java/dev/langchain4j/service/AiServiceTokenStreamTest.java +++ b/langchain4j/src/test/java/dev/langchain4j/service/AiServiceTokenStreamTest.java @@ -5,11 +5,14 @@ import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.mockito.Mockito.mock; import dev.langchain4j.data.message.ChatMessage; +import dev.langchain4j.guardrail.GuardrailRequestParams; +import dev.langchain4j.model.chat.ChatModel; import dev.langchain4j.model.chat.StreamingChatModel; import dev.langchain4j.model.chat.response.ChatResponse; import dev.langchain4j.rag.content.Content; import java.util.ArrayList; import java.util.List; +import java.util.Map; import java.util.function.Consumer; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -105,14 +108,24 @@ class AiServiceTokenStreamTest { } private AiServiceTokenStream setupAiServiceTokenStream() { - StreamingChatModel model = mock(StreamingChatModel.class); + StreamingChatModel streamingModel = mock(StreamingChatModel.class); + ChatModel chatModel = mock(ChatModel.class); + AiServiceContext context = new AiServiceContext(getClass()); - context.streamingChatModel = model; + context.streamingChatModel = streamingModel; + context.chatModel = chatModel; + return new AiServiceTokenStream(AiServiceTokenStreamParameters.builder() .messages(messages) .retrievedContents(content) .context(context) .memoryId(memoryId) + .commonGuardrailParams(GuardrailRequestParams.builder() + .chatMemory(null) + .augmentationResult(null) + .userMessageTemplate("") + .variables(Map.of()) + .build()) .build()); } } diff --git a/langchain4j/src/test/java/dev/langchain4j/service/common/AbstractStreamingAiServiceIT.java b/langchain4j/src/test/java/dev/langchain4j/service/common/AbstractStreamingAiServiceIT.java index bc18faece7..14eca9ef75 100644 --- a/langchain4j/src/test/java/dev/langchain4j/service/common/AbstractStreamingAiServiceIT.java +++ b/langchain4j/src/test/java/dev/langchain4j/service/common/AbstractStreamingAiServiceIT.java @@ -48,9 +48,7 @@ public abstract class AbstractStreamingAiServiceIT { // given model = spy(model); - Assistant assistant = - AiServices.builder(Assistant.class).streamingChatModel(model).build(); - + Assistant assistant = AiServices.create(Assistant.class, model); StringBuilder answerBuilder = new StringBuilder(); CompletableFuture futureAnswer = new CompletableFuture<>(); CompletableFuture futureChatResponse = new CompletableFuture<>(); diff --git a/langchain4j/src/test/java/dev/langchain4j/service/guardrail/AiServiceGuardrailTests.java b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/AiServiceGuardrailTests.java new file mode 100644 index 0000000000..2138278e08 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/AiServiceGuardrailTests.java @@ -0,0 +1,278 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.ChatMessageType; +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.GuardrailException; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.service.AiServices; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class AiServiceGuardrailTests { + @Test + void noGuardrails() { + var noGuardrails = Assistant.create(); + + assertThat(noGuardrails.chat("Hello!")).isEqualTo("Request: Hello!; Response: Hi!"); + assertThat(noGuardrails.chat2("Hello!")).isEqualTo("Request: Hello!; Response: Hi!"); + } + + @ParameterizedTest + @MethodSource("classLevelAssistants") + void classLevelAssistants(String testDescription, Assistant assistant) { + assertThatExceptionOfType(GuardrailException.class) + .isThrownBy(() -> assistant.chat("Hello!")) + .withMessageContaining( + "The guardrail %s failed with this message: Request: Hello! from %s; Response: Hi! failure from %s", + OutputGuardrailFail.class.getName(), + InputGuardrailSuccess.class.getSimpleName(), + OutputGuardrailFail.class.getSimpleName()); + assertThatExceptionOfType(GuardrailException.class) + .isThrownBy(() -> assistant.chat2("Hello!")) + .withMessageContaining( + "The guardrail %s failed with this message: Request: Hello! from %s; Response: Hi! failure from %s", + OutputGuardrailFail.class.getName(), + InputGuardrailSuccess.class.getSimpleName(), + OutputGuardrailFail.class.getSimpleName()); + } + + @Test + void methodLevelAssistant() { + var assistant = MethodLevelAssistant.create(); + + assertThat(assistant.chat("Hello!")) + .isEqualTo( + "Request: Hello! from %s; Response: Hi! from %s", + InputGuardrailSuccess.class.getSimpleName(), OutputGuardrailSuccess.class.getSimpleName()); + assertThatExceptionOfType(GuardrailException.class) + .isThrownBy(() -> assistant.chat2("Hello!")) + .withMessage( + "The guardrail %s failed with this message: Hello! failure from %s", + InputGuardrailFail.class.getName(), InputGuardrailFail.class.getSimpleName()); + } + + @Test + void anotherMethodLevelAssistant() { + var assistant = MethodLevelAssistant1.create(); + + assertThatExceptionOfType(GuardrailException.class) + .isThrownBy(() -> assistant.chat("Hello!")) + .withMessage( + "The guardrail %s failed with this message: Hello! from %s failure from %s", + InputGuardrailFail.class.getName(), + InputGuardrailSuccess.class.getSimpleName(), + InputGuardrailFail.class.getSimpleName()); + + assertThat(assistant.chat2("Hello!")).isEqualTo("Request: Hello!; Response: Hi!"); + } + + @Test + void classAndMethodLevelAssistant() { + var assistant = ClassAndMethodLevelAssistant.create(); + + assertThatExceptionOfType(GuardrailException.class) + .isThrownBy(() -> assistant.chat("Hello!")) + .withMessage( + "The guardrail %s failed with this message: Hello! from %s failure from %s", + InputGuardrailFail.class.getName(), + InputGuardrailSuccess.class.getSimpleName(), + InputGuardrailFail.class.getSimpleName()); + + assertThat(assistant.chat2("Hello!")) + .isEqualTo( + "Request: Hello! from %s; Response: Hi! from %s", + InputGuardrailSuccess.class.getSimpleName(), OutputGuardrailSuccess.class.getSimpleName()); + } + + static Stream classLevelAssistants() { + return Stream.of( + Arguments.of("assistant with class-level annotations", ClassLevelAssistant.create()), + Arguments.of("assistant with method-level annotations", SameMethodLevelAssistant.create()), + Arguments.of( + "assistant with guardrail classes defined (class-level)", + ClassLevelAssistant.createUsingClassNames()), + Arguments.of( + "assistant with guardrail instances defined (class-level)", + ClassLevelAssistant.createUsingClassInstances()), + Arguments.of( + "assistant with guardrail classes defined (method-level)", + SameMethodLevelAssistant.createUsingClassNames()), + Arguments.of( + "assistant with guardrail instances defined (method-level)", + SameMethodLevelAssistant.createUsingClassInstances())); + } + + interface Assistant { + String chat(String message); + + String chat2(String message); + + static T create(Class clazz) { + return AiServices.create(clazz, new MyChatModel()); + } + + static Assistant create() { + return create(Assistant.class); + } + } + + @InputGuardrails(InputGuardrailSuccess.class) + @OutputGuardrails(OutputGuardrailFail.class) + interface ClassLevelAssistant extends Assistant { + static Assistant create() { + return AiServices.create(ClassLevelAssistant.class, new MyChatModel()); + } + + static Assistant createUsingClassNames() { + return AiServices.builder(ClassLevelAssistant.class) + .chatModel(new MyChatModel()) + .inputGuardrailClasses(InputGuardrailSuccess.class) + .outputGuardrailClasses(OutputGuardrailFail.class) + .build(); + } + + static Assistant createUsingClassInstances() { + return AiServices.builder(ClassLevelAssistant.class) + .chatModel(new MyChatModel()) + .inputGuardrails(new InputGuardrailSuccess()) + .outputGuardrails(new OutputGuardrailFail()) + .build(); + } + } + + interface SameMethodLevelAssistant extends Assistant { + @InputGuardrails(InputGuardrailSuccess.class) + @OutputGuardrails(OutputGuardrailFail.class) + @Override + String chat(String message); + + @InputGuardrails(InputGuardrailSuccess.class) + @OutputGuardrails(OutputGuardrailFail.class) + @Override + String chat2(String message); + + static Assistant create() { + return Assistant.create(SameMethodLevelAssistant.class); + } + + static Assistant createUsingClassNames() { + return AiServices.builder(SameMethodLevelAssistant.class) + .chatModel(new MyChatModel()) + .inputGuardrailClasses(InputGuardrailSuccess.class) + .outputGuardrailClasses(OutputGuardrailFail.class) + .build(); + } + + static Assistant createUsingClassInstances() { + return AiServices.builder(SameMethodLevelAssistant.class) + .chatModel(new MyChatModel()) + .inputGuardrails(new InputGuardrailSuccess()) + .outputGuardrails(new OutputGuardrailFail()) + .build(); + } + } + + interface MethodLevelAssistant extends Assistant { + @InputGuardrails(InputGuardrailSuccess.class) + @OutputGuardrails(OutputGuardrailSuccess.class) + @Override + String chat(String message); + + @InputGuardrails(InputGuardrailFail.class) + @OutputGuardrails(OutputGuardrailFail.class) + @Override + String chat2(String message); + + static Assistant create() { + return Assistant.create(MethodLevelAssistant.class); + } + } + + interface MethodLevelAssistant1 extends Assistant { + @InputGuardrails({InputGuardrailSuccess.class, InputGuardrailFail.class}) + @OutputGuardrails( + value = {OutputGuardrailSuccess.class, OutputGuardrailFail.class}, + maxRetries = 10) + @Override + String chat(String message); + + static Assistant create() { + return Assistant.create(MethodLevelAssistant1.class); + } + } + + @InputGuardrails(InputGuardrailSuccess.class) + @OutputGuardrails(OutputGuardrailSuccess.class) + interface ClassAndMethodLevelAssistant extends Assistant { + @InputGuardrails({InputGuardrailSuccess.class, InputGuardrailFail.class}) + @OutputGuardrails( + value = {OutputGuardrailSuccess.class, OutputGuardrailFail.class}, + maxRetries = 10) + @Override + String chat(String message); + + static Assistant create() { + return Assistant.create(ClassAndMethodLevelAssistant.class); + } + } + + public static class InputGuardrailSuccess implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return successWith(userMessage.singleText() + " from " + getClass().getSimpleName()); + } + } + + public static class InputGuardrailFail implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return failure( + userMessage.singleText() + " failure from " + getClass().getSimpleName()); + } + } + + public static class OutputGuardrailSuccess implements OutputGuardrail { + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + return successWith(responseFromLLM.text() + " from " + getClass().getSimpleName()); + } + } + + public static class OutputGuardrailFail implements OutputGuardrail { + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + return failure( + responseFromLLM.text() + " failure from " + getClass().getSimpleName()); + } + } + + public static class MyChatModel implements ChatModel { + private static String getUserMessage(ChatRequest chatRequest) { + return chatRequest.messages().stream() + .filter(message -> message.type() == ChatMessageType.USER) + .findFirst() + .map(chatMessage -> ((UserMessage) chatMessage).singleText()) + .orElseThrow(() -> new IllegalArgumentException("No user message found")); + } + + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + return ChatResponse.builder() + .aiMessage(AiMessage.from("Request: %s; Response: Hi!".formatted(getUserMessage(chatRequest)))) + .build(); + } + } +} diff --git a/langchain4j/src/test/java/dev/langchain4j/service/guardrail/DefaultGuardrailServiceTests.java b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/DefaultGuardrailServiceTests.java new file mode 100644 index 0000000000..be44343728 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/DefaultGuardrailServiceTests.java @@ -0,0 +1,357 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThat; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.guardrail.OutputGuardrail; +import dev.langchain4j.guardrail.OutputGuardrailResult; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class DefaultGuardrailServiceTests { + @Test + void noGuardrails() { + var guardrailService = (DefaultGuardrailService) + GuardrailService.builder(NoGuardrailAssistant.class).build(); + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(0); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(0); + assertThat(guardrailService.getInputConfig("chat")).isEmpty(); + assertThat(guardrailService.getInputGuardrails("chat")).isEmpty(); + assertThat(guardrailService.getOutputConfig("chat")).isEmpty(); + assertThat(guardrailService.getOutputGuardrails("chat")).isEmpty(); + assertThat(guardrailService.getInputConfig("chat2")).isEmpty(); + assertThat(guardrailService.getInputGuardrails("chat2")).isEmpty(); + assertThat(guardrailService.getOutputConfig("chat2")).isEmpty(); + assertThat(guardrailService.getOutputGuardrails("chat2")).isEmpty(); + } + + @Test + void classLevelGuardrailsNoBuilders() { + var gs = GuardrailService.builder(ClassLevelAssistant.class).build(); + assertThat(gs).isInstanceOf(DefaultGuardrailService.class); + var guardrailService = (DefaultGuardrailService) gs; + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getInputConfig("chat")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat")).singleElement().isExactlyInstanceOf(OG1.class); + + assertThat(guardrailService.getInputConfig("chat2")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat2")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat2")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat2")) + .singleElement() + .isExactlyInstanceOf(OG1.class); + } + + @ParameterizedTest + @MethodSource("classLevelGuardrailBuilders") + void classLevelGuardrails(String testDescription, GuardrailService gs) { + assertThat(gs).isInstanceOf(DefaultGuardrailService.class); + var guardrailService = (DefaultGuardrailService) gs; + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getInputConfig("chat")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat")).singleElement().isExactlyInstanceOf(OG1.class); + + assertThat(guardrailService.getInputConfig("chat2")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat2")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat2")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat2")) + .singleElement() + .isExactlyInstanceOf(OG1.class); + } + + @Test + void methodLevelGuardrails() { + var guardrailService = (DefaultGuardrailService) + GuardrailService.builder(MethodLevelAssistant.class).build(); + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getInputConfig("chat")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat")).singleElement().isExactlyInstanceOf(OG1.class); + + assertThat(guardrailService.getInputConfig("chat2")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat2")).singleElement().isExactlyInstanceOf(IG2.class); + + assertThat(guardrailService.getOutputConfig("chat2")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat2")) + .singleElement() + .isExactlyInstanceOf(OG2.class); + } + + @Test + void classAndMethodLevelGuardrailsNoBuilders() { + var gs = GuardrailService.builder(ClassAndMethodLevelAssistant.class).build(); + assertThat(gs).isInstanceOf(DefaultGuardrailService.class); + var guardrailService = (DefaultGuardrailService) gs; + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getInputConfig("chat")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(IG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(IG2.class)); + + assertThat(guardrailService.getOutputConfig("chat")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(10); + + assertThat(guardrailService.getOutputGuardrails("chat")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(OG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(OG2.class)); + + assertThat(guardrailService.getInputConfig("chat2")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat2")).singleElement().isExactlyInstanceOf(IG1.class); + + assertThat(guardrailService.getOutputConfig("chat2")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(dev.langchain4j.guardrail.config.OutputGuardrailsConfig.MAX_RETRIES_DEFAULT); + + assertThat(guardrailService.getOutputGuardrails("chat2")) + .singleElement() + .isExactlyInstanceOf(OG1.class); + } + + @ParameterizedTest + @MethodSource("classAndMethodLevelGuardrailBuilders") + void classAndMethodLevelGuardrails(String testDescription, GuardrailService gs) { + assertThat(gs).isInstanceOf(DefaultGuardrailService.class); + var guardrailService = (DefaultGuardrailService) gs; + + assertThat(guardrailService.getInputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getOutputGuardrailMethodCount()).isEqualTo(2); + assertThat(guardrailService.getInputConfig("chat")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(IG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(IG2.class)); + + assertThat(guardrailService.getOutputConfig("chat")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(10); + + assertThat(guardrailService.getOutputGuardrails("chat")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(OG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(OG2.class)); + + assertThat(guardrailService.getInputConfig("chat2")).isNotEmpty(); + + assertThat(guardrailService.getInputGuardrails("chat2")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(IG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(IG2.class)); + + assertThat(guardrailService.getOutputConfig("chat2")) + .get() + .extracting(dev.langchain4j.guardrail.config.OutputGuardrailsConfig::maxRetries) + .isEqualTo(10); + + assertThat(guardrailService.getOutputGuardrails("chat2")) + .hasSize(2) + .satisfiesExactly( + guardrail -> assertThat(guardrail).isInstanceOf(OG1.class), + guardrail -> assertThat(guardrail).isInstanceOf(OG2.class)); + } + + static Stream classLevelGuardrailBuilders() { + return Stream.of( + Arguments.of( + "assistant with annotations", + GuardrailService.builder(ClassLevelAssistant.class).build()), + Arguments.of( + "assistant with guardrail classes defined", + GuardrailService.builder(NoGuardrailAssistant.class) + .inputGuardrailClasses(IG1.class) + .outputGuardrailClasses(OG1.class) + .build()), + Arguments.of( + "assistant with guardrail instances defined", + GuardrailService.builder(NoGuardrailAssistant.class) + .inputGuardrails(new IG1()) + .outputGuardrails(new OG1()) + .build())); + } + + static Stream classAndMethodLevelGuardrailBuilders() { + return Stream.of( + Arguments.of( + "assistant with annotations", + GuardrailService.builder(MethodLevelAssistant1.class) + .inputGuardrailClasses(IG1.class, IG2.class) + .outputGuardrailClasses(OG1.class, OG2.class) + .outputGuardrailsConfig( + dev.langchain4j.guardrail.config.OutputGuardrailsConfig.builder() + .maxRetries(10) + .build()) + .build()), + Arguments.of( + "assistant with guardrail classes defined", + GuardrailService.builder(NoGuardrailAssistant.class) + .inputGuardrailClasses(IG1.class, IG2.class) + .outputGuardrailClasses(OG1.class, OG2.class) + .outputGuardrailsConfig( + dev.langchain4j.guardrail.config.OutputGuardrailsConfig.builder() + .maxRetries(10) + .build()) + .build()), + Arguments.of( + "assistant with guardrail instances defined", + GuardrailService.builder(NoGuardrailAssistant.class) + .inputGuardrails(new IG1(), new IG2()) + .outputGuardrails(new OG1(), new OG2()) + .outputGuardrailsConfig( + dev.langchain4j.guardrail.config.OutputGuardrailsConfig.builder() + .maxRetries(10) + .build()) + .build())); + } + + interface NoGuardrailAssistant { + String chat(String message); + + String chat2(String message); + + static void doSomething() {} + } + + @InputGuardrails(IG1.class) + @OutputGuardrails(OG1.class) + interface ClassLevelAssistant { + String chat(String message); + + String chat2(String message); + + static void doSomething() {} + } + + interface MethodLevelAssistant { + @InputGuardrails(IG1.class) + @OutputGuardrails(OG1.class) + String chat(String message); + + @InputGuardrails(IG2.class) + @OutputGuardrails(OG2.class) + String chat2(String message); + + static void doSomething() {} + } + + interface MethodLevelAssistant1 { + @InputGuardrails({IG1.class, IG2.class}) + @OutputGuardrails( + value = {OG1.class, OG2.class}, + maxRetries = 10) + String chat(String message); + + String chat2(String message); + + static void doSomething() {} + } + + @InputGuardrails(IG1.class) + @OutputGuardrails(OG1.class) + interface ClassAndMethodLevelAssistant { + @InputGuardrails({IG1.class, IG2.class}) + @OutputGuardrails( + value = {OG1.class, OG2.class}, + maxRetries = 10) + String chat(String message); + + String chat2(String message); + + static void doSomething() {} + } + + public static class IG1 implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return success(); + } + } + + public static class IG2 implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage userMessage) { + return success(); + } + } + + public static class OG1 implements OutputGuardrail { + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + return success(); + } + } + + public static class OG2 implements OutputGuardrail { + @Override + public OutputGuardrailResult validate(AiMessage responseFromLLM) { + return success(); + } + } +} diff --git a/langchain4j/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java new file mode 100644 index 0000000000..f64d90a813 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/InputGuardrailChainTests.java @@ -0,0 +1,107 @@ +package dev.langchain4j.service.guardrail; + +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import dev.langchain4j.data.message.AiMessage; +import dev.langchain4j.data.message.UserMessage; +import dev.langchain4j.guardrail.InputGuardrail; +import dev.langchain4j.guardrail.InputGuardrailException; +import dev.langchain4j.guardrail.InputGuardrailResult; +import dev.langchain4j.model.chat.ChatModel; +import dev.langchain4j.model.chat.request.ChatRequest; +import dev.langchain4j.model.chat.response.ChatResponse; +import dev.langchain4j.service.AiServices; +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.MethodSource; + +class InputGuardrailChainTests { + @ParameterizedTest + @MethodSource("aiServices") + void failsTheChain(AnAiService aiService) { + assertThatThrownBy(() -> aiService.failingFirstTwo("foo")) + .isInstanceOf(InputGuardrailException.class) + .hasCauseInstanceOf(ValidationException.class) + .hasRootCauseMessage("boom"); + } + + static List aiServices() { + return List.of( + createAiServiceWithClassNames(), + createAiServiceWithInstances(), + createMethodLevelAnnotationAiService(), + createClassLevelAnnotationAiService()); + } + + private static MyAnnotationMethodLevelAiService createMethodLevelAnnotationAiService() { + return AiServices.create(MyAnnotationMethodLevelAiService.class, new MyChatModel()); + } + + private static MyAnnotationClassLevelAiService createClassLevelAnnotationAiService() { + return AiServices.create(MyAnnotationClassLevelAiService.class, new MyChatModel()); + } + + private static MyAiService createAiServiceWithClassNames() { + return AiServices.builder(MyAiService.class) + .chatModel(new MyChatModel()) + .inputGuardrailClasses(FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class) + .build(); + } + + private static MyAiService createAiServiceWithInstances() { + return AiServices.builder(MyAiService.class) + .chatModel(new MyChatModel()) + .inputGuardrails(new FirstGuardrail(), new FailingGuardrail(), new SecondGuardrail()) + .build(); + } + + @InputGuardrails({FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class}) + public interface MyAnnotationClassLevelAiService extends AnAiService {} + + public interface MyAiService extends AnAiService {} + + public interface MyAnnotationMethodLevelAiService extends AnAiService { + @InputGuardrails({FirstGuardrail.class, FailingGuardrail.class, SecondGuardrail.class}) + @Override + String failingFirstTwo(String message); + } + + public interface AnAiService { + String failingFirstTwo(String message); + } + + public static class FirstGuardrail implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage um) { + return success(); + } + } + + public static class SecondGuardrail implements InputGuardrail { + @Override + public InputGuardrailResult validate(UserMessage um) { + return success(); + } + } + + public static class FailingGuardrail implements InputGuardrail { + AtomicInteger spy = new AtomicInteger(0); + + @Override + public InputGuardrailResult validate(UserMessage um) { + if (spy.incrementAndGet() == 1) { + return fatal("boom", new ValidationException("boom")); + } + + return success(); + } + } + + public static class MyChatModel implements ChatModel { + @Override + public ChatResponse doChat(ChatRequest chatRequest) { + return ChatResponse.builder().aiMessage(AiMessage.from("Hi!")).build(); + } + } +} diff --git a/langchain4j/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java new file mode 100644 index 0000000000..9b5ec3d529 --- /dev/null +++ b/langchain4j/src/test/java/dev/langchain4j/service/guardrail/ValidationException.java @@ -0,0 +1,7 @@ +package dev.langchain4j.service.guardrail; + +public class ValidationException extends RuntimeException { + public ValidationException(String message) { + super(message); + } +} diff --git a/pom.xml b/pom.xml index aa06321ac2..33ab926641 100644 --- a/pom.xml +++ b/pom.xml @@ -14,6 +14,7 @@ langchain4j-bom langchain4j-core + langchain4j-test langchain4j langchain4j-kotlin @@ -99,6 +100,9 @@ experimental/langchain4j-experimental-sql + + integration-tests + @@ -174,6 +178,9 @@ aggregate false + + integration-tests + default