Refactoring of ChatClient to add fluent API and introduce Model as dependent object

* Rename the ModelClient class hierarchy into Model:
  - Rename ModelClient into Model. Update all code and doc references.
  - Rename ChatClient to ChatModel. Update all ChatClient suffixes and chatClient fields and variables in code and doc.
  - Rename EmbeddingClient into EmbeddingModel. Update the XxxEmbeddingClient class and variable suffixes and embeddingClient variables and fields in code and docs.
  - Rename ImageClient into ImageModel.
  - Rename SpeechClient into SpeechModel.
  - Rename TranscriptionClient into TranscriptionModel.
  - Update all javadocs and antora pages. Update the related diagrams.

* Create fluent API in ChatClient interface that now includes streaming support
* Add OpenAI FunctionCallbackWrapper2IT auto-config tests.
* Add ChatClientTest mockito testing.
* Add ChatModel#getDefaultOptions(), and remove @FunctionalInterface

* ChatModel enums extend the new ModelDescription interface.
* Implement fromOptions copy method in every ChatOptions implementation.
* Extend ChatClient to use the model default options if not provided explicitly.

* Update readme to provide guidance on how to adapt to breaking changes.

Co-authored-by: Christian Tzolov <ctzolov@vmware.com>
Co-authored-by: Mark Pollack <mpollack@vmware.com>
This commit is contained in:
Josh Long
2024-05-18 21:30:51 +02:00
committed by Mark Pollack
parent bce45c2d2f
commit fbfc87e814
426 changed files with 4815 additions and 2912 deletions

View File

@@ -32,7 +32,7 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.filter.converter.PgVectorFilterExpressionConverter;
import org.springframework.beans.factory.InitializingBean;
@@ -66,7 +66,7 @@ public class PgVectorStore implements VectorStore, InitializingBean {
private final JdbcTemplate jdbcTemplate;
private final EmbeddingClient embeddingClient;
private final EmbeddingModel embeddingModel;
private int dimensions;
@@ -197,21 +197,21 @@ public class PgVectorStore implements VectorStore, InitializingBean {
}
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) {
this(jdbcTemplate, embeddingClient, INVALID_EMBEDDING_DIMENSION, PgVectorStore.PgDistanceType.COSINE_DISTANCE,
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel) {
this(jdbcTemplate, embeddingModel, INVALID_EMBEDDING_DIMENSION, PgVectorStore.PgDistanceType.COSINE_DISTANCE,
false, PgIndexType.NONE);
}
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions) {
this(jdbcTemplate, embeddingClient, dimensions, PgVectorStore.PgDistanceType.COSINE_DISTANCE, false,
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, int dimensions) {
this(jdbcTemplate, embeddingModel, dimensions, PgVectorStore.PgDistanceType.COSINE_DISTANCE, false,
PgIndexType.NONE);
}
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions,
public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, int dimensions,
PgDistanceType distanceType, boolean removeExistingVectorStoreTable, PgIndexType createIndexMethod) {
this.jdbcTemplate = jdbcTemplate;
this.embeddingClient = embeddingClient;
this.embeddingModel = embeddingModel;
this.dimensions = dimensions;
this.distanceType = distanceType;
this.removeExistingVectorStoreTable = removeExistingVectorStoreTable;
@@ -237,7 +237,7 @@ public class PgVectorStore implements VectorStore, InitializingBean {
var document = documents.get(i);
var content = document.getContent();
var json = toJson(document.getMetadata());
var pGvector = new PGvector(toFloatArray(embeddingClient.embed(document)));
var pGvector = new PGvector(toFloatArray(embeddingModel.embed(document)));
StatementCreatorUtils.setParameterValue(ps, 1, SqlTypeValue.TYPE_UNKNOWN,
UUID.fromString(document.getId()));
@@ -320,7 +320,7 @@ public class PgVectorStore implements VectorStore, InitializingBean {
}
private PGvector getQueryEmbedding(String query) {
List<Double> embedding = this.embeddingClient.embed(query);
List<Double> embedding = this.embeddingModel.embed(query);
return new PGvector(toFloatArray(embedding));
}
@@ -366,13 +366,13 @@ public class PgVectorStore implements VectorStore, InitializingBean {
}
try {
int embeddingDimensions = this.embeddingClient.dimensions();
int embeddingDimensions = this.embeddingModel.dimensions();
if (embeddingDimensions > 0) {
return embeddingDimensions;
}
}
catch (Exception e) {
logger.warn("Failed to obtain the embedding dimensions from the embedding client and fall backs to default:"
logger.warn("Failed to obtain the embedding dimensions from the embedding model and fall backs to default:"
+ OPENAI_EMBEDDING_DIMENSION_SIZE, e);
}
return OPENAI_EMBEDDING_DIMENSION_SIZE;

View File

@@ -20,7 +20,7 @@ import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.jdbc.core.JdbcTemplate;
import static org.assertj.core.api.Assertions.assertThat;
@@ -36,7 +36,7 @@ import static org.mockito.Mockito.when;
public class PgVectorEmbeddingDimensionsTests {
@Mock
private EmbeddingClient embeddingClient;
private EmbeddingModel embeddingModel;
@Mock
private JdbcTemplate jdbcTemplate;
@@ -46,32 +46,32 @@ public class PgVectorEmbeddingDimensionsTests {
final int explicitDimensions = 696;
var dim = new PgVectorStore(jdbcTemplate, embeddingClient, explicitDimensions).embeddingDimensions();
var dim = new PgVectorStore(jdbcTemplate, embeddingModel, explicitDimensions).embeddingDimensions();
assertThat(dim).isEqualTo(explicitDimensions);
verify(embeddingClient, never()).dimensions();
verify(embeddingModel, never()).dimensions();
}
@Test
public void embeddingClientDimensions() {
when(embeddingClient.dimensions()).thenReturn(969);
public void embeddingModelDimensions() {
when(embeddingModel.dimensions()).thenReturn(969);
var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions();
var dim = new PgVectorStore(jdbcTemplate, embeddingModel).embeddingDimensions();
assertThat(dim).isEqualTo(969);
verify(embeddingClient, only()).dimensions();
verify(embeddingModel, only()).dimensions();
}
@Test
public void fallBackToDefaultDimensions() {
when(embeddingClient.dimensions()).thenThrow(new RuntimeException());
when(embeddingModel.dimensions()).thenThrow(new RuntimeException());
var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions();
var dim = new PgVectorStore(jdbcTemplate, embeddingModel).embeddingDimensions();
assertThat(dim).isEqualTo(PgVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE);
verify(embeddingClient, only()).dimensions();
verify(embeddingModel, only()).dimensions();
}
}

View File

@@ -35,9 +35,9 @@ import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.OpenAiEmbeddingClient;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.vectorstore.PgVectorStore.PgIndexType;
import org.springframework.ai.vectorstore.filter.FilterExpressionTextParser.FilterExpressionParseException;
import org.springframework.beans.factory.annotation.Value;
@@ -306,8 +306,8 @@ public class PgVectorStoreIT {
PgVectorStore.PgDistanceType distanceType;
@Bean
public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) {
return new PgVectorStore(jdbcTemplate, embeddingClient, PgVectorStore.INVALID_EMBEDDING_DIMENSION,
public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel) {
return new PgVectorStore(jdbcTemplate, embeddingModel, PgVectorStore.INVALID_EMBEDDING_DIMENSION,
distanceType, true, PgIndexType.HNSW);
}
@@ -329,8 +329,8 @@ public class PgVectorStoreIT {
}
@Bean
public EmbeddingClient embeddingClient() {
return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}