18c9354095
- 注入数据库、运行时配置、统一响应、文件存储和平台 AI 能力,补充业务读写工具与客户快捷操作契约。 - 移除模块内重复的组织、客户、工单、标签、技能、旧工作流、MCP 和迁移实现,将身份权限与业务主体交由宿主管理。 - 使用 libSQL 重构向量存储,并完善图片消息、访客身份、排队调度、企业微信和支持聊天页面。 - 统一 HTTP、DTO 与 WebSocket 的 snake_case 协议,补齐模块初始化、业务动作和公共载荷等回归测试。
100 lines
2.5 KiB
Go
100 lines
2.5 KiB
Go
package ai
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
|
|
openai "github.com/openai/openai-go/v3"
|
|
|
|
"code.tczkiot.com/wlw/ai-agent/internal/models"
|
|
"code.tczkiot.com/wlw/ai-agent/internal/pkg/enums"
|
|
"code.tczkiot.com/wlw/ai-agent/internal/pkg/errorsx"
|
|
)
|
|
|
|
type EmbeddingResult struct {
|
|
Vector []float32
|
|
TokensUsed int
|
|
ModelName string
|
|
Dimension int
|
|
}
|
|
|
|
type embedding struct{}
|
|
|
|
var Embedding = &embedding{}
|
|
|
|
func (s *embedding) GetModel(ctx context.Context) (*models.AIConfig, error) {
|
|
return resolveDefaultAIConfig(ctx, enums.AIModelTypeEmbedding)
|
|
}
|
|
|
|
func (s *embedding) GenerateEmbedding(ctx context.Context, text string) (*EmbeddingResult, error) {
|
|
if text == "" {
|
|
return nil, errorsx.InvalidParamI18n("error.e0215")
|
|
}
|
|
ctx = ensurePlatformAIRequestScope(ctx)
|
|
|
|
result, err := s.callEmbeddingAPI(ctx, text)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return result, nil
|
|
}
|
|
|
|
func (s *embedding) GenerateBatchEmbeddings(ctx context.Context, texts []string) ([]EmbeddingResult, error) {
|
|
if len(texts) == 0 {
|
|
return nil, errorsx.InvalidParamI18n("error.e0216")
|
|
}
|
|
ctx = ensurePlatformAIRequestScope(ctx)
|
|
|
|
results := make([]EmbeddingResult, 0, len(texts))
|
|
for _, text := range texts {
|
|
result, err := s.callEmbeddingAPI(ctx, text)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to generate embedding for text: %w", err)
|
|
}
|
|
results = append(results, *result)
|
|
}
|
|
|
|
return results, nil
|
|
}
|
|
|
|
func (s *embedding) callEmbeddingAPI(ctx context.Context, text string) (*EmbeddingResult, error) {
|
|
config, err := s.GetModel(ctx)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
client := newOpenAIClient(*config)
|
|
embeddingResp, err := client.Embeddings.New(ctx, openai.EmbeddingNewParams{
|
|
Input: openai.EmbeddingNewParamsInputUnion{
|
|
OfString: openai.String(text),
|
|
},
|
|
Model: openai.EmbeddingModel(config.ModelName),
|
|
}, platformRequestOptions(ctx, *config, "embedding")...)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to call embedding api: %w", err)
|
|
}
|
|
|
|
if len(embeddingResp.Data) == 0 {
|
|
return nil, fmt.Errorf("no embedding data in response")
|
|
}
|
|
vector := make([]float32, 0, len(embeddingResp.Data[0].Embedding))
|
|
for _, item := range embeddingResp.Data[0].Embedding {
|
|
vector = append(vector, float32(item))
|
|
}
|
|
|
|
return &EmbeddingResult{
|
|
Vector: vector,
|
|
TokensUsed: int(embeddingResp.Usage.TotalTokens),
|
|
ModelName: embeddingResp.Model,
|
|
Dimension: len(vector),
|
|
}, nil
|
|
}
|
|
|
|
func (s *embedding) GetDimension(ctx context.Context) (int, error) {
|
|
model, err := s.GetModel(ctx)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return model.Dimension, nil
|
|
}
|