Files
goclaw/internal/store/pg/pairing.go
T
yatulandClaude Opus 5 ea4890c257 fix(pairing): address review — docs, error mapping, approve feedback, unreadable expiry
- docs: device.pair.update and the `permanent` option on approve in
  docs/04-gateway-protocol.md, docs/19-websocket-rpc.md and
  websocket-protocol.md; the paired-device TTL row in docs/09-security.md
  now mentions the admin opt-out.
- store.ErrPairedDeviceNotFound: SetPairingPermanent wraps it in both stores.
  device.pair.update maps it to NOT_FOUND and any other store error to
  INTERNAL, so a DB failure no longer reads as "not found".
- web UI: approve and make-permanent/set-expiry now toast the server error
  and reload the list in `finally`. A partially applied approve (paired, but
  the permanent write failed) shows up in the table instead of leaving the
  dialog dead-ended.
- SQLite ListPaired: a stored expiry that fails to parse stays 0 (expires,
  date unknown) rather than being mistaken for permanent; the UI renders it
  as "--" instead of a 1970 date.

Tests: gateway handler error mapping (NOT_FOUND / INTERNAL / OK), sentinel
checks in the PG and SQLite store tests, SQLite unreadable-expiry case.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-18 19:26:43 +04:00

347 lines
11 KiB
Go

package pg
import (
"context"
"crypto/rand"
"database/sql"
"encoding/json"
"fmt"
"time"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
const (
codeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
codeLength = 8
codeTTL = 60 * time.Minute
pairedDeviceTTL = 30 * 24 * time.Hour // 30 days
maxPendingPerAccount = 3
)
// PGPairingStore implements store.PairingStore backed by Postgres.
type PGPairingStore struct {
db *sql.DB
onRequest func(code, senderID, channel, chatID string)
}
func NewPGPairingStore(db *sql.DB) *PGPairingStore {
return &PGPairingStore{db: db}
}
// SetOnRequest sets a callback fired after a new pairing request is created.
func (s *PGPairingStore) SetOnRequest(cb func(code, senderID, channel, chatID string)) {
s.onRequest = cb
}
func (s *PGPairingStore) RequestPairing(ctx context.Context, senderID, channel, chatID, accountID string, metadata map[string]string) (string, error) {
tid := tenantIDForInsert(ctx)
// Prune expired
s.db.ExecContext(ctx, "DELETE FROM pairing_requests WHERE expires_at < $1", time.Now())
// Check max pending (per tenant)
var count int64
s.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM pairing_requests WHERE account_id = $1 AND tenant_id = $2", accountID, tid).Scan(&count)
if count >= maxPendingPerAccount {
return "", fmt.Errorf("max pending pairing requests (%d) exceeded", maxPendingPerAccount)
}
// Check existing (per tenant)
var existingCode string
err := s.db.QueryRowContext(ctx, "SELECT code FROM pairing_requests WHERE sender_id = $1 AND channel = $2 AND tenant_id = $3", senderID, channel, tid).Scan(&existingCode)
if err == nil {
return existingCode, nil
}
metaJSON := []byte("{}")
if len(metadata) > 0 {
metaJSON, _ = json.Marshal(metadata)
}
code := generatePairingCode()
now := time.Now()
_, err = s.db.ExecContext(ctx,
`INSERT INTO pairing_requests (id, code, sender_id, channel, chat_id, account_id, expires_at, created_at, metadata, tenant_id)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)`,
uuid.Must(uuid.NewV7()), code, senderID, channel, chatID, accountID, now.Add(codeTTL), now, metaJSON, tid,
)
if err != nil {
return "", fmt.Errorf("create pairing request: %w", err)
}
if s.onRequest != nil {
go s.onRequest(code, senderID, channel, chatID)
}
return code, nil
}
// ApprovePairing looks up by code (globally unique random token) and creates paired device.
// The approver's tenant context determines paired_devices.tenant_id.
func (s *PGPairingStore) ApprovePairing(ctx context.Context, code, approvedBy string) (*store.PairedDeviceData, error) {
// Prune expired
s.db.ExecContext(ctx, "DELETE FROM pairing_requests WHERE expires_at < $1", time.Now())
var reqID uuid.UUID
var senderID, channel, chatID string
var metaJSON []byte
var reqTenantID uuid.UUID
// Code lookup is cross-tenant (random token, approver is WS/HTTP user).
// Also check expires_at to close race between prune DELETE and this SELECT.
err := s.db.QueryRowContext(ctx,
"SELECT id, sender_id, channel, chat_id, COALESCE(metadata, '{}'), tenant_id FROM pairing_requests WHERE code = $1 AND expires_at > NOW()", code,
).Scan(&reqID, &senderID, &channel, &chatID, &metaJSON, &reqTenantID)
if err != nil {
return nil, fmt.Errorf("pairing code %s not found or expired", code)
}
// Remove from pending
s.db.ExecContext(ctx, "DELETE FROM pairing_requests WHERE id = $1", reqID)
// Add to paired — use the request's tenant (the channel that initiated pairing)
now := time.Now()
expiresAt := now.Add(pairedDeviceTTL)
_, err = s.db.ExecContext(ctx,
`INSERT INTO paired_devices (id, sender_id, channel, chat_id, paired_by, paired_at, metadata, expires_at, tenant_id)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)`,
uuid.Must(uuid.NewV7()), senderID, channel, chatID, approvedBy, now, metaJSON, expiresAt, reqTenantID,
)
if err != nil {
return nil, fmt.Errorf("create paired device: %w", err)
}
var meta map[string]string
if len(metaJSON) > 0 {
json.Unmarshal(metaJSON, &meta)
}
expiresAtMs := expiresAt.UnixMilli()
return &store.PairedDeviceData{
SenderID: senderID,
Channel: channel,
ChatID: chatID,
PairedAt: now.UnixMilli(),
PairedBy: approvedBy,
ExpiresAt: &expiresAtMs,
Metadata: meta,
}, nil
}
func (s *PGPairingStore) DenyPairing(ctx context.Context, code string) error {
// Code lookup is cross-tenant (random token)
result, err := s.db.ExecContext(ctx, "DELETE FROM pairing_requests WHERE code = $1", code)
if err != nil {
return err
}
n, _ := result.RowsAffected()
if n == 0 {
return fmt.Errorf("pairing code %s not found or expired", code)
}
return nil
}
func (s *PGPairingStore) RevokePairing(ctx context.Context, senderID, channel string) error {
tid := tenantIDForInsert(ctx)
result, err := s.db.ExecContext(ctx, "DELETE FROM paired_devices WHERE sender_id = $1 AND channel = $2 AND tenant_id = $3", senderID, channel, tid)
if err != nil {
return err
}
n, _ := result.RowsAffected()
if n == 0 {
return fmt.Errorf("paired device not found: %s/%s", channel, senderID)
}
return nil
}
// SetPairingPermanent clears (permanent=true) or restarts (permanent=false)
// the expiry of a live pairing. An already expired pairing is not revived.
func (s *PGPairingStore) SetPairingPermanent(ctx context.Context, senderID, channel string, permanent bool) error {
tid := tenantIDForInsert(ctx)
var expiresAt *time.Time
if !permanent {
t := time.Now().Add(pairedDeviceTTL)
expiresAt = &t
}
result, err := s.db.ExecContext(ctx,
`UPDATE paired_devices SET expires_at = $1
WHERE sender_id = $2 AND channel = $3 AND tenant_id = $4 AND (expires_at IS NULL OR expires_at > NOW())`,
expiresAt, senderID, channel, tid,
)
if err != nil {
return err
}
n, _ := result.RowsAffected()
if n == 0 {
return fmt.Errorf("%w: %s/%s", store.ErrPairedDeviceNotFound, channel, senderID)
}
return nil
}
func (s *PGPairingStore) IsPaired(ctx context.Context, senderID, channel string) (bool, error) {
tid := tenantIDForInsert(ctx)
var count int64
err := s.db.QueryRowContext(ctx,
"SELECT COUNT(*) FROM paired_devices WHERE sender_id = $1 AND channel = $2 AND tenant_id = $3 AND (expires_at IS NULL OR expires_at > NOW())",
senderID, channel, tid,
).Scan(&count)
if err != nil {
return false, fmt.Errorf("pairing check query: %w", err)
}
return count > 0, nil
}
// pairingRequestRow is an sqlx scan struct for pairing_requests.
// Domain struct uses int64 (Unix ms) for timestamps, DB stores time.Time.
type pairingRequestRow struct {
Code string `json:"code" db:"code"`
SenderID string `json:"sender_id" db:"sender_id"`
Channel string `json:"channel" db:"channel"`
ChatID string `json:"chat_id" db:"chat_id"`
AccountID string `json:"account_id" db:"account_id"`
CreatedAt time.Time `json:"created_at" db:"created_at"`
ExpiresAt time.Time `json:"expires_at" db:"expires_at"`
Metadata []byte `json:"metadata" db:"metadata"`
}
// pairedDeviceRow is an sqlx scan struct for paired_devices.
type pairedDeviceRow struct {
SenderID string `json:"sender_id" db:"sender_id"`
Channel string `json:"channel" db:"channel"`
ChatID string `json:"chat_id" db:"chat_id"`
PairedBy string `json:"paired_by" db:"paired_by"`
PairedAt time.Time `json:"paired_at" db:"paired_at"`
ExpiresAt *time.Time `json:"expires_at" db:"expires_at"`
Metadata []byte `json:"metadata" db:"metadata"`
}
func (s *PGPairingStore) ListPending(ctx context.Context) []store.PairingRequestData {
tid := tenantIDForInsert(ctx)
// Prune expired
s.db.ExecContext(ctx, "DELETE FROM pairing_requests WHERE expires_at < $1", time.Now())
var rows []pairingRequestRow
err := pkgSqlxDB.SelectContext(ctx, &rows,
`SELECT code, sender_id, channel, chat_id, account_id, created_at, expires_at, COALESCE(metadata, '{}') AS metadata
FROM pairing_requests WHERE tenant_id = $1 ORDER BY created_at DESC`, tid)
if err != nil {
return []store.PairingRequestData{}
}
result := make([]store.PairingRequestData, len(rows))
for i, r := range rows {
result[i] = store.PairingRequestData{
Code: r.Code, SenderID: r.SenderID, Channel: r.Channel,
ChatID: r.ChatID, AccountID: r.AccountID,
CreatedAt: r.CreatedAt.UnixMilli(), ExpiresAt: r.ExpiresAt.UnixMilli(),
}
if len(r.Metadata) > 0 {
json.Unmarshal(r.Metadata, &result[i].Metadata)
}
}
return result
}
func (s *PGPairingStore) ListPaired(ctx context.Context) []store.PairedDeviceData {
tid := tenantIDForInsert(ctx)
// Prune expired paired devices
s.db.ExecContext(ctx, "DELETE FROM paired_devices WHERE expires_at IS NOT NULL AND expires_at < NOW()")
var rows []pairedDeviceRow
err := pkgSqlxDB.SelectContext(ctx, &rows,
`SELECT sender_id, channel, chat_id, paired_by, paired_at, expires_at, COALESCE(metadata, '{}') AS metadata
FROM paired_devices WHERE tenant_id = $1 ORDER BY paired_at DESC`, tid)
if err != nil {
return []store.PairedDeviceData{}
}
result := make([]store.PairedDeviceData, len(rows))
for i, r := range rows {
result[i] = store.PairedDeviceData{
SenderID: r.SenderID, Channel: r.Channel, ChatID: r.ChatID,
PairedBy: r.PairedBy, PairedAt: r.PairedAt.UnixMilli(),
}
if r.ExpiresAt != nil {
ms := r.ExpiresAt.UnixMilli()
result[i].ExpiresAt = &ms
}
if len(r.Metadata) > 0 {
json.Unmarshal(r.Metadata, &result[i].Metadata)
}
}
return result
}
func (s *PGPairingStore) MigrateGroupChatID(ctx context.Context, channel, oldChatID, newChatID string) error {
tid := tenantIDForInsert(ctx)
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin migrate tx: %w", err)
}
defer tx.Rollback()
// 1. paired_devices: update sender_id and chat_id
if _, err := tx.ExecContext(ctx,
`UPDATE paired_devices
SET sender_id = REPLACE(sender_id, $1, $2),
chat_id = REPLACE(chat_id, $1, $2)
WHERE sender_id LIKE '%' || $1 || '%'
AND channel = $3
AND tenant_id = $4`,
oldChatID, newChatID, channel, tid,
); err != nil {
return fmt.Errorf("migrate paired_devices: %w", err)
}
// 2. sessions: update session_key and user_id
if _, err := tx.ExecContext(ctx,
`UPDATE sessions
SET session_key = REPLACE(session_key, ':' || $1, ':' || $2),
user_id = REPLACE(user_id, ':' || $1, ':' || $2)
WHERE session_key LIKE '%:telegram:%:' || $1 || '%'
AND tenant_id = $3`,
oldChatID, newChatID, tid,
); err != nil {
return fmt.Errorf("migrate sessions: %w", err)
}
// 3. channel_contacts: update sender_id
if _, err := tx.ExecContext(ctx,
`UPDATE channel_contacts
SET sender_id = REPLACE(sender_id, $1, $2)
WHERE sender_id LIKE '%' || $1 || '%'
AND channel_type = 'telegram'
AND tenant_id = $3`,
oldChatID, newChatID, tid,
); err != nil {
return fmt.Errorf("migrate channel_contacts: %w", err)
}
// 4. channel_pending_messages: update history_key
if _, err := tx.ExecContext(ctx,
`UPDATE channel_pending_messages
SET history_key = REPLACE(history_key, $1, $2)
WHERE history_key LIKE '%' || $1 || '%'
AND channel_name = $3
AND tenant_id = $4`,
oldChatID, newChatID, channel, tid,
); err != nil {
return fmt.Errorf("migrate channel_pending_messages: %w", err)
}
return tx.Commit()
}
func generatePairingCode() string {
b := make([]byte, codeLength)
rand.Read(b)
code := make([]byte, codeLength)
for i := range code {
code[i] = codeAlphabet[int(b[i])%len(codeAlphabet)]
}
return string(code)
}