From 917836710960f6fdba180c1f8f8d5908c795a2a5 Mon Sep 17 00:00:00 2001 From: Konstantin Pavlov <1517853+kpavlov@users.noreply.github.com> Date: Thu, 31 Oct 2024 17:34:05 +0200 Subject: [PATCH] FIX BUILD: Fix code coverage and some improvements (#2010) ## Issue Build was failing due to low test coverage after [this change](https://github.com/langchain4j/langchain4j/pull/1987/files#diff-9a5519a26a4b6d2fd412c877de4459c2a119e793fece9385f558a93a7e0aee5aR273-R275). The hotfix is to enable tracing for logger with mockStatic for in the affected DefaultRetrievalAugmentorTest. Other changes: * **Refine integration test run conditions to exclude experimental builds and require the presence of the `OPENAI_API_KEY` secret.** * [Upgrade mockito instrumentation](https://javadoc.io/doc/org.mockito/mockito-core/latest/org/mockito/Mockito.html#0.3) (javaagent) setup to make it compatible with new JDKs * Updated `pom.xml` to upgrade `jacoco-maven-plugin` version. * Enhanced `README.md` with additional badges for nightly build and Codacy dashboard. (`[README.mdL3-R6](diffhunk://#diff-b335630551682c19a781afebcf4d07bf978fb1f8ac04c6bf87428ed5106870f5L3-R6)`) * Introduced new tests for comparison filters * Added tests for logical filters * Added configuration for external dependencies in `.idea/externalDependencies.xml` to include SonarLint and SpotBugs plugins. ## General checklist - [x] There are no breaking changes - [x] I have added unit and integration tests for my change - [ ] I have manually run all the unit and integration tests in the module I have added/changed, and they are all green - [ ] I have manually run all the unit and integration tests in the [core](https://github.com/langchain4j/langchain4j/tree/main/langchain4j-core) and [main](https://github.com/langchain4j/langchain4j/tree/main/langchain4j) modules, and they are all green - [ ] I have added/updated the [documentation](https://github.com/langchain4j/langchain4j/tree/main/docs/docs) - [ ] I have added an example in the [examples repo](https://github.com/langchain4j/langchain4j-examples) (only for "big" features) - [ ] I have added/updated [Spring Boot starter(s)](https://github.com/langchain4j/langchain4j-spring) (if applicable) --- .github/workflows/main.yaml | 11 +- .idea/externalDependencies.xml | 7 ++ README.md | 5 +- .../GraalVmJavaScriptExecutionEngine.java | 10 +- .../graalvm/GraalVmPythonExecutionEngine.java | 10 +- .../code/judge0/Judge0JavaScriptEngine.java | 18 +-- langchain4j-core/pom.xml | 21 +++- .../rag/DefaultRetrievalAugmentorTest.java | 118 +++++++++++------- .../router/LanguageModelQueryRouterTest.java | 9 +- .../comparison/AbstractComparisonTest.java | 27 ++++ .../filter/comparison/IsEqualToTest.java | 52 ++++++++ .../IsGreaterThanOrEqualToTest.java | 30 +++++ .../filter/comparison/IsGreaterThanTest.java | 34 +++++ .../comparison/IsLessThanOrEqualToTest.java | 31 +++++ .../filter/comparison/IsLessThanTest.java | 31 +++++ .../filter/comparison/IsNotEqualToTest.java | 71 +++++++++++ .../embedding/filter/logical/AndTest.java | 75 +++++++++++ .../embedding/filter/logical/NotTest.java | 78 ++++++++++++ .../embedding/filter/logical/OrTest.java | 74 +++++++++++ langchain4j-parent/pom.xml | 18 ++- 20 files changed, 650 insertions(+), 80 deletions(-) create mode 100644 .idea/externalDependencies.xml create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/AbstractComparisonTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsEqualToTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanOrEqualToTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanOrEqualToTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsNotEqualToTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/AndTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/NotTest.java create mode 100644 langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/OrTest.java diff --git a/.github/workflows/main.yaml b/.github/workflows/main.yaml index 0a728c3f15..e230955bde 100644 --- a/.github/workflows/main.yaml +++ b/.github/workflows/main.yaml @@ -46,6 +46,8 @@ jobs: - java_version: 23 mvn_opts: '' experimental: true + env: + OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} steps: - uses: actions/checkout@v4 with: @@ -69,11 +71,10 @@ jobs: ${{ matrix.mvn_opts }} - name: Integration test with JDK ${{ matrix.java_version }} - ## allow builds to run only: - ## - on the main branch when it is not a pull request, - ## - for pull requests from the same repository - ## - when triggered manually - if: ${{ !matrix.experimental }} && (github.event_name == 'pull_request' && github.head_repo.full_name == github.repository) || (github.event_name == 'push' && github.ref == 'refs/heads/main') || github.event_name == 'workflow_dispatch' + ## The step or job will only run if the `experimental` variable + ## in the matrix is false (not set to true) + ## and the OPENAI_API_KEY secret is available and not empty. + if: ${{ !matrix.experimental && env.OPENAI_API_KEY != '' }} run: | mvn -B -U verify \ -Dgib.disable=false -Dgib.referenceBranch=__branch_before \ diff --git a/.idea/externalDependencies.xml b/.idea/externalDependencies.xml new file mode 100644 index 0000000000..5718858203 --- /dev/null +++ b/.idea/externalDependencies.xml @@ -0,0 +1,7 @@ + + + + + + + diff --git a/README.md b/README.md index 7e32513d4d..8571b451ba 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,9 @@ # LangChain for Java: Supercharge your Java application with the power of LLMs -[![Build Status](https://img.shields.io/github/actions/workflow/status/langchain4j/langchain4j/main.yaml?branch=main&style=for-the-badge&label=GITHUB%20ACTIONS&logo=github)](https://github.com/langchain4j/langchain4j/actions/workflows/main.yaml) +[![Build Status](https://img.shields.io/github/actions/workflow/status/langchain4j/langchain4j/main.yaml?branch=main&style=for-the-badge&label=CI%20BUILD&logo=github)](https://github.com/langchain4j/langchain4j/actions/workflows/main.yaml) +[![Nightly Build](https://img.shields.io/github/actions/workflow/status/langchain4j/langchain4j/nightly.yaml?branch=main&style=for-the-badge&label=NIGHTLY%20BUILD&logo=github)](https://github.com/langchain4j/langchain4j/actions/workflows/nightly.yaml) +[![CODACY](https://img.shields.io/badge/Codacy-Dashboard-blue?style=for-the-badge&logo=codacy)](https://app.codacy.com/gh/langchain4j/langchain4j/dashboard) + [![Discord](https://dcbadge.vercel.app/api/server/JzTFvyjG6R?style=for-the-badge)](https://discord.gg/JzTFvyjG6R) [![X](https://img.shields.io/badge/@langchain4j-follow-blue?logo=x&style=for-the-badge)](https://x.com/langchain4j) [![Maven Version](https://img.shields.io/maven-central/v/dev.langchain4j/langchain4j?logo=apachemaven&style=for-the-badge)](https://search.maven.org/#search|gav|1|g:"dev.langchain4j"%20AND%20a:"langchain4j") diff --git a/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmJavaScriptExecutionEngine.java b/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmJavaScriptExecutionEngine.java index f95039a265..f5526ce681 100644 --- a/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmJavaScriptExecutionEngine.java +++ b/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmJavaScriptExecutionEngine.java @@ -22,11 +22,11 @@ public class GraalVmJavaScriptExecutionEngine implements CodeExecutionEngine { public String execute(String code) { OutputStream outputStream = new ByteArrayOutputStream(); try (Context context = Context.newBuilder("js") - .sandbox(CONSTRAINED) - .allowHostAccess(UNTRUSTED) - .out(outputStream) - .err(outputStream) - .build()) { + .sandbox(CONSTRAINED) + .allowHostAccess(UNTRUSTED) + .out(outputStream) + .err(outputStream) + .build()) { Object result = context.eval("js", code).as(Object.class); return String.valueOf(result); } diff --git a/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmPythonExecutionEngine.java b/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmPythonExecutionEngine.java index b9c778677f..2b6daa5eb2 100644 --- a/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmPythonExecutionEngine.java +++ b/code-execution-engines/langchain4j-code-execution-engine-graalvm-polyglot/src/main/java/dev/langchain4j/code/graalvm/GraalVmPythonExecutionEngine.java @@ -22,11 +22,11 @@ public class GraalVmPythonExecutionEngine implements CodeExecutionEngine { public String execute(String code) { OutputStream outputStream = new ByteArrayOutputStream(); try (Context context = Context.newBuilder("python") - .sandbox(TRUSTED) - .allowHostAccess(UNTRUSTED) - .out(outputStream) - .err(outputStream) - .build()) { + .sandbox(TRUSTED) + .allowHostAccess(UNTRUSTED) + .out(outputStream) + .err(outputStream) + .build()) { Object result = context.eval("python", code).as(Object.class); return String.valueOf(result); } diff --git a/code-execution-engines/langchain4j-code-execution-engine-judge0/src/main/java/dev/langchain4j/code/judge0/Judge0JavaScriptEngine.java b/code-execution-engines/langchain4j-code-execution-engine-judge0/src/main/java/dev/langchain4j/code/judge0/Judge0JavaScriptEngine.java index 0c6a350bec..5f2e172c8e 100644 --- a/code-execution-engines/langchain4j-code-execution-engine-judge0/src/main/java/dev/langchain4j/code/judge0/Judge0JavaScriptEngine.java +++ b/code-execution-engines/langchain4j-code-execution-engine-judge0/src/main/java/dev/langchain4j/code/judge0/Judge0JavaScriptEngine.java @@ -25,11 +25,11 @@ class Judge0JavaScriptEngine implements CodeExecutionEngine { this.apiKey = apiKey; this.languageId = languageId; this.client = new OkHttpClient.Builder() - .connectTimeout(timeout) - .readTimeout(timeout) - .writeTimeout(timeout) - .callTimeout(timeout) - .build(); + .connectTimeout(timeout) + .readTimeout(timeout) + .writeTimeout(timeout) + .callTimeout(timeout) + .build(); } @Override @@ -42,10 +42,10 @@ class Judge0JavaScriptEngine implements CodeExecutionEngine { RequestBody requestBody = RequestBody.create(Json.toJson(submission), MEDIA_TYPE); Request request = new Request.Builder() - .url("https://judge0-ce.p.rapidapi.com/submissions?base64_encoded=true&wait=true&fields=*") - .addHeader("X-RapidAPI-Key", apiKey) - .post(requestBody) - .build(); + .url("https://judge0-ce.p.rapidapi.com/submissions?base64_encoded=true&wait=true&fields=*") + .addHeader("X-RapidAPI-Key", apiKey) + .post(requestBody) + .build(); try { Response response = client.newCall(request).execute(); diff --git a/langchain4j-core/pom.xml b/langchain4j-core/pom.xml index 9f811704ed..63e2f095ad 100644 --- a/langchain4j-core/pom.xml +++ b/langchain4j-core/pom.xml @@ -34,6 +34,7 @@ slf4j-api + org.junit.jupiter junit-jupiter-engine @@ -85,7 +86,6 @@ - org.apache.maven.plugins maven-jar-plugin @@ -101,7 +101,7 @@ org.jacoco jacoco-maven-plugin - 0.8.11 + 0.8.12 prepare-agent @@ -209,4 +209,21 @@ + + + + org.jacoco + jacoco-maven-plugin + + + + + report + + + + + + + diff --git a/langchain4j-core/src/test/java/dev/langchain4j/rag/DefaultRetrievalAugmentorTest.java b/langchain4j-core/src/test/java/dev/langchain4j/rag/DefaultRetrievalAugmentorTest.java index cbe5a4dc4f..f4a8d2f40a 100644 --- a/langchain4j-core/src/test/java/dev/langchain4j/rag/DefaultRetrievalAugmentorTest.java +++ b/langchain4j-core/src/test/java/dev/langchain4j/rag/DefaultRetrievalAugmentorTest.java @@ -12,9 +12,14 @@ import dev.langchain4j.rag.query.router.DefaultQueryRouter; import dev.langchain4j.rag.query.router.QueryRouter; import dev.langchain4j.rag.query.transformer.DefaultQueryTransformer; import dev.langchain4j.rag.query.transformer.QueryTransformer; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.MethodSource; +import org.mockito.MockedStatic; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import java.util.Collection; import java.util.HashMap; @@ -32,14 +37,31 @@ import static java.util.stream.Collectors.toList; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.Mockito.any; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.mockStatic; import static org.mockito.Mockito.spy; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.verifyNoInteractions; import static org.mockito.Mockito.verifyNoMoreInteractions; +import static org.mockito.Mockito.when; class DefaultRetrievalAugmentorTest { + private static MockedStatic loggerFactoryMock; + + @BeforeAll + static void mockLogger() { + loggerFactoryMock = mockStatic(LoggerFactory.class); + Logger logger = mock(Logger.class); + when(LoggerFactory.getLogger(DefaultRetrievalAugmentor.class)).thenReturn(logger); + when(logger.isTraceEnabled()).thenReturn(true); + } + + @AfterAll + static void releaseLogger() { + loggerFactoryMock.close(); + } + @ParameterizedTest @MethodSource("executors") void should_augment_user_message__multiple_queries_multiple_retrievers(Executor executor) { @@ -64,12 +86,12 @@ class DefaultRetrievalAugmentorTest { ContentInjector contentInjector = spy(new TestContentInjector()); RetrievalAugmentor retrievalAugmentor = DefaultRetrievalAugmentor.builder() - .queryTransformer(queryTransformer) - .queryRouter(queryRouter) - .contentAggregator(contentAggregator) - .contentInjector(contentInjector) - .executor(executor) - .build(); + .queryTransformer(queryTransformer) + .queryRouter(queryRouter) + .contentAggregator(contentAggregator) + .contentInjector(contentInjector) + .executor(executor) + .build(); UserMessage userMessage = UserMessage.from("query"); @@ -80,7 +102,7 @@ class DefaultRetrievalAugmentorTest { // then assertThat(augmented.singleText()).isEqualTo( - """ + """ query content 1 content 2 @@ -109,25 +131,25 @@ class DefaultRetrievalAugmentorTest { Map>> queryToContents = new HashMap<>(); queryToContents.put(query1, asList( - asList(content1, content2), - asList(content3, content4) + asList(content1, content2), + asList(content3, content4) )); queryToContents.put(query2, asList( - asList(content1, content2), - asList(content3, content4) + asList(content1, content2), + asList(content3, content4) )); verify(contentAggregator).aggregate(queryToContents); verifyNoMoreInteractions(contentAggregator); verify(contentInjector).inject(asList( - content1, content2, content3, content4, - content1, content2, content3, content4 + content1, content2, content3, content4, + content1, content2, content3, content4 ), userMessage); verify(contentInjector).inject(asList( - content1, content2, content3, content4, - content1, content2, content3, content4 + content1, content2, content3, content4, + content1, content2, content3, content4 ), (ChatMessage) userMessage); verifyNoMoreInteractions(contentInjector); } @@ -155,12 +177,12 @@ class DefaultRetrievalAugmentorTest { Executor executor = spy(new TestExecutor()); RetrievalAugmentor retrievalAugmentor = DefaultRetrievalAugmentor.builder() - .queryTransformer(queryTransformer) - .queryRouter(queryRouter) - .contentAggregator(contentAggregator) - .contentInjector(contentInjector) - .executor(executor) - .build(); + .queryTransformer(queryTransformer) + .queryRouter(queryRouter) + .contentAggregator(contentAggregator) + .contentInjector(contentInjector) + .executor(executor) + .build(); UserMessage userMessage = UserMessage.from("query"); @@ -171,7 +193,7 @@ class DefaultRetrievalAugmentorTest { // then assertThat(augmented.singleText()).isEqualTo( - """ + """ query content 1 content 2 @@ -194,8 +216,8 @@ class DefaultRetrievalAugmentorTest { Map>> queryToContents = new HashMap<>(); queryToContents.put(query, asList( - asList(content1, content2), - asList(content3, content4) + asList(content1, content2), + asList(content3, content4) )); verify(contentAggregator).aggregate(queryToContents); @@ -236,12 +258,12 @@ class DefaultRetrievalAugmentorTest { Executor executor = mock(Executor.class); RetrievalAugmentor retrievalAugmentor = DefaultRetrievalAugmentor.builder() - .queryTransformer(queryTransformer) - .queryRouter(queryRouter) - .contentAggregator(contentAggregator) - .contentInjector(contentInjector) - .executor(executor) - .build(); + .queryTransformer(queryTransformer) + .queryRouter(queryRouter) + .contentAggregator(contentAggregator) + .contentInjector(contentInjector) + .executor(executor) + .build(); UserMessage userMessage = UserMessage.from("query"); @@ -252,7 +274,7 @@ class DefaultRetrievalAugmentorTest { // then assertThat(augmented.singleText()).isEqualTo( - """ + """ query content 1 content 2""" @@ -289,9 +311,9 @@ class DefaultRetrievalAugmentorTest { QueryRouter queryRouter = spy(new TestQueryRouter(retrievers)); RetrievalAugmentor retrievalAugmentor = DefaultRetrievalAugmentor.builder() - .queryRouter(queryRouter) - .executor(executor) - .build(); + .queryRouter(queryRouter) + .executor(executor) + .build(); UserMessage userMessage = UserMessage.from("query"); @@ -309,14 +331,14 @@ class DefaultRetrievalAugmentorTest { static Stream executors() { return Stream.builder() - .add(Executors.newCachedThreadPool()) - .add(Executors.newFixedThreadPool(1)) - .add(Executors.newFixedThreadPool(2)) - .add(Executors.newFixedThreadPool(3)) - .add(Executors.newFixedThreadPool(4)) - .add(Runnable::run) // same thread executor - .add(null) // to use default Executor in DefaultRetrievalAugmentor - .build(); + .add(Executors.newCachedThreadPool()) + .add(Executors.newFixedThreadPool(1)) + .add(Executors.newFixedThreadPool(2)) + .add(Executors.newFixedThreadPool(3)) + .add(Executors.newFixedThreadPool(4)) + .add(Runnable::run) // same thread executor + .add(null) // to use default Executor in DefaultRetrievalAugmentor + .build(); } static class TestQueryTransformer implements QueryTransformer { @@ -366,10 +388,10 @@ class DefaultRetrievalAugmentorTest { @Override public List aggregate(Map>> queryToContents) { return queryToContents.values() - .stream() - .flatMap(Collection::stream) - .flatMap(List::stream) - .collect(toList()); + .stream() + .flatMap(Collection::stream) + .flatMap(List::stream) + .collect(toList()); } } @@ -378,8 +400,8 @@ class DefaultRetrievalAugmentorTest { @Override public UserMessage inject(List contents, UserMessage userMessage) { String joinedContents = contents.stream() - .map(it -> it.textSegment().text()) - .collect(joining("\n")); + .map(it -> it.textSegment().text()) + .collect(joining("\n")); return UserMessage.from(userMessage.text() + "\n" + joinedContents); } } diff --git a/langchain4j-core/src/test/java/dev/langchain4j/rag/query/router/LanguageModelQueryRouterTest.java b/langchain4j-core/src/test/java/dev/langchain4j/rag/query/router/LanguageModelQueryRouterTest.java index 400f242a93..6040cd121e 100644 --- a/langchain4j-core/src/test/java/dev/langchain4j/rag/query/router/LanguageModelQueryRouterTest.java +++ b/langchain4j-core/src/test/java/dev/langchain4j/rag/query/router/LanguageModelQueryRouterTest.java @@ -159,9 +159,10 @@ class LanguageModelQueryRouterTest { ChatModelMock model = ChatModelMock.thatAlwaysResponds("Sorry, I don't know"); - Map retrieverToDescription = new LinkedHashMap<>(); - retrieverToDescription.put(catArticlesRetriever, "articles about cats"); - retrieverToDescription.put(dogArticlesRetriever, "articles about dogs"); + final var retrieverToDescription = Map.of( + catArticlesRetriever, "articles about cats", + dogArticlesRetriever, "articles about dogs" + ); QueryRouter router = new LanguageModelQueryRouter(model, retrieverToDescription); @@ -290,4 +291,4 @@ class LanguageModelQueryRouterTest { .isExactlyInstanceOf(RuntimeException.class) .hasMessageContaining("Something went wrong"); } -} \ No newline at end of file +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/AbstractComparisonTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/AbstractComparisonTest.java new file mode 100644 index 0000000000..79f5e5320e --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/AbstractComparisonTest.java @@ -0,0 +1,27 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import dev.langchain4j.store.embedding.filter.Filter; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +abstract class AbstractComparisonTest { + + protected T subject; + + @Test + void shouldReturnsFalseWhenObjectIsNotMetadata() { + assertThat(subject.test(new Object())).isFalse(); + } + + @Test + void shouldReturnsFalseWhenMetadataDoesNotContainKey() { + Metadata metadata = mock(Metadata.class); + when(metadata.containsKey("key")).thenReturn(false); + + assertThat(subject.test(metadata)).isFalse(); + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsEqualToTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsEqualToTest.java new file mode 100644 index 0000000000..abae8a9ab4 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsEqualToTest.java @@ -0,0 +1,52 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.Test; + +import java.util.HashMap; +import java.util.Map; +import java.util.UUID; + +import static org.assertj.core.api.Assertions.assertThat; + + class IsEqualToTest { + + @Test + void testShouldReturnFalseWhenNotMetadata() { + IsEqualTo isEqualTo = new IsEqualTo("key", "value"); + assertThat(isEqualTo.test("notMetadata")).isFalse(); + } + + @Test + void testShouldReturnFalseWhenKeyNotFound() { + IsEqualTo isEqualTo = new IsEqualTo("key", "value"); + Metadata metadata = new Metadata(Map.of()); + assertThat(isEqualTo.test(metadata)).isFalse(); + } + + @Test + void testShouldReturnTrueWhenValuesAreNumbers() { + IsEqualTo isEqualTo = new IsEqualTo("key", 2); + Metadata metadata = new Metadata(Map.of("key", 2)); + assertThat(isEqualTo.test(metadata)).isTrue(); + } + + @Test + void testShouldReturnTrueWhenValuesAreStrings() { + IsEqualTo isEqualTo = new IsEqualTo("key", "value"); + Metadata metadata = new Metadata(new HashMap<>() {{ + put("key", "value"); + }}); + assertThat(isEqualTo.test(metadata)).isTrue(); + } + + @Test + void testShouldReturnTrueWhenActualValueIsUUIDAsString() { + UUID uuid = UUID.randomUUID(); + IsEqualTo isEqualTo = new IsEqualTo("key", uuid); + Metadata metadata = new Metadata(new HashMap<>() {{ + put("key", uuid.toString()); + }}); + assertThat(isEqualTo.test(metadata)).isTrue(); + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanOrEqualToTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanOrEqualToTest.java new file mode 100644 index 0000000000..1c3de8c594 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanOrEqualToTest.java @@ -0,0 +1,30 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; + +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class IsGreaterThanOrEqualToTest extends AbstractComparisonTest { + + @BeforeEach + void setUp() { + subject = new IsGreaterThanOrEqualTo("key", 5); + } + + @ParameterizedTest + @CsvSource({ + "4, false", + "5, true", + "6, true" + }) + void testComparisonValue(Integer value, boolean expectedResult) { + Metadata metadata = Metadata.from(Map.of("key", value)); + assertThat(subject.test(metadata)).isEqualTo(expectedResult); + } + +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanTest.java new file mode 100644 index 0000000000..2fba3d142a --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsGreaterThanTest.java @@ -0,0 +1,34 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.mockito.MockedConstruction; +import org.mockito.Mockito; + +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class IsGreaterThanTest extends AbstractComparisonTest{ + + @BeforeEach + void beforeEach() { + subject = new IsGreaterThan("key", 5); + } + + @ParameterizedTest + @CsvSource({ + "0, false", + "4, false", + "5, false", + "6, true" + }) + void testComparisonValue(Integer value, boolean expectedResult) { + Metadata metadata = Metadata.from(Map.of("key", value)); + assertThat(subject.test(metadata)).isEqualTo(expectedResult); + } + +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanOrEqualToTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanOrEqualToTest.java new file mode 100644 index 0000000000..37d589084c --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanOrEqualToTest.java @@ -0,0 +1,31 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; + +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class IsLessThanOrEqualToTest extends AbstractComparisonTest { + + @BeforeEach + void setUp() { + subject = new IsLessThanOrEqualTo("key", 5); + } + + @ParameterizedTest + @CsvSource({ + "0, true", + "4, true", + "5, true", + "6, false" + }) + void testComparisonValue(Integer value, boolean expectedResult) { + Metadata metadata = Metadata.from(Map.of("key", value)); + assertThat(subject.test(metadata)).isEqualTo(expectedResult); + } + +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanTest.java new file mode 100644 index 0000000000..ea86d7c2ef --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsLessThanTest.java @@ -0,0 +1,31 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; + +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class IsLessThanTest extends AbstractComparisonTest { + + @BeforeEach + void setUp() { + subject = new IsLessThan("key", 5); + } + + @ParameterizedTest + @CsvSource({ + "0, true", + "4, true", + "5, false", + "6, false" + }) + void testComparisonValue(Integer value, boolean expectedResult) { + Metadata metadata = Metadata.from(Map.of("key", value)); + assertThat(subject.test(metadata)).isEqualTo(expectedResult); + } + +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsNotEqualToTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsNotEqualToTest.java new file mode 100644 index 0000000000..a9d0cd349b --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/comparison/IsNotEqualToTest.java @@ -0,0 +1,71 @@ +package dev.langchain4j.store.embedding.filter.comparison; + +import dev.langchain4j.data.document.Metadata; +import org.junit.jupiter.api.Test; + +import java.util.HashMap; +import java.util.UUID; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class IsNotEqualToTest { + + @Test + void testIsNotEqualToFilter() { + String key = "testKey"; + String unequalValue = "notEqual"; + + IsNotEqualTo subject = new IsNotEqualTo(key, "testValue"); + + assertIsNotEqualToObject(subject); + assertMetadataDoesNotContainKey(subject); + assertValuesAreNumbersNotEqual(key); + assertUuidAndStringComparison(key); + assertUnequalStringValues(key, unequalValue); + assertEqualStringValues(key); + } + + private void assertIsNotEqualToObject(IsNotEqualTo subject) { + // Testing when object is not an instance of Metadata + assertFalse(subject.test(new Object())); + } + + private void assertMetadataDoesNotContainKey(IsNotEqualTo subject) { + // Testing when Metadata does not contain key + assertTrue(subject.test(new Metadata(new HashMap<>()))); + } + + private void assertValuesAreNumbersNotEqual(String key) { + // Testing when actual value is Number + Metadata metadata = new Metadata(new HashMap<>()); + metadata.put(key, 123); + IsNotEqualTo subject = new IsNotEqualTo(key, 1234); + assertTrue(subject.test(metadata)); + } + + private void assertUuidAndStringComparison(String key) { + // Testing when comparisonValue is instance of UUID and actualValue is instance of String + UUID uuid = UUID.randomUUID(); + IsNotEqualTo subject = new IsNotEqualTo(key, uuid); + Metadata metadata = new Metadata(new HashMap<>()); + metadata.put(key, uuid.toString() + "extra"); + assertTrue(subject.test(metadata)); + } + + private void assertUnequalStringValues(String key, String unequalValue) { + // Testing when values are not equal + IsNotEqualTo subject = new IsNotEqualTo(key, unequalValue); + Metadata metadata = new Metadata(new HashMap<>()); + metadata.put(key, "testValue"); + assertTrue(subject.test(metadata)); + } + + private void assertEqualStringValues(String key) { + // Testing when values are equal + IsNotEqualTo subject = new IsNotEqualTo(key, "testValue"); + Metadata metadata = new Metadata(new HashMap<>()); + metadata.put(key, "testValue"); + assertFalse(subject.test(metadata)); + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/AndTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/AndTest.java new file mode 100644 index 0000000000..7f2280bf4b --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/AndTest.java @@ -0,0 +1,75 @@ +package dev.langchain4j.store.embedding.filter.logical; + +import dev.langchain4j.store.embedding.filter.Filter; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +@ExtendWith(MockitoExtension.class) +class AndTest { + + @Mock + private Filter mockFilterPasses; + + @Mock + private Filter mockFilterFails; + + @Test + void testBothFiltersPass() { + when(mockFilterPasses.test(any())).thenReturn(true); + + And andFilter = new And(mockFilterPasses, mockFilterPasses); + + assertThat(andFilter.test(new Object())).isTrue(); + } + + @Test + void testLeftFilterFails() { + when(mockFilterFails.test(any())).thenReturn(false); + + And andFilter = new And(mockFilterFails, mockFilterPasses); + + assertThat(andFilter.test(new Object())).isFalse(); + verifyNoInteractions(mockFilterPasses); + } + + @Test + void testRightFilterFails() { + when(mockFilterPasses.test(any())).thenReturn(true); + when(mockFilterFails.test(any())).thenReturn(false); + + And andFilter = new And(mockFilterPasses, mockFilterFails); + + assertThat(andFilter.test(new Object())).isFalse(); + } + + @Test + void testBothFiltersFail() { + when(mockFilterFails.test(any())).thenReturn(false); + + And andFilter = new And(mockFilterFails, mockFilterFails); + + assertThat(andFilter.test(new Object())).isFalse(); + } + + @Test + void testEqualsAndHashCode() { + And andFilter = new And(mockFilterPasses, mockFilterPasses); + And sameAndFilter = new And(mockFilterPasses, mockFilterPasses); + + assertThat(andFilter) + .isEqualTo(sameAndFilter) + .hasSameHashCodeAs(sameAndFilter); + } + + @Test + void checkToStringImplemented() { + And andFilter = new And(mockFilterPasses, mockFilterPasses); + + assertThat(andFilter).hasToString("And(left=mockFilterPasses, right=mockFilterPasses)"); + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/NotTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/NotTest.java new file mode 100644 index 0000000000..16a9f23bd1 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/NotTest.java @@ -0,0 +1,78 @@ +package dev.langchain4j.store.embedding.filter.logical; + +import dev.langchain4j.store.embedding.filter.Filter; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mock; +import org.mockito.Mockito; +import org.mockito.junit.jupiter.MockitoExtension; +import org.junit.jupiter.api.extension.ExtendWith; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; + +@ExtendWith(MockitoExtension.class) +class NotTest { + + @Mock + private Filter filter; + + private Not subject; + + @BeforeEach + void beforeEach() { + subject = new Not(filter); + } + + @Test + void shouldReturnFalseWhenFilterReturnsTrue() { + Mockito.when(filter.test(any())).thenReturn(true); + + boolean result = subject.test(new Object()); + + assertThat(result).isFalse(); + } + + @Test + void shouldReturnTrueWhenFilterReturnsFalse() { + Mockito.when(filter.test(any())).thenReturn(false); + + boolean result = subject.test(new Object()); + + assertThat(result).isTrue(); + } + + @Test + void shouldReturnCorrectExpression() { + Filter result = subject.expression(); + + assertThat(result).isEqualTo(filter); + } + + @Test + void shouldHaveCorrectToStringImplementation() { + assertThat(subject).hasToString("Not(expression=" + filter + ")"); + } + + @Test + void shouldReturnTrueForEqualsWithSameExpression() { + Not anotherWithSameExp = new Not(filter); + + assertThat(subject) + .isEqualTo(anotherWithSameExp) + .hasSameHashCodeAs(anotherWithSameExp); + } + + @Test + void shouldReturnFalseForEqualsWithDifferentExpression() { + Not anotherWithDiffExp = new Not(Mockito.mock(Filter.class)); + assertThat(subject).isNotEqualTo(anotherWithDiffExp); + } + + @Test + void shouldHaveConsistentHashCode() { + int initialHashCode = subject.hashCode(); + + assertThat(subject.hashCode()).isEqualTo(initialHashCode); + } +} diff --git a/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/OrTest.java b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/OrTest.java new file mode 100644 index 0000000000..595adf6da7 --- /dev/null +++ b/langchain4j-core/src/test/java/dev/langchain4j/store/embedding/filter/logical/OrTest.java @@ -0,0 +1,74 @@ +package dev.langchain4j.store.embedding.filter.logical; + +import dev.langchain4j.store.embedding.filter.Filter; +import org.junit.jupiter.api.Test; +import org.mockito.Mock; +import org.mockito.Mockito; +import org.mockito.junit.jupiter.MockitoExtension; +import org.junit.jupiter.api.extension.ExtendWith; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.verifyNoInteractions; + +@ExtendWith(MockitoExtension.class) +class OrTest { + + @Mock + Filter mockFilterTrue; + @Mock + Filter mockFilterFalse; + + @Test + void testShouldReturnTrueWhenLeftIsTrue() { + Mockito.when(mockFilterTrue.test(Mockito.any())).thenReturn(true); + + Or or = new Or(mockFilterTrue, mockFilterFalse); + + assertThat(or.test(new Object())).isTrue(); + verifyNoInteractions(mockFilterFalse); + } + + @Test + void testShouldReturnTrueWhenRightIsTrue() { + Mockito.when(mockFilterTrue.test(Mockito.any())).thenReturn(true); + Mockito.when(mockFilterFalse.test(Mockito.any())).thenReturn(false); + + Or or = new Or(mockFilterFalse, mockFilterTrue); + + assertThat(or.test(new Object())).isTrue(); + } + + @Test + void testShouldReturnFalseWhenBothAreFalse() { + Mockito.when(mockFilterFalse.test(Mockito.any())).thenReturn(false); + + Or or = new Or(mockFilterFalse, mockFilterFalse); + + assertThat(or.test(new Object())).isFalse(); + } + + @Test + void testEqualsMethodShouldReturnTrueForNonDistinctObjects() { + Or or1 = new Or(mockFilterTrue, mockFilterFalse); + Or or2 = new Or(mockFilterTrue, mockFilterFalse); + + assertThat(or1).isEqualTo(or2); + } + + @Test + void testEqualsAndHashCode() { + Or orFilter = new Or(mockFilterTrue, mockFilterFalse); + Or sameOrFilter = new Or(mockFilterTrue, mockFilterFalse); + + assertThat(orFilter) + .isEqualTo(sameOrFilter) + .hasSameHashCodeAs(sameOrFilter); + } + + @Test + void testToStringMethodShouldReturnExpectedValue() { + Or or = new Or(mockFilterTrue, mockFilterFalse); + + assertThat(or).hasToString("Or(left=" + mockFilterTrue + ", right=" + mockFilterFalse + ")"); + } +} diff --git a/langchain4j-parent/pom.xml b/langchain4j-parent/pom.xml index 3acd8d0749..47b2138f2a 100644 --- a/langchain4j-parent/pom.xml +++ b/langchain4j-parent/pom.xml @@ -18,6 +18,7 @@ ${java.version} UTF-8 1714382357 + 0.23.0 1.0.0-beta.11 @@ -33,7 +34,7 @@ 2.10.1 5.11.2 1.20.2 - 1.14.10 + 1.15.7 5.14.1 3.24.2 2.6.2 @@ -436,6 +437,14 @@ + + + org.mockito + mockito-core + test + + + @@ -493,6 +502,12 @@ org.apache.maven.plugins maven-dependency-plugin + + + properties + + + detect-unused-dependencies @@ -517,6 +532,7 @@ maven-surefire-plugin 3.3.1 + @{argLine} -javaagent:${org.mockito:mockito-core:jar} info