87 lines
2.1 KiB
Go
87 lines
2.1 KiB
Go
|
|
package main
|
||
|
|
|
||
|
|
import (
|
||
|
|
"database/sql"
|
||
|
|
"fmt"
|
||
|
|
"log"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
_ "github.com/go-sql-driver/mysql"
|
||
|
|
)
|
||
|
|
|
||
|
|
type PromptDB struct {
|
||
|
|
db *sql.DB
|
||
|
|
}
|
||
|
|
|
||
|
|
// NewPromptDB 创建 MySQL 连接并自动建表
|
||
|
|
// dsn 格式: user:password@tcp(host:port)/dbname?parseTime=true
|
||
|
|
func NewPromptDB(dsn string) (*PromptDB, error) {
|
||
|
|
db, err := sql.Open("mysql", dsn)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to open mysql: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// 连接池配置
|
||
|
|
db.SetMaxOpenConns(10)
|
||
|
|
db.SetMaxIdleConns(5)
|
||
|
|
db.SetConnMaxLifetime(5 * time.Minute)
|
||
|
|
|
||
|
|
// 等待连接就绪
|
||
|
|
if err := db.Ping(); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to ping mysql: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// 自动建表
|
||
|
|
if err := migrate(db); err != nil {
|
||
|
|
return nil, fmt.Errorf("failed to migrate: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
log.Println("mysql connected and migrated")
|
||
|
|
return &PromptDB{db: db}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func migrate(db *sql.DB) error {
|
||
|
|
query := `
|
||
|
|
CREATE TABLE IF NOT EXISTS prompts (
|
||
|
|
id BIGINT AUTO_INCREMENT PRIMARY KEY,
|
||
|
|
session_id VARCHAR(128) NOT NULL,
|
||
|
|
project_name VARCHAR(255) NOT NULL DEFAULT '',
|
||
|
|
prompt TEXT NOT NULL,
|
||
|
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||
|
|
INDEX idx_session_id (session_id),
|
||
|
|
INDEX idx_project_name (project_name),
|
||
|
|
INDEX idx_created_at (created_at)
|
||
|
|
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;`
|
||
|
|
|
||
|
|
_, err := db.Exec(query)
|
||
|
|
if err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
// 兼容旧表:如果已存在但缺少 project_name 列,则补充
|
||
|
|
_, err = db.Exec("ALTER TABLE prompts ADD COLUMN project_name VARCHAR(255) NOT NULL DEFAULT '' AFTER session_id")
|
||
|
|
if err != nil {
|
||
|
|
// 忽略 "Duplicate column" 错误,说明列已存在
|
||
|
|
log.Printf("alter table (expected if column exists): %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// SavePrompt 将一条 prompt 写入数据库
|
||
|
|
func (p *PromptDB) SavePrompt(sessionID, projectName, prompt string) error {
|
||
|
|
_, err := p.db.Exec(
|
||
|
|
"INSERT INTO prompts (session_id, project_name, prompt) VALUES (?, ?, ?)",
|
||
|
|
sessionID, projectName, prompt,
|
||
|
|
)
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
// Close 关闭数据库连接
|
||
|
|
func (p *PromptDB) Close() error {
|
||
|
|
if p.db != nil {
|
||
|
|
return p.db.Close()
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|