package repository import ( "context" "errors" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/yuxingu/digital-psychology/apps/api/internal/model" ) // AskRepo persists ask threads and messages. type AskRepo struct { Pool *pgxpool.Pool } // CreateThread inserts a thread bound to a profile. func (r *AskRepo) CreateThread(ctx context.Context, userID, profileID uuid.UUID, scene *string) (*model.AskThread, error) { t := &model.AskThread{} err := r.Pool.QueryRow(ctx, ` INSERT INTO ask_threads(user_id, profile_id, scene) VALUES ($1,$2,$3) RETURNING id, user_id, profile_id, scene, created_at`, userID, profileID, scene, ).Scan(&t.ID, &t.UserID, &t.ProfileID, &t.Scene, &t.CreatedAt) return t, err } // GetThreadForUser loads a thread owned by user. func (r *AskRepo) GetThreadForUser(ctx context.Context, userID, threadID uuid.UUID) (*model.AskThread, error) { t := &model.AskThread{} err := r.Pool.QueryRow(ctx, ` SELECT id, user_id, profile_id, scene, created_at FROM ask_threads WHERE id=$1 AND user_id=$2 AND deleted_at IS NULL`, threadID, userID, ).Scan(&t.ID, &t.UserID, &t.ProfileID, &t.Scene, &t.CreatedAt) return t, err } // ListMessages returns messages in a thread (oldest first). func (r *AskRepo) ListMessages(ctx context.Context, threadID uuid.UUID) ([]model.AskMessage, error) { rows, err := r.Pool.Query(ctx, ` SELECT id, thread_id, role, content, created_at FROM ask_messages WHERE thread_id=$1 AND deleted_at IS NULL ORDER BY created_at ASC`, threadID) if err != nil { return nil, err } defer rows.Close() var out []model.AskMessage for rows.Next() { var m model.AskMessage if err := rows.Scan(&m.ID, &m.ThreadID, &m.Role, &m.Content, &m.CreatedAt); err != nil { return nil, err } out = append(out, m) } return out, rows.Err() } // InsertMessage stores one message. func (r *AskRepo) InsertMessage(ctx context.Context, threadID uuid.UUID, role, content string) (*model.AskMessage, error) { m := &model.AskMessage{} err := r.Pool.QueryRow(ctx, ` INSERT INTO ask_messages(thread_id, role, content) VALUES ($1,$2,$3) RETURNING id, thread_id, role, content, created_at`, threadID, role, content, ).Scan(&m.ID, &m.ThreadID, &m.Role, &m.Content, &m.CreatedAt) return m, err } // CountUserAssistantMessages counts all assistant replies for quota (free tier). func (r *AskRepo) CountUserAssistantMessages(ctx context.Context, userID uuid.UUID) (int, error) { var n int err := r.Pool.QueryRow(ctx, ` SELECT COUNT(*) FROM ask_messages m JOIN ask_threads t ON t.id = m.thread_id WHERE t.user_id=$1 AND m.role='assistant' AND m.deleted_at IS NULL AND t.deleted_at IS NULL`, userID, ).Scan(&n) return n, err } // ConsumeMembershipQuota decrements ask_quota_left when membership is active. // Returns false if no active membership or quota is already 0. func (r *AskRepo) ConsumeMembershipQuota(ctx context.Context, userID uuid.UUID) (ok bool, left int, err error) { err = r.Pool.QueryRow(ctx, ` UPDATE memberships SET ask_quota_left = ask_quota_left - 1, updated_at=now() WHERE user_id=$1 AND status='active' AND expires_at > now() AND deleted_at IS NULL AND ask_quota_left > 0 RETURNING ask_quota_left`, userID, ).Scan(&left) if errors.Is(err, pgx.ErrNoRows) { return false, 0, nil } if err != nil { return false, 0, err } return true, left, nil }