package repo import ( "context" "database/sql" "errors" "fmt" "time" "github.com/agentmail/gateway/internal/db" "github.com/agentmail/gateway/internal/models" "github.com/google/uuid" ) // ---------- 密钥认证 ---------- // // 两类密钥共享一个全局唯一的 token 命名空间:验证时先查 agent_keys 再查 user_keys。 // 这样一个 token 永远只有一种身份,不会出现「同一串字符既能注册 Agent 又能读人类邮箱」。 var ( // ErrKeyNotFound 密钥不存在 ErrKeyNotFound = errors.New("key not found") // ErrKeyUsed 一次性密钥已被使用 ErrKeyUsed = errors.New("key already used") // ErrKeyExpired 定时密钥已过期 ErrKeyExpired = errors.New("key expired") // ErrKeyTypeInvalid 密钥类型不受支持 ErrKeyTypeInvalid = errors.New("invalid key type") // ErrKeyNeedsExpiry timed 密钥缺少有效的 expires_hours ErrKeyNeedsExpiry = errors.New("timed key requires positive expires_hours") // ErrKeyTooShort 登记的客户端密钥长度不足 ErrKeyTooShort = errors.New("key token too short") ) // expiryFor 依据密钥类型算出过期时间。 // 只有 timed 需要 expires_at;permanent 与 one_time 都是 NULL, // 各自的失效条件由 key_type 本身表达,不混用 expires_at。 func expiryFor(keyType string, hours int) (*time.Time, error) { if !models.ValidKeyType(keyType) { return nil, ErrKeyTypeInvalid } if keyType != models.KeyTimed { return nil, nil } if hours <= 0 { return nil, ErrKeyNeedsExpiry } t := time.Now().Add(time.Duration(hours) * time.Hour) return &t, nil } // checkKeyUsable 判断一条密钥记录当前是否可用。 func checkKeyUsable(keyType string, expiresAt, usedAt *time.Time) error { switch keyType { case models.KeyOneTime: if usedAt != nil { return ErrKeyUsed } case models.KeyTimed: if expiresAt == nil || time.Now().After(*expiresAt) { return ErrKeyExpired } } return nil } // ---------- Agent 密钥(管理员签发) ---------- // ErrKeyTokenTaken 登记的密钥已被占用 var ErrKeyTokenTaken = errors.New("key token already registered") // CreateAgentKey 签发一条 Agent 接入密钥。agentName 为空表示待绑定。 // // presetToken 非空时登记客户端已在本地生成的密钥(插件首装场景), // 这样密钥全文只从客户端往服务器走一次,不需要反方向传递;留空则由服务器生成。 func CreateAgentKey(ctx context.Context, agentName, keyType, label string, expiresHours int, createdBy uuid.UUID, presetToken string) (*models.AgentKey, error) { expires, err := expiryFor(keyType, expiresHours) if err != nil { return nil, err } token := presetToken if token == "" { if token, err = newToken(); err != nil { return nil, err } } else if len(token) < 32 { // 太短的客户端密钥不接受,否则等于把弱口令当凭证 return nil, ErrKeyTooShort } // 已退役的名字不可重建 —— 登记密钥会级联建 agents 行,绕过注册检查。 if agentName != "" { if retired, rErr := IsRetiredAgentName(ctx, agentName); rErr != nil { return nil, rErr } else if retired { return nil, fmt.Errorf("该名字已退役,不可重建") } } var namePtr *string if agentName != "" { namePtr = &agentName } k := &models.AgentKey{ Token: token, TokenHint: models.TokenHint(token), AgentName: namePtr, KeyType: keyType, Label: label, ExpiresAt: expires, CreatedBy: &createdBy, } err = db.DB.QueryRowContext(ctx, ` INSERT INTO agent_keys (key_token, agent_name, key_type, label, expires_at, created_by) VALUES ($1, $2, $3, $4, $5, $6) RETURNING key_id, created_at `, token, namePtr, keyType, label, expires, createdBy).Scan(&k.ID, &k.CreatedAt) if db.IsUniqueViolation(err) { return nil, ErrKeyTokenTaken } if err != nil { return nil, err } return k, nil } // ListAgentKeys 列出 Agent 密钥;agentName 非空时按 Agent 过滤。 // 返回值不含 token 全文,只有 hint。 func ListAgentKeys(ctx context.Context, agentName string) ([]models.AgentKey, error) { q := `SELECT key_id, key_token, agent_name, key_type, label, expires_at, used_at, created_by, created_at FROM agent_keys` args := []any{} if agentName != "" { q += ` WHERE agent_name = $1` args = append(args, agentName) } q += ` ORDER BY created_at DESC` rows, err := db.DB.QueryContext(ctx, q, args...) if err != nil { return nil, err } defer rows.Close() keys := []models.AgentKey{} for rows.Next() { var k models.AgentKey var token string if err := rows.Scan(&k.ID, &token, &k.AgentName, &k.KeyType, &k.Label, &k.ExpiresAt, &k.UsedAt, &k.CreatedBy, &k.CreatedAt); err != nil { return nil, err } k.TokenHint = models.TokenHint(token) // 不回传全文 keys = append(keys, k) } return keys, rows.Err() } // DeleteAgentKey 吊销一条 Agent 密钥。 func DeleteAgentKey(ctx context.Context, id uuid.UUID) error { tag, err := db.DB.ExecContext(ctx, `DELETE FROM agent_keys WHERE key_id = $1`, id) if err != nil { return err } if n, _ := tag.RowsAffected(); n == 0 { return ErrKeyNotFound } return nil } // BindAgentKey 把一条密钥绑定到指定 Agent 名。 func BindAgentKey(ctx context.Context, id uuid.UUID, agentName string) error { tag, err := db.DB.ExecContext(ctx, `UPDATE agent_keys SET agent_name = $2 WHERE key_id = $1`, id, agentName) if err != nil { return err } if n, _ := tag.RowsAffected(); n == 0 { return ErrKeyNotFound } return nil } // VerifyAgentKey 校验 Agent 密钥并返回它绑定的 Agent 名(未绑定时返回空串)。 // // 一次性密钥在校验通过时立刻写 used_at —— 用 WHERE used_at IS NULL 保证并发下 // 只有一个请求能把它标记掉,避免两个 Agent 拿同一把一次性密钥同时注册成功。 func VerifyAgentKey(ctx context.Context, token string) (string, error) { var ( id uuid.UUID agentName *string keyType string expiresAt *time.Time usedAt *time.Time ) err := db.DB.QueryRowContext(ctx, ` SELECT key_id, agent_name, key_type, expires_at, used_at FROM agent_keys WHERE key_token = $1 `, token).Scan(&id, &agentName, &keyType, &expiresAt, &usedAt) if errors.Is(err, sql.ErrNoRows) { return "", ErrKeyNotFound } if err != nil { return "", err } if err := checkKeyUsable(keyType, expiresAt, usedAt); err != nil { return "", err } if keyType == models.KeyOneTime { tag, err := db.DB.ExecContext(ctx, `UPDATE agent_keys SET used_at = NOW() WHERE key_id = $1 AND used_at IS NULL`, id) if err != nil { return "", err } if n, _ := tag.RowsAffected(); n == 0 { return "", ErrKeyUsed // 并发下被别人抢先用掉了 } } if agentName == nil { return "", nil } return *agentName, nil } // ClaimAgentKey 在待绑定密钥首次注册时把它落定到该 Agent 名。 // 已绑定的密钥不受影响(WHERE agent_name IS NULL),因此不能借一把已绑定的密钥改注册别的 Agent。 func ClaimAgentKey(ctx context.Context, token, agentName string) error { _, err := db.DB.ExecContext(ctx, `UPDATE agent_keys SET agent_name = $2 WHERE key_token = $1 AND agent_name IS NULL`, token, agentName) return err } // ---------- 用户密钥(用户自助签发) ---------- // CreateUserKey 为用户签发一条客户端连接密钥。 func CreateUserKey(ctx context.Context, userID uuid.UUID, label, keyType string, expiresHours int) (*models.UserKey, error) { expires, err := expiryFor(keyType, expiresHours) if err != nil { return nil, err } token, err := newToken() if err != nil { return nil, err } k := &models.UserKey{ Token: token, TokenHint: models.TokenHint(token), UserID: userID, Label: label, KeyType: keyType, ExpiresAt: expires, } err = db.DB.QueryRowContext(ctx, ` INSERT INTO user_keys (key_token, user_id, label, key_type, expires_at) VALUES ($1, $2, $3, $4, $5) RETURNING key_id, created_at `, token, userID, label, keyType, expires).Scan(&k.ID, &k.CreatedAt) if err != nil { return nil, err } return k, nil } // ListUserKeys 列出某用户的连接密钥(不含 token 全文)。 func ListUserKeys(ctx context.Context, userID uuid.UUID) ([]models.UserKey, error) { rows, err := db.DB.QueryContext(ctx, ` SELECT key_id, key_token, user_id, label, key_type, expires_at, used_at, created_at FROM user_keys WHERE user_id = $1 ORDER BY created_at DESC `, userID) if err != nil { return nil, err } defer rows.Close() keys := []models.UserKey{} for rows.Next() { var k models.UserKey var token string if err := rows.Scan(&k.ID, &token, &k.UserID, &k.Label, &k.KeyType, &k.ExpiresAt, &k.UsedAt, &k.CreatedAt); err != nil { return nil, err } k.TokenHint = models.TokenHint(token) keys = append(keys, k) } return keys, rows.Err() } // DeleteUserKey 删除自己的一条密钥。带 user_id 条件,避免删掉别人的。 func DeleteUserKey(ctx context.Context, userID, keyID uuid.UUID) error { tag, err := db.DB.ExecContext(ctx, `DELETE FROM user_keys WHERE key_id = $1 AND user_id = $2`, keyID, userID) if err != nil { return err } if n, _ := tag.RowsAffected(); n == 0 { return ErrKeyNotFound } return nil } // VerifyUserKey 校验用户密钥并返回对应用户。 // 用户必须仍处于 active 状态——禁用账号后其密钥应当立即失效。 func VerifyUserKey(ctx context.Context, token string) (*models.User, error) { var ( id uuid.UUID userID uuid.UUID keyType string expiresAt *time.Time usedAt *time.Time ) err := db.DB.QueryRowContext(ctx, ` SELECT key_id, user_id, key_type, expires_at, used_at FROM user_keys WHERE key_token = $1 `, token).Scan(&id, &userID, &keyType, &expiresAt, &usedAt) if errors.Is(err, sql.ErrNoRows) { return nil, ErrKeyNotFound } if err != nil { return nil, err } if err := checkKeyUsable(keyType, expiresAt, usedAt); err != nil { return nil, err } if keyType == models.KeyOneTime { tag, err := db.DB.ExecContext(ctx, `UPDATE user_keys SET used_at = NOW() WHERE key_id = $1 AND used_at IS NULL`, id) if err != nil { return nil, err } if n, _ := tag.RowsAffected(); n == 0 { return nil, ErrKeyUsed } } u, err := GetUserByID(ctx, userID) if err != nil { return nil, err } if u.Status != "active" { return nil, ErrKeyNotFound // 账号已禁用,密钥一并失效 } return u, nil }