package research import ( "context" "encoding/json" "github.com/jackc/pgx/v5/pgxpool" ) // pgxRepository 是基于 pgx 连接池的 Repository 实现。 type pgxRepository struct { pool *pgxpool.Pool } func NewPgxRepository(pool *pgxpool.Pool) Repository { return &pgxRepository{pool: pool} } func (r *pgxRepository) Insert(ctx context.Context, id, userID string, appID *string, topic string, config map[string]any) error { if config == nil { config = map[string]any{} } configJSON, err := json.Marshal(config) if err != nil { return err } _, err = r.pool.Exec(ctx, `INSERT INTO research_tasks (id, user_id, app_id, topic, config) VALUES ($1, $2, $3, $4, $5)`, id, userID, appID, topic, configJSON, ) return err } func (r *pgxRepository) Get(ctx context.Context, userID, taskID string) (*Task, error) { var t Task var sources []byte err := r.pool.QueryRow(ctx, `SELECT id, topic, status, progress, status_message, error_message, report, sources, tokens_used, created_at FROM research_tasks WHERE id = $1 AND user_id = $2`, taskID, userID, ).Scan(&t.ID, &t.Topic, &t.Status, &t.Progress, &t.StatusMessage, &t.ErrorMessage, &t.Report, &sources, &t.TokensUsed, &t.CreatedAt) if err != nil { return nil, err } t.Sources = sources return &t, nil } func (r *pgxRepository) List(ctx context.Context, userID string, limit int) ([]Task, error) { rows, err := r.pool.Query(ctx, `SELECT id, topic, status, progress, status_message, error_message, tokens_used, created_at FROM research_tasks WHERE user_id = $1 ORDER BY created_at DESC LIMIT $2`, userID, limit, ) if err != nil { return nil, err } defer rows.Close() var tasks []Task for rows.Next() { var t Task if err := rows.Scan(&t.ID, &t.Topic, &t.Status, &t.Progress, &t.StatusMessage, &t.ErrorMessage, &t.TokensUsed, &t.CreatedAt); err != nil { continue } tasks = append(tasks, t) } return tasks, nil } func (r *pgxRepository) Cancel(ctx context.Context, userID, taskID string) (bool, error) { tag, err := r.pool.Exec(ctx, `UPDATE research_tasks SET status = 'canceled', updated_at = NOW() WHERE id = $1 AND user_id = $2 AND status IN ('pending','planning','searching','reading','synthesizing')`, taskID, userID, ) if err != nil { return false, err } return tag.RowsAffected() > 0, nil }