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:
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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")));
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user