package db import ( "context" "fmt" "time" "github.com/emil/deepres/internal/models" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) type DB struct { pool *pgxpool.Pool } func New(databaseURL string) (*DB, error) { pool, err := pgxpool.New(context.Background(), databaseURL) if err != nil { return nil, fmt.Errorf("connect to db: %w", err) } return &DB{pool: pool}, nil } func (db *DB) Close() { db.pool.Close() } // ---- Users ---- func (db *DB) GetOrCreateUser(ctx context.Context, tgID int64, username, name string) (*models.User, error) { u := &models.User{} err := db.pool.QueryRow(ctx, ` INSERT INTO users (tg_id, username, name) VALUES ($1, $2, $3) ON CONFLICT (tg_id) DO UPDATE SET username=$2, name=$3 RETURNING tg_id, username, name, stars_balance, created_at `, tgID, username, name).Scan(&u.TelegramID, &u.Username, &u.Name, &u.StarsBalance, &u.CreatedAt) if err != nil { return nil, fmt.Errorf("get or create user: %w", err) } return u, nil } func (db *DB) GetUser(ctx context.Context, tgID int64) (*models.User, error) { u := &models.User{} err := db.pool.QueryRow(ctx, "SELECT tg_id, username, name, stars_balance, created_at FROM users WHERE tg_id=$1", tgID, ).Scan(&u.TelegramID, &u.Username, &u.Name, &u.StarsBalance, &u.CreatedAt) if err == pgx.ErrNoRows { return nil, nil } if err != nil { return nil, fmt.Errorf("get user: %w", err) } return u, nil } // ---- Research ---- func (db *DB) CreateResearch(ctx context.Context, r *models.Research) error { r.ID = uuid.New() r.Status = "pending" r.CreatedAt = time.Now() _, err := db.pool.Exec(ctx, ` INSERT INTO researches (id, tg_id, query, mode, status, stars_cost, created_at) VALUES ($1,$2,$3,$4,$5,$6,$7) `, r.ID, r.TelegramID, r.Query, r.Mode, r.Status, r.StarsCost, r.CreatedAt) return err } func (db *DB) UpdateResearch(ctx context.Context, id uuid.UUID, status, summary string, reportPath string, sources []byte, errMsg *string) error { now := time.Now() _, err := db.pool.Exec(ctx, ` UPDATE researches SET status=$2, summary=$3, report_path=$4, sources=$5, error_message=$6, completed_at=$7 WHERE id=$1 `, id, status, summary, reportPath, sources, errMsg, now) return err } func (db *DB) GetResearch(ctx context.Context, id uuid.UUID) (*models.Research, error) { r := &models.Research{} err := db.pool.QueryRow(ctx, "SELECT id,tg_id,query,mode,status,stars_cost,summary,report_path,sources,error_message,created_at,completed_at FROM researches WHERE id=$1", id). Scan(&r.ID, &r.TelegramID, &r.Query, &r.Mode, &r.Status, &r.StarsCost, &r.Summary, &r.ReportPath, &r.SourcesJSON, &r.ErrorMessage, &r.CreatedAt, &r.CompletedAt) if err != nil { return nil, err } return r, nil } // ---- Transactions ---- // AddTransaction добавляет транзакцию и обновляет баланс атомарно func (db *DB) AddTransaction(ctx context.Context, t *models.Transaction) error { tx, err := db.pool.Begin(ctx) if err != nil { return fmt.Errorf("begin tx: %w", err) } defer tx.Rollback(ctx) t.ID = uuid.New() t.CreatedAt = time.Now() _, err = tx.Exec(ctx, ` INSERT INTO transactions (id, tg_id, amount, type, description, telegram_charge_id, created_at) VALUES ($1,$2,$3,$4,$5,$6,$7) `, t.ID, t.TelegramID, t.Amount, t.Type, t.Description, t.TelegramChargeID, t.CreatedAt) if err != nil { return fmt.Errorf("insert transaction: %w", err) } // Обновляем баланс _, err = tx.Exec(ctx, "UPDATE users SET stars_balance = stars_balance + $1 WHERE tg_id=$2", t.Amount, t.TelegramID) if err != nil { return fmt.Errorf("update balance: %w", err) } return tx.Commit(ctx) } // DeductResearchCost списывает звёзды за ресерч атомарно func (db *DB) DeductResearchCost(ctx context.Context, userID int64, researchID string, cost int, description string) error { tx, err := db.pool.Begin(ctx) if err != nil { return fmt.Errorf("begin tx: %w", err) } defer tx.Rollback(ctx) // Проверяем баланс var balance int err = tx.QueryRow(ctx, "SELECT stars_balance FROM users WHERE tg_id=$1 FOR UPDATE", userID).Scan(&balance) if err != nil { return fmt.Errorf("get balance: %w", err) } if balance < cost { return fmt.Errorf("insufficient balance: %d < %d", balance, cost) } // Создаём транзакцию _, err = tx.Exec(ctx, ` INSERT INTO transactions (id, tg_id, amount, type, description, created_at) VALUES (gen_random_uuid(), $1, $2, 'research', $3, now()) `, userID, -cost, description) if err != nil { return fmt.Errorf("insert transaction: %w", err) } // Списываем _, err = tx.Exec(ctx, "UPDATE users SET stars_balance = stars_balance - $1 WHERE tg_id=$2", cost, userID) if err != nil { return fmt.Errorf("deduct balance: %w", err) } return tx.Commit(ctx) } func (db *DB) UpdateStarsBalance(ctx context.Context, tgID int64, delta int) error { _, err := db.pool.Exec(ctx, "UPDATE users SET stars_balance = stars_balance + $1 WHERE tg_id=$2", delta, tgID) return err } func (db *DB) GetTodayFastCount(ctx context.Context, tgID int64) (int, error) { var count int err := db.pool.QueryRow(ctx, ` SELECT COUNT(*) FROM researches WHERE tg_id=$1 AND mode='fast' AND created_at > CURRENT_DATE `, tgID).Scan(&count) return count, err }