Files

87 lines
2.1 KiB
Go
Raw Permalink Normal View History

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
}