mirror of
https://github.com/vxcontrol/langchaingo.git
synced 2026-07-19 13:45:47 -04:00
googleai: fix user provided embedding model getting overridden
Fixes issue where user-provided embedding model was being overridden by default model
This commit is contained in:
committed by
GitHub
parent
78296871ac
commit
ea5209cc18
@@ -0,0 +1,64 @@
|
||||
package palmclient
|
||||
|
||||
import (
|
||||
"google.golang.org/api/option"
|
||||
)
|
||||
|
||||
const (
|
||||
embeddingModelName = "text-embedding-005"
|
||||
TextModelName = "text-bison"
|
||||
ChatModelName = "chat-bison"
|
||||
)
|
||||
|
||||
// Options are the Palm client options
|
||||
type Options struct {
|
||||
EmbeddingModelName string
|
||||
TextModelName string
|
||||
ChatModelName string
|
||||
ClientOptions []option.ClientOption
|
||||
}
|
||||
|
||||
// Option is an option
|
||||
type Option func(*Options)
|
||||
|
||||
// WithEmbeddingModelName sets the default embedding model
|
||||
func WithEmbeddingModelName(modelName string) Option {
|
||||
return func(o *Options) {
|
||||
if modelName != "" {
|
||||
o.EmbeddingModelName = modelName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithTextModelName sets the default text model
|
||||
func WithTextModelName(modelName string) Option {
|
||||
return func(o *Options) {
|
||||
if modelName != "" {
|
||||
o.TextModelName = modelName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithChatModelName sets the default chat model
|
||||
func WithChatModelName(modelName string) Option {
|
||||
return func(o *Options) {
|
||||
if modelName != "" {
|
||||
o.ChatModelName = modelName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithClientOptions sets the client options for the Google API client
|
||||
func WithClientOptions(opts ...option.ClientOption) Option {
|
||||
return func(o *Options) {
|
||||
o.ClientOptions = append(o.ClientOptions, opts...)
|
||||
}
|
||||
}
|
||||
|
||||
func defaultOptions() Options {
|
||||
return Options{
|
||||
EmbeddingModelName: embeddingModelName,
|
||||
TextModelName: TextModelName,
|
||||
ChatModelName: ChatModelName,
|
||||
}
|
||||
}
|
||||
@@ -28,40 +28,48 @@ var defaultParameters = map[string]interface{}{ //nolint:gochecknoglobals
|
||||
}
|
||||
|
||||
const (
|
||||
embeddingModelName = "text-embedding-005"
|
||||
TextModelName = "text-bison"
|
||||
ChatModelName = "chat-bison"
|
||||
|
||||
defaultMaxConns = 4
|
||||
)
|
||||
|
||||
// PaLMClient represents a Vertex AI based PaLM API client.
|
||||
type PaLMClient struct {
|
||||
client *aiplatform.PredictionClient
|
||||
projectID string
|
||||
client *aiplatform.PredictionClient
|
||||
projectID string
|
||||
embeddingModelName string
|
||||
textModelName string
|
||||
chatModelName string
|
||||
}
|
||||
|
||||
// New returns a new Vertex AI based PaLM API client.
|
||||
func New(ctx context.Context, projectID, location string, opts ...option.ClientOption) (*PaLMClient, error) {
|
||||
func New(ctx context.Context, projectID, location string, opts ...Option) (*PaLMClient, error) {
|
||||
numConns := runtime.GOMAXPROCS(0)
|
||||
if numConns > defaultMaxConns {
|
||||
numConns = defaultMaxConns
|
||||
}
|
||||
|
||||
pOpt := defaultOptions()
|
||||
for _, o := range opts {
|
||||
o(&pOpt)
|
||||
}
|
||||
|
||||
o := []option.ClientOption{
|
||||
option.WithGRPCConnectionPool(numConns),
|
||||
option.WithEndpoint(fmt.Sprintf("%s-aiplatform.googleapis.com:443", location)),
|
||||
}
|
||||
opts = append(o, opts...)
|
||||
pOpt.ClientOptions = append(o, pOpt.ClientOptions...)
|
||||
// PredictionClient only support GRPC.
|
||||
opts = append(opts, option.WithHTTPClient(nil))
|
||||
pOpt.ClientOptions = append(pOpt.ClientOptions, option.WithHTTPClient(nil))
|
||||
|
||||
client, err := aiplatform.NewPredictionClient(ctx, opts...)
|
||||
client, err := aiplatform.NewPredictionClient(ctx, pOpt.ClientOptions...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PaLMClient{
|
||||
client: client,
|
||||
projectID: projectID,
|
||||
client: client,
|
||||
projectID: projectID,
|
||||
embeddingModelName: pOpt.EmbeddingModelName,
|
||||
textModelName: pOpt.TextModelName,
|
||||
chatModelName: pOpt.ChatModelName,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -92,7 +100,7 @@ func (c *PaLMClient) CreateCompletion(ctx context.Context, r *CompletionRequest)
|
||||
"top_k": r.TopK,
|
||||
"stopSequences": convertArray(r.StopSequences),
|
||||
}
|
||||
predictions, err := c.batchPredict(ctx, TextModelName, r.Prompts, params)
|
||||
predictions, err := c.batchPredict(ctx, c.textModelName, r.Prompts, params)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -118,7 +126,7 @@ type EmbeddingRequest struct {
|
||||
// CreateEmbedding creates embeddings.
|
||||
func (c *PaLMClient) CreateEmbedding(ctx context.Context, r *EmbeddingRequest) ([][]float32, error) {
|
||||
params := map[string]interface{}{}
|
||||
responses, err := c.batchPredict(ctx, embeddingModelName, r.Input, params)
|
||||
responses, err := c.batchPredict(ctx, c.embeddingModelName, r.Input, params)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -327,7 +335,7 @@ func (c *PaLMClient) chat(ctx context.Context, r *ChatRequest) ([]*structpb.Valu
|
||||
structpb.NewStructValue(instance),
|
||||
}
|
||||
resp, err := c.client.Predict(ctx, &aiplatformpb.PredictRequest{
|
||||
Endpoint: c.projectLocationPublisherModelPath(c.projectID, "us-central1", "google", ChatModelName),
|
||||
Endpoint: c.projectLocationPublisherModelPath(c.projectID, "us-central1", "google", c.chatModelName),
|
||||
Instances: instances,
|
||||
Parameters: structpb.NewStructValue(mergedParams),
|
||||
})
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
|
||||
"cloud.google.com/go/vertexai/genai"
|
||||
"google.golang.org/api/option"
|
||||
"google.golang.org/grpc"
|
||||
)
|
||||
|
||||
// Options is a set of options for GoogleAI and Vertex clients.
|
||||
@@ -99,6 +100,15 @@ func WithHTTPClient(httpClient *http.Client) Option {
|
||||
}
|
||||
}
|
||||
|
||||
// WithGRPCConn appends a ClientOption that uses the provided gRPC client connection to
|
||||
// make requests.
|
||||
// This is useful for testing embeddings in vertex clients.
|
||||
func WithGRPCConn(conn *grpc.ClientConn) Option {
|
||||
return func(opts *Options) {
|
||||
opts.ClientOptions = append(opts.ClientOptions, option.WithGRPCConn(conn))
|
||||
}
|
||||
}
|
||||
|
||||
// WithCloudProject passes the GCP cloud project name to the client. This is
|
||||
// useful for vertex clients.
|
||||
func WithCloudProject(p string) Option {
|
||||
|
||||
@@ -116,5 +116,9 @@ func newClient(opts ...Option) (*palmclient.PaLMClient, error) {
|
||||
return nil, ErrMissingLocation
|
||||
}
|
||||
|
||||
return palmclient.New(context.TODO(), options.projectID, options.location, options.clientOptions...)
|
||||
palmOptions := []palmclient.Option{
|
||||
palmclient.WithClientOptions(options.clientOptions...),
|
||||
}
|
||||
|
||||
return palmclient.New(context.TODO(), options.location, options.projectID, palmOptions...)
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cloud.google.com/go/aiplatform/apiv1/aiplatformpb"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tmc/langchaingo/embeddings"
|
||||
@@ -24,6 +26,11 @@ import (
|
||||
"github.com/tmc/langchaingo/llms"
|
||||
"github.com/tmc/langchaingo/llms/googleai"
|
||||
"github.com/tmc/langchaingo/llms/googleai/vertex"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/grpc/test/bufconn"
|
||||
"google.golang.org/protobuf/types/known/structpb"
|
||||
)
|
||||
|
||||
func newGoogleAIClient(t *testing.T, opts ...googleai.Option) *googleai.GoogleAI {
|
||||
@@ -160,6 +167,21 @@ func TestVertexShared(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestVertex_WithCustomEmbeddingModel tests custom embedding models passed as an option.
|
||||
// TODO: refactor testConfig to have a opts provider func so it this can be moved to a test config.
|
||||
func TestVertex_WithCustomEmbeddingModel(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Parallel()
|
||||
const modelName = "custom-embedding-model"
|
||||
opts := getCustomEmbeddingModelTestOptionsWithGRPC(t, modelName)
|
||||
|
||||
llm, err := vertex.New(context.Background(), opts...)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = llm.CreateEmbedding(context.Background(), []string{"test"})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func testMultiContentText(t *testing.T, llm llms.Model) {
|
||||
t.Helper()
|
||||
t.Parallel()
|
||||
@@ -617,6 +639,41 @@ func getHTTPTestClientOptions() []googleai.Option {
|
||||
return []googleai.Option{googleai.WithRest(), googleai.WithHTTPClient(client)}
|
||||
}
|
||||
|
||||
// getCustomEmbeddingModelTestOptionsWithGRPC creates options to connect to a fake gRPC server.
|
||||
func getCustomEmbeddingModelTestOptionsWithGRPC(t *testing.T, model string) []googleai.Option {
|
||||
t.Helper()
|
||||
|
||||
// Create an in-memory "network connection"
|
||||
lis := bufconn.Listen(1024 * 1024)
|
||||
// Create the mock gRPC server and register the fake prediction service
|
||||
grpcServer := grpc.NewServer()
|
||||
mockPredictionServer := &mockPredictionServer{
|
||||
predictFunc: getPredictHandlerFuncWithCustomEmbeddingModel(model)}
|
||||
aiplatformpb.RegisterPredictionServiceServer(grpcServer, mockPredictionServer)
|
||||
|
||||
go func() {
|
||||
if err := grpcServer.Serve(lis); err != nil {
|
||||
t.Logf("gRPC server exited with error: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
t.Cleanup(func() { grpcServer.Stop() })
|
||||
|
||||
// Create a client connection to the fake server
|
||||
conn, err := grpc.DialContext(context.Background(), "bufnet",
|
||||
grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) {
|
||||
return lis.Dial()
|
||||
}),
|
||||
grpc.WithInsecure(),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
return []googleai.Option{
|
||||
googleai.WithDefaultEmbeddingModel(model),
|
||||
googleai.WithGRPCConn(conn),
|
||||
}
|
||||
}
|
||||
|
||||
type testRequestInterceptor struct{}
|
||||
|
||||
func (i *testRequestInterceptor) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
@@ -666,3 +723,57 @@ func checkMatch(t *testing.T, got string, wants ...string) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// PredictHandlerFunc is a handler func that matches the gRPC Predict method.
|
||||
type PredictHandlerFunc func(context.Context, *aiplatformpb.PredictRequest) (*aiplatformpb.PredictResponse, error)
|
||||
|
||||
// mockPredictionServer is a mock gRPC prediction server.
|
||||
type mockPredictionServer struct {
|
||||
aiplatformpb.UnimplementedPredictionServiceServer
|
||||
predictFunc PredictHandlerFunc
|
||||
}
|
||||
|
||||
// NewMockPredictionServer creates a new server with a custom Predict handler.
|
||||
func NewMockPredictionServer(handler PredictHandlerFunc) *mockPredictionServer {
|
||||
return &mockPredictionServer{
|
||||
predictFunc: handler,
|
||||
}
|
||||
}
|
||||
|
||||
// Predict implements the UnimplementedPredictionServiceServer. It calls predictFunc.
|
||||
func (s *mockPredictionServer) Predict(ctx context.Context, req *aiplatformpb.PredictRequest) (*aiplatformpb.PredictResponse, error) {
|
||||
if s.predictFunc == nil {
|
||||
return nil, status.Error(codes.Unimplemented, "Predict handler was not provided")
|
||||
}
|
||||
|
||||
return s.predictFunc(ctx, req)
|
||||
}
|
||||
|
||||
// getPredictHandlerFuncWithCustomEmbeddingModel returns a predictFunc which checks that the embedding request has been made
|
||||
// with the custom provided embedding model.
|
||||
func getPredictHandlerFuncWithCustomEmbeddingModel(embeddingModel string) func(ctx context.Context, req *aiplatformpb.PredictRequest) (*aiplatformpb.PredictResponse, error) {
|
||||
return func(ctx context.Context, req *aiplatformpb.PredictRequest) (*aiplatformpb.PredictResponse, error) {
|
||||
expectedEndpointSuffix := "/models/" + embeddingModel
|
||||
if !strings.HasSuffix(req.Endpoint, expectedEndpointSuffix) {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "model name mismatch, expected suffix '%s'", expectedEndpointSuffix)
|
||||
}
|
||||
|
||||
// Create a dummy embedding response that the client code will accept.
|
||||
// The client expects a structure like: { "embeddings": { "values": [0.1, 0.2, ...] } }
|
||||
embeddingStruct, err := structpb.NewStruct(map[string]interface{}{
|
||||
"embeddings": map[string]interface{}{
|
||||
"values": []interface{}{0.1, 0.2, 0.3, 0.4},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to create embedding struct: %v", err)
|
||||
}
|
||||
|
||||
predictionValue := structpb.NewStructValue(embeddingStruct)
|
||||
response := &aiplatformpb.PredictResponse{
|
||||
Predictions: []*structpb.Value{predictionValue},
|
||||
}
|
||||
|
||||
return response, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,11 +42,15 @@ func New(ctx context.Context, opts ...googleai.Option) (*Vertex, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
palmOpts := []palmclient.Option{
|
||||
palmclient.WithEmbeddingModelName(clientOptions.DefaultEmbeddingModel),
|
||||
palmclient.WithClientOptions(clientOptions.ClientOptions...),
|
||||
}
|
||||
palmClient, err := palmclient.New(
|
||||
ctx,
|
||||
clientOptions.CloudProject,
|
||||
clientOptions.CloudLocation,
|
||||
clientOptions.ClientOptions...)
|
||||
palmOpts...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user