-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdb.go
More file actions
410 lines (347 loc) · 11.4 KB
/
Copy pathdb.go
File metadata and controls
410 lines (347 loc) · 11.4 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
// db.go
package knowledgesdk
import (
"context"
"fmt"
"log"
"strings"
"time"
"github.com/sashabaranov/go-openai"
"github.com/spf13/viper"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// Config 表示SDK配置
type Config struct {
// 数据库配置
DBHost string
DBPort int
DBName string
DBUser string
DBPassword string
// 向量嵌入服务配置
APIKey string
BaseURL string // 用于兼容不同的模型服务
EmbeddingModel string // 如"text-embedding-ada-002"
}
// KnowledgeSDK 是主SDK结构体
type KnowledgeSDK struct {
config Config
db *gorm.DB
openaiClient *openai.Client
modelDimension int // 存储模型向量维度
}
// NewKnowledgeSDK 创建一个新的SDK实例
func NewKnowledgeSDK(config *Config) (*KnowledgeSDK, error) {
// 连接数据库
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=disable",
config.DBHost, config.DBPort, config.DBUser, config.DBPassword, config.DBName)
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
Logger: logger.Default.LogMode(logger.Warn),
})
if err != nil {
return nil, fmt.Errorf("连接数据库失败: %w", err)
}
// 设置连接池参数
sqlDB, err := db.DB()
if err != nil {
return nil, fmt.Errorf("获取DB实例失败: %w", err)
}
sqlDB.SetMaxIdleConns(10)
sqlDB.SetMaxOpenConns(100)
sqlDB.SetConnMaxLifetime(time.Hour)
// 创建OpenAI客户端
clientConfig := openai.DefaultConfig(config.APIKey)
if config.BaseURL != "" {
clientConfig.BaseURL = config.BaseURL
}
client := openai.NewClientWithConfig(clientConfig)
// 初始化SDK
sdk := &KnowledgeSDK{
config: *config,
db: db,
openaiClient: client,
}
// 初始化数据库
if err := sdk.initDatabase(); err != nil {
return nil, fmt.Errorf("初始化数据库失败: %w", err)
}
// 检测并设置向量维度
dim, err := sdk.detectEmbeddingDimension()
if err != nil {
return nil, fmt.Errorf("检测向量维度失败: %w", err)
}
sdk.modelDimension = dim
return sdk, nil
}
// initDatabase 初始化数据库
func (k *KnowledgeSDK) initDatabase() error {
// 启用必要的PostgreSQL扩展
if err := k.db.Exec("CREATE EXTENSION IF NOT EXISTS vector").Error; err != nil {
return fmt.Errorf("创建vector扩展失败: %w", err)
}
if err := k.db.Exec("CREATE EXTENSION IF NOT EXISTS pg_bigm").Error; err != nil {
return fmt.Errorf("创建pg_bigm扩展失败: %w", err)
}
k.db.AutoMigrate(&KnowledgeBase{}, &Document{}, &Chunk{})
// 创建全文索引
if err := k.db.Exec(`
CREATE INDEX IF NOT EXISTS idx_chunk_content_bigm
ON knowledge_chunks USING gin (content gin_bigm_ops)
`).Error; err != nil {
return fmt.Errorf("创建全文索引失败: %w", err)
}
// 创建向量索引
if err := k.db.Exec(`
CREATE INDEX IF NOT EXISTS idx_chunk_embedding_hnsw
ON knowledge_chunks USING hnsw (embedding vector_cosine_ops)
`).Error; err != nil {
return fmt.Errorf("创建向量索引失败: %w", err)
}
return nil
}
// EmbeddingToPgVector 将向量转换为PostgresSQL向量格式
func (k *KnowledgeSDK) EmbeddingToPgVector(embedding []float32) string {
if len(embedding) == 0 {
return "[]"
}
parts := make([]string, len(embedding))
for i, v := range embedding {
parts[i] = fmt.Sprintf("%.6f", v)
}
return "[" + strings.Join(parts, ",") + "]"
}
// detectEmbeddingDimension 检测嵌入模型的向量维度
func (k *KnowledgeSDK) detectEmbeddingDimension() (int, error) {
// 使用OpenAI客户端获取向量维度
resp, err := k.openaiClient.CreateEmbeddings(
context.Background(),
openai.EmbeddingRequest{
Input: []string{"测试文本"},
Model: openai.EmbeddingModel(k.config.EmbeddingModel),
},
)
if err != nil {
return 0, fmt.Errorf("获取向量维度失败: %w", err)
}
if len(resp.Data) == 0 || len(resp.Data[0].Embedding) == 0 {
return 0, fmt.Errorf("API返回空向量")
}
modelDim := len(resp.Data[0].Embedding)
// 检查数据库向量维度
var dbDim int
err = k.db.Raw(`
SELECT a.atttypmod
FROM pg_attribute a
JOIN pg_class c ON a.attrelid = c.oid
JOIN pg_namespace n ON c.relnamespace = n.oid
WHERE c.relname = 'knowledge_chunks' AND a.attname = 'embedding' AND n.nspname = 'public'
`).Scan(&dbDim).Error
if err != nil {
return 0, fmt.Errorf("查询数据库向量维度失败: %w", err)
}
// 如果数据库向量维度与模型不匹配,则更新
if dbDim != modelDim && dbDim > 0 {
log.Printf("数据库向量维度 (%d) 与模型维度 (%d) 不匹配,正在更新...", dbDim, modelDim)
// 开始事务
tx := k.db.Begin()
if tx.Error != nil {
return 0, fmt.Errorf("开始事务失败: %w", tx.Error)
}
// 创建临时表,使用新的向量维度
err = tx.Exec(fmt.Sprintf(`
CREATE TABLE knowledge_chunks_temp (
document_id UUID NOT NULL,
chunk_index INTEGER NOT NULL,
content TEXT NOT NULL,
embedding VECTOR(%d),
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
last_indexed_at TIMESTAMPTZ,
PRIMARY KEY (document_id, chunk_index)
)
`, modelDim)).Error
if err != nil {
tx.Rollback()
return 0, fmt.Errorf("创建临时表失败: %w", err)
}
// 复制数据到临时表(除了向量数据)
err = tx.Exec(`
INSERT INTO knowledge_chunks_temp (document_id, chunk_index, content, updated_at, last_indexed_at)
SELECT document_id, chunk_index, content, updated_at, last_indexed_at FROM knowledge_chunks
`).Error
if err != nil {
tx.Rollback()
return 0, fmt.Errorf("复制数据到临时表失败: %w", err)
}
// 删除旧表
err = tx.Exec(`DROP TABLE knowledge_chunks`).Error
if err != nil {
tx.Rollback()
return 0, fmt.Errorf("删除原表失败: %w", err)
}
// 重命名临时表
err = tx.Exec(`ALTER TABLE knowledge_chunks_temp RENAME TO knowledge_chunks`).Error
if err != nil {
tx.Rollback()
return 0, fmt.Errorf("重命名临时表失败: %w", err)
}
// 尝试重新创建外键约束
err = tx.Exec(`
ALTER TABLE knowledge_chunks
ADD CONSTRAINT fk_document
FOREIGN KEY (document_id) REFERENCES documents(document_id) ON DELETE CASCADE
`).Error
if err != nil {
log.Printf("警告: 无法重新创建外键约束: %v", err)
// 即使无法添加外键约束,我们也继续执行
}
// 重新创建索引
err = tx.Exec(`
CREATE INDEX idx_chunk_content_bigm ON knowledge_chunks USING gin (content gin_bigm_ops);
CREATE INDEX idx_chunk_embedding_hnsw ON knowledge_chunks USING hnsw (embedding vector_cosine_ops);
`).Error
if err != nil {
tx.Rollback()
return 0, fmt.Errorf("重新创建索引失败: %w", err)
}
// 提交事务
if err := tx.Commit().Error; err != nil {
return 0, fmt.Errorf("提交事务失败: %w", err)
}
log.Printf("数据库向量维度已成功更新为 %d", modelDim)
}
return modelDim, nil
}
// GetDB 获取GORM数据库连接
func (k *KnowledgeSDK) GetDB() *gorm.DB {
return k.db
}
// GetOpenAIClient 获取OpenAI客户端
func (k *KnowledgeSDK) GetOpenAIClient() *openai.Client {
return k.openaiClient
}
// GetEmbeddingModel 获取当前使用的嵌入模型名称
func (k *KnowledgeSDK) GetEmbeddingModel() string {
return k.config.EmbeddingModel
}
// GetModelDimension 获取模型向量维度
func (k *KnowledgeSDK) GetModelDimension() int {
return k.modelDimension
}
// NewKnowledgeSDKFromConfig 从配置文件创建SDK实例
func NewKnowledgeSDKFromConfig(configPath string) (*KnowledgeSDK, error) {
// 使用viper加载配置
cfg, err := loadConfigWithViper(configPath)
if err != nil {
return nil, fmt.Errorf("加载配置失败: %w", err)
}
// 创建SDK配置
sdkConfig := Config{
DBHost: cfg.GetString("pgvector.host"),
DBPort: cfg.GetInt("pgvector.port"),
DBName: cfg.GetString("pgvector.db_name"),
DBUser: cfg.GetString("pgvector.user"),
DBPassword: cfg.GetString("pgvector.password"),
APIKey: cfg.GetString("embedding.api_key"),
BaseURL: cfg.GetString("embedding.endpoint"),
EmbeddingModel: cfg.GetString("embedding.model_name"),
}
return NewKnowledgeSDK(&sdkConfig)
}
// NewKnowledgeSDKFromEnv 从环境变量创建SDK实例
func NewKnowledgeSDKFromEnv() (*KnowledgeSDK, error) {
// 使用viper加载环境变量
cfg := viper.New()
// 设置环境变量前缀
cfg.SetEnvPrefix("KNOWLEDGE")
cfg.AutomaticEnv()
// 替换环境变量中的分隔符
cfg.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
// 读取.env文件
cfg.SetConfigFile(".env")
cfg.ReadInConfig() // 忽略错误,.env文件可能不存在
// 设置默认值
setDefaultValues(cfg)
// 创建SDK配置
sdkConfig := Config{
DBHost: cfg.GetString("pgvector.host"),
DBPort: cfg.GetInt("pgvector.port"),
DBName: cfg.GetString("pgvector.db_name"),
DBUser: cfg.GetString("pgvector.user"),
DBPassword: cfg.GetString("pgvector.password"),
APIKey: cfg.GetString("embedding.api_key"),
BaseURL: cfg.GetString("embedding.endpoint"),
EmbeddingModel: cfg.GetString("embedding.model_name"),
}
return NewKnowledgeSDK(&sdkConfig)
}
// loadConfigWithViper 使用viper加载配置文件
func loadConfigWithViper(configPath string) (*viper.Viper, error) {
cfg := viper.New()
// 设置配置文件路径和名称
if configPath != "" {
cfg.SetConfigFile(configPath)
} else {
cfg.SetConfigName("config")
cfg.SetConfigType("yaml")
cfg.AddConfigPath(".")
cfg.AddConfigPath("./config")
cfg.AddConfigPath("/etc/knowledgesdk")
}
// 支持环境变量
cfg.SetEnvPrefix("KNOWLEDGE")
cfg.AutomaticEnv()
cfg.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
// 读取配置文件
if err := cfg.ReadInConfig(); err != nil {
return nil, fmt.Errorf("读取配置文件失败: %w", err)
}
// 读取.env文件作为补充
envCfg := viper.New()
envCfg.SetConfigFile(".env")
if err := envCfg.ReadInConfig(); err == nil {
// 将.env中的配置合并到主配置中
for _, key := range envCfg.AllKeys() {
if !cfg.IsSet(key) {
cfg.Set(key, envCfg.Get(key))
}
}
}
// 设置默认值
setDefaultValues(cfg)
return cfg, nil
}
// setDefaultValues 设置默认配置值
func setDefaultValues(cfg *viper.Viper) {
// 服务器默认配置
cfg.SetDefault("server.host", "0.0.0.0")
cfg.SetDefault("server.port", 8080)
cfg.SetDefault("server.debug", false)
// 数据库默认配置
cfg.SetDefault("pgvector.host", "localhost")
cfg.SetDefault("pgvector.port", 5432)
cfg.SetDefault("pgvector.user", "postgres")
cfg.SetDefault("pgvector.password", "password")
cfg.SetDefault("pgvector.db_name", "knowledge_base")
cfg.SetDefault("pgvector.ssl_mode", "disable")
// 嵌入服务默认配置
cfg.SetDefault("embedding.model_name", "text-embedding-ada-002")
cfg.SetDefault("embedding.endpoint", "https://api.openai.com/v1")
// LLM默认配置
cfg.SetDefault("llm.temperature", 0.7)
cfg.SetDefault("llm.max_tokens", 1000)
}
// GetConfigFromViper 从viper获取配置值的辅助函数
func GetConfigFromViper(cfg *viper.Viper) Config {
return Config{
DBHost: cfg.GetString("pgvector.host"),
DBPort: cfg.GetInt("pgvector.port"),
DBName: cfg.GetString("pgvector.db_name"),
DBUser: cfg.GetString("pgvector.user"),
DBPassword: cfg.GetString("pgvector.password"),
APIKey: cfg.GetString("embedding.api_key"),
BaseURL: cfg.GetString("embedding.endpoint"),
EmbeddingModel: cfg.GetString("embedding.model_name"),
}
}