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
-[](https://github.com/langchain4j/langchain4j/actions/workflows/main.yaml)
+[](https://github.com/langchain4j/langchain4j/actions/workflows/main.yaml)
+[](https://github.com/langchain4j/langchain4j/actions/workflows/nightly.yaml)
+[](https://app.codacy.com/gh/langchain4j/langchain4j/dashboard)
+
[](https://discord.gg/JzTFvyjG6R)
[](https://x.com/langchain4j)
[](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