Files
gochat/backend/internal/llm/eino_provider.go
T

273 lines
7.5 KiB
Go

package llm
import (
"context"
"fmt"
"time"
"github.com/cloudwego/eino/components/embedding"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
applogger "github.com/gochat/gochat/pkg/logger"
)
// EinoProvider implements the Provider interface by delegating to Eino's
// model.ChatModel + embedding.Embedder. This allows GoChat to use Eino's
// rich ecosystem of model implementations (OpenAI, Claude, Gemini, DeepSeek,
// Ark, Qwen, etc.) without changing any business logic code.
//
// The EinoProvider is constructed with:
// - chatModel: a model.ToolCallingChatModel (or model.ChatModel) from eino-ext
// - embedder: an embedding.Embedder from eino-ext
//
// If embedder is nil, CreateEmbedding returns an error. This happens when
// using providers like Anthropic that don't offer an embeddings API — in
// that case, configure a separate OpenAI-compatible embedding provider.
type EinoProvider struct {
chatModel model.BaseChatModel
embedder Embedder
// fallback embedder for providers that don't support embeddings (e.g. Anthropic)
// if non-nil, used when embedder is nil
fallbackEmbedder Embedder
}
// Embedder is the interface EinoProvider expects for embedding generation.
// This matches eino's embedding.Embedder interface exactly.
type Embedder = embedding.Embedder
// NewEinoProvider creates a new EinoProvider.
// chatModel must implement model.BaseChatModel (Generate + Stream).
// embedder is optional (can be nil if only chat completion is needed).
func NewEinoProvider(chatModel model.BaseChatModel, embedder Embedder) *EinoProvider {
return &EinoProvider{
chatModel: chatModel,
embedder: embedder,
}
}
// SetFallbackEmbedder sets a fallback embedder used when the primary embedder
// is nil (e.g. when using Anthropic as chat model but OpenAI for embeddings).
func (p *EinoProvider) SetFallbackEmbedder(embedder Embedder) {
p.fallbackEmbedder = embedder
}
// ChatCompletion sends a synchronous chat completion request via Eino.
func (p *EinoProvider) ChatCompletion(ctx context.Context, req ChatRequest) (*ChatResponse, error) {
// Convert GoChat messages to Eino schema.Message
messages := make([]*schema.Message, 0, len(req.Messages))
for _, msg := range req.Messages {
role := schema.RoleType(msg.Role)
if role == "" {
role = schema.User
}
messages = append(messages, &schema.Message{
Role: role,
Content: msg.Content,
})
// Handle tool call fields
if len(msg.ToolCalls) > 0 {
messages[len(messages)-1].ToolCalls = convertToSchemaToolCalls(msg.ToolCalls)
}
if msg.ToolCallID != "" {
messages[len(messages)-1].ToolCallID = msg.ToolCallID
messages[len(messages)-1].ToolName = msg.Name
}
}
// Build options
opts := []model.Option{}
if req.Temperature > 0 {
opts = append(opts, WithTemperature(req.Temperature))
}
if req.MaxTokens > 0 {
opts = append(opts, WithMaxTokens(req.MaxTokens))
}
// Call Eino model
resp, err := p.chatModel.Generate(ctx, messages, opts...)
if err != nil {
applogger.L().Errorf("EinoProvider ChatCompletion: %v", err)
return nil, fmt.Errorf("eino chat completion: %w", err)
}
// Convert Eino response to GoChat ChatResponse
result := &ChatResponse{
ID: "",
Object: "chat.completion",
Created: time.Now().Unix(),
Model: req.Model,
Choices: []ChatChoice{},
}
if resp != nil {
finishReason := ""
if resp.ResponseMeta != nil {
finishReason = resp.ResponseMeta.FinishReason
}
choice := ChatChoice{
Index: 0,
Message: ChatMessage{Role: string(resp.Role), Content: resp.Content},
FinishReason: finishReason,
}
// Convert tool calls if present
if len(resp.ToolCalls) > 0 {
choice.Message.ToolCalls = convertFromSchemaToolCalls(resp.ToolCalls)
}
result.Choices = append(result.Choices, choice)
// Map usage metadata if available
if resp.ResponseMeta != nil && resp.ResponseMeta.Usage != nil {
result.Usage = TokenUsage{
PromptTokens: resp.ResponseMeta.Usage.PromptTokens,
CompletionTokens: resp.ResponseMeta.Usage.CompletionTokens,
TotalTokens: resp.ResponseMeta.Usage.TotalTokens,
}
}
}
return result, nil
}
// ChatCompletionStream sends a streaming chat completion request via Eino.
// Chunks are delivered via the onChunk callback.
func (p *EinoProvider) ChatCompletionStream(ctx context.Context, req ChatRequest, onChunk func(StreamChunk) error) error {
// Convert messages
messages := make([]*schema.Message, 0, len(req.Messages))
for _, msg := range req.Messages {
role := schema.RoleType(msg.Role)
if role == "" {
role = schema.User
}
messages = append(messages, &schema.Message{
Role: role,
Content: msg.Content,
})
}
opts := []model.Option{}
if req.Temperature > 0 {
opts = append(opts, WithTemperature(req.Temperature))
}
if req.MaxTokens > 0 {
opts = append(opts, WithMaxTokens(req.MaxTokens))
}
reader, err := p.chatModel.Stream(ctx, messages, opts...)
if err != nil {
return fmt.Errorf("eino stream: %w", err)
}
if reader == nil {
return fmt.Errorf("eino stream returned nil reader")
}
defer reader.Close()
for {
chunk, err := reader.Recv()
if err != nil {
if err.Error() == "EOF" || err.Error() == "io: EOF" {
break
}
return fmt.Errorf("eino stream recv: %w", err)
}
if chunk == nil {
break
}
streamChunk := StreamChunk{
Model: req.Model,
Choices: []StreamChoice{
{
Index: 0,
Delta: StreamDelta{Content: chunk.Content},
},
},
}
if chunk.ResponseMeta != nil {
streamChunk.Choices[0].FinishReason = chunk.ResponseMeta.FinishReason
}
if err := onChunk(streamChunk); err != nil {
return fmt.Errorf("chunk callback: %w", err)
}
}
return nil
}
// CreateEmbedding generates embeddings using the Eino embedder.
func (p *EinoProvider) CreateEmbedding(ctx context.Context, req EmbeddingRequest) (*EmbeddingResponse, error) {
embedder := p.embedder
if embedder == nil {
embedder = p.fallbackEmbedder
}
if embedder == nil {
return nil, fmt.Errorf("no embedder configured for this provider")
}
vectors, err := embedder.EmbedStrings(ctx, req.Input)
if err != nil {
return nil, fmt.Errorf("eino embedding: %w", err)
}
data := make([]EmbeddingData, 0, len(vectors))
for i, vec := range vectors {
data = append(data, EmbeddingData{
Object: "embedding",
Index: i,
Embedding: vec,
})
}
return &EmbeddingResponse{
Object: "list",
Data: data,
Model: req.Model,
}, nil
}
// --- Helper functions for tool call conversion ---
func convertToSchemaToolCalls(calls []ToolCall) []schema.ToolCall {
result := make([]schema.ToolCall, 0, len(calls))
for _, c := range calls {
result = append(result, schema.ToolCall{
ID: c.ID,
Type: c.Type,
Function: schema.FunctionCall{
Name: c.Function.Name,
Arguments: c.Function.Arguments,
},
})
}
return result
}
func convertFromSchemaToolCalls(calls []schema.ToolCall) []ToolCall {
result := make([]ToolCall, 0, len(calls))
for _, c := range calls {
result = append(result, ToolCall{
ID: c.ID,
Type: c.Type,
Function: ToolCallFunction{
Name: c.Function.Name,
Arguments: c.Function.Arguments,
},
})
}
return result
}
// --- Model option helpers ---
// These wrap eino's model.Option to provide a simple API without importing
// eino's option package directly in business code.
// WithTemperature sets the temperature option for the model.
func WithTemperature(temp float64) model.Option {
return model.WithTemperature(float32(temp))
}
// WithMaxTokens sets the max tokens option for the model.
func WithMaxTokens(maxTokens int) model.Option {
return model.WithMaxTokens(maxTokens)
}