352 lines
10 KiB
Go
352 lines
10 KiB
Go
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
|
||
}
|