mirror of
https://github.com/OpenByteInc/QuantDinger.git
synced 2026-09-28 23:32:55 +08:00
799 lines
30 KiB
Python
799 lines
30 KiB
Python
"""
|
|
Agent Gateway authentication, scopes, audit, idempotency, and rate limiting.
|
|
|
|
This module is intentionally **separate** from `app.utils.auth` (human JWT).
|
|
Agent tokens authenticate machine clients (external AI agents, MCP servers,
|
|
custom automations) against `/api/agent/v1/...` and are subject to
|
|
capability-class scoping, per-token rate limits, and an append-only audit log.
|
|
|
|
Contract reference: docs/agent/agent-openapi.json
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import os
|
|
import secrets
|
|
import threading
|
|
import time
|
|
from contextlib import contextmanager
|
|
from datetime import datetime
|
|
from functools import wraps
|
|
from typing import Any, Callable, Iterable, Optional
|
|
|
|
from flask import current_app, g, jsonify, make_response, request
|
|
|
|
from app.config.redis_urls import cache_redis_url
|
|
from app.utils.db import get_db_connection
|
|
from app.utils.logger import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
TOKEN_PREFIX = "qd_agent_"
|
|
|
|
# Capability classes are documented in docs/agent/README.md.
|
|
SCOPE_R = "R" # Read
|
|
SCOPE_W = "W" # Workspace write
|
|
SCOPE_B = "B" # Backtest / simulation
|
|
SCOPE_N = "N" # Notifications & misc side-effects
|
|
SCOPE_C = "C" # Credentials (admin only)
|
|
SCOPE_T = "T" # Trading / capital
|
|
ALL_SCOPES = (SCOPE_R, SCOPE_W, SCOPE_B, SCOPE_N, SCOPE_C, SCOPE_T)
|
|
|
|
|
|
_schema_ready = False
|
|
_schema_lock = threading.Lock()
|
|
|
|
|
|
def _ensure_schema() -> None:
|
|
"""Idempotent runtime guard.
|
|
|
|
The canonical schema lives in `migrations/init.sql` and is applied by the
|
|
Postgres container's first-boot script. For installations that upgraded
|
|
in-place we still want the agent tables to materialize on first use so the
|
|
gateway never fails with "relation does not exist".
|
|
"""
|
|
global _schema_ready
|
|
if _schema_ready:
|
|
return
|
|
with _schema_lock:
|
|
if _schema_ready:
|
|
return
|
|
ddl = """
|
|
CREATE TABLE IF NOT EXISTS qd_agent_tokens (
|
|
id SERIAL PRIMARY KEY,
|
|
user_id INTEGER NOT NULL REFERENCES qd_users(id) ON DELETE CASCADE,
|
|
name VARCHAR(80) NOT NULL,
|
|
token_prefix VARCHAR(24) NOT NULL,
|
|
token_hash VARCHAR(128) NOT NULL,
|
|
scopes TEXT NOT NULL DEFAULT 'R',
|
|
markets TEXT NOT NULL DEFAULT '*',
|
|
instruments TEXT NOT NULL DEFAULT '*',
|
|
paper_only BOOLEAN NOT NULL DEFAULT TRUE,
|
|
rate_limit_per_min INTEGER NOT NULL DEFAULT 60,
|
|
max_order_notional DECIMAL(24,8) NOT NULL DEFAULT 1000,
|
|
max_daily_notional DECIMAL(24,8) NOT NULL DEFAULT 5000,
|
|
status VARCHAR(20) NOT NULL DEFAULT 'active',
|
|
expires_at TIMESTAMP,
|
|
last_used_at TIMESTAMP,
|
|
created_at TIMESTAMP DEFAULT NOW()
|
|
);
|
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_agent_tokens_hash ON qd_agent_tokens(token_hash);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_tokens_user ON qd_agent_tokens(user_id);
|
|
|
|
CREATE TABLE IF NOT EXISTS qd_agent_jobs (
|
|
id BIGSERIAL PRIMARY KEY,
|
|
job_id VARCHAR(40) NOT NULL UNIQUE,
|
|
user_id INTEGER NOT NULL REFERENCES qd_users(id) ON DELETE CASCADE,
|
|
agent_token_id INTEGER REFERENCES qd_agent_tokens(id) ON DELETE SET NULL,
|
|
kind VARCHAR(40) NOT NULL,
|
|
status VARCHAR(20) NOT NULL DEFAULT 'queued',
|
|
request JSONB NOT NULL DEFAULT '{}'::jsonb,
|
|
result JSONB,
|
|
error TEXT,
|
|
progress JSONB,
|
|
idempotency_key VARCHAR(120),
|
|
created_at TIMESTAMP DEFAULT NOW(),
|
|
started_at TIMESTAMP,
|
|
finished_at TIMESTAMP
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_jobs_user ON qd_agent_jobs(user_id);
|
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_agent_jobs_idem
|
|
ON qd_agent_jobs(agent_token_id, kind, idempotency_key)
|
|
WHERE idempotency_key IS NOT NULL;
|
|
|
|
CREATE TABLE IF NOT EXISTS qd_agent_audit (
|
|
id BIGSERIAL PRIMARY KEY,
|
|
user_id INTEGER NOT NULL,
|
|
agent_token_id INTEGER,
|
|
agent_name VARCHAR(80),
|
|
route VARCHAR(160) NOT NULL,
|
|
method VARCHAR(8) NOT NULL,
|
|
scope_class VARCHAR(4) NOT NULL,
|
|
status_code INTEGER NOT NULL,
|
|
idempotency_key VARCHAR(120),
|
|
request_summary JSONB,
|
|
response_summary JSONB,
|
|
duration_ms INTEGER,
|
|
created_at TIMESTAMP DEFAULT NOW()
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_audit_user ON qd_agent_audit(user_id, created_at DESC);
|
|
|
|
CREATE TABLE IF NOT EXISTS qd_agent_paper_orders (
|
|
id BIGSERIAL PRIMARY KEY,
|
|
order_uid VARCHAR(40) NOT NULL UNIQUE,
|
|
user_id INTEGER NOT NULL REFERENCES qd_users(id) ON DELETE CASCADE,
|
|
agent_token_id INTEGER REFERENCES qd_agent_tokens(id) ON DELETE SET NULL,
|
|
market VARCHAR(40) NOT NULL,
|
|
symbol VARCHAR(60) NOT NULL,
|
|
side VARCHAR(8) NOT NULL,
|
|
order_type VARCHAR(16) NOT NULL DEFAULT 'market',
|
|
qty DECIMAL(28,10) NOT NULL,
|
|
limit_price DECIMAL(28,10),
|
|
fill_price DECIMAL(28,10),
|
|
fill_value DECIMAL(28,10),
|
|
status VARCHAR(16) NOT NULL DEFAULT 'filled',
|
|
note TEXT,
|
|
created_at TIMESTAMP DEFAULT NOW()
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_paper_orders_user
|
|
ON qd_agent_paper_orders(user_id, created_at DESC);
|
|
|
|
ALTER TABLE qd_agent_tokens
|
|
ADD COLUMN IF NOT EXISTS max_order_notional DECIMAL(24,8) NOT NULL DEFAULT 1000;
|
|
ALTER TABLE qd_agent_tokens
|
|
ADD COLUMN IF NOT EXISTS max_daily_notional DECIMAL(24,8) NOT NULL DEFAULT 5000;
|
|
|
|
CREATE TABLE IF NOT EXISTS qd_agent_idempotency (
|
|
id BIGSERIAL PRIMARY KEY,
|
|
agent_token_id INTEGER NOT NULL REFERENCES qd_agent_tokens(id) ON DELETE CASCADE,
|
|
method VARCHAR(8) NOT NULL,
|
|
route VARCHAR(200) NOT NULL,
|
|
idempotency_key VARCHAR(120) NOT NULL,
|
|
request_hash VARCHAR(64) NOT NULL,
|
|
status VARCHAR(20) NOT NULL DEFAULT 'started',
|
|
response_body JSONB,
|
|
response_status INTEGER,
|
|
created_at TIMESTAMP DEFAULT NOW(),
|
|
updated_at TIMESTAMP DEFAULT NOW(),
|
|
UNIQUE(agent_token_id, method, route, idempotency_key)
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_idempotency_created
|
|
ON qd_agent_idempotency(created_at);
|
|
|
|
CREATE TABLE IF NOT EXISTS qd_agent_notional_reservations (
|
|
id BIGSERIAL PRIMARY KEY,
|
|
user_id INTEGER NOT NULL REFERENCES qd_users(id) ON DELETE CASCADE,
|
|
agent_token_id INTEGER NOT NULL REFERENCES qd_agent_tokens(id) ON DELETE CASCADE,
|
|
idempotency_key VARCHAR(120) NOT NULL,
|
|
notional DECIMAL(24,8) NOT NULL,
|
|
status VARCHAR(20) NOT NULL DEFAULT 'reserved',
|
|
created_at TIMESTAMP DEFAULT NOW(),
|
|
updated_at TIMESTAMP DEFAULT NOW(),
|
|
UNIQUE(agent_token_id, idempotency_key)
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_notional_daily
|
|
ON qd_agent_notional_reservations(agent_token_id, created_at);
|
|
"""
|
|
try:
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
for stmt in [s.strip() for s in ddl.split(";") if s.strip()]:
|
|
cur.execute(stmt)
|
|
db.commit()
|
|
cur.close()
|
|
_schema_ready = True
|
|
except Exception as exc:
|
|
logger.warning(f"agent_auth: schema ensure failed (will retry): {exc}")
|
|
|
|
|
|
def ensure_agent_gateway_schema() -> None:
|
|
"""Ensure agent gateway tables exist (idempotent).
|
|
|
|
Admin JWT routes (e.g. token issuance) bypass ``agent_required``, which
|
|
normally triggers ``_ensure_schema()`` on first agent call. Without this,
|
|
a fresh or partially migrated DB can hit ``INSERT`` before tables exist and
|
|
return an unhandled 500.
|
|
"""
|
|
_ensure_schema()
|
|
|
|
|
|
# ─────────────────────────── token primitives ───────────────────────────
|
|
|
|
def _hash_token(token: str) -> str:
|
|
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
|
|
|
|
|
def generate_token() -> tuple[str, str, str]:
|
|
"""Generate a new agent token.
|
|
|
|
Returns:
|
|
(full_token, token_prefix, token_hash). Only the hash is stored;
|
|
the full token is shown to the operator exactly once.
|
|
"""
|
|
body = secrets.token_urlsafe(32).rstrip("=")
|
|
full = f"{TOKEN_PREFIX}{body}"
|
|
prefix = full[: len(TOKEN_PREFIX) + 8] # qd_agent_XXXXXXXX
|
|
return full, prefix, _hash_token(full)
|
|
|
|
|
|
def parse_scopes(raw: str | Iterable[str] | None) -> set[str]:
|
|
if raw is None:
|
|
return {SCOPE_R}
|
|
if isinstance(raw, str):
|
|
items = [p.strip().upper() for p in raw.split(",") if p.strip()]
|
|
else:
|
|
items = [str(p).strip().upper() for p in raw if str(p).strip()]
|
|
return {p for p in items if p in ALL_SCOPES}
|
|
|
|
|
|
def parse_csv_list(raw: str | None, default: str = "*") -> list[str]:
|
|
if not raw:
|
|
return [default]
|
|
items = [p.strip() for p in str(raw).split(",") if p.strip()]
|
|
return items or [default]
|
|
|
|
|
|
def list_matches(item: str, allowlist: list[str]) -> bool:
|
|
if not allowlist or "*" in allowlist:
|
|
return True
|
|
needle = (item or "").strip().upper()
|
|
return any(needle == a.strip().upper() for a in allowlist)
|
|
|
|
|
|
# ─────────────────────────── distributed rate limit ───────────────────────────
|
|
|
|
_rate_state: dict[int, list[float]] = {}
|
|
_rate_lock = threading.Lock()
|
|
_redis_rate_client = None
|
|
_redis_rate_lock = threading.Lock()
|
|
_redis_rate_warned = False
|
|
|
|
|
|
def _memory_rate_limit(key: str, limit_per_min: int) -> dict[str, int | bool]:
|
|
now = time.time()
|
|
window_start = now - 60.0
|
|
state_key = hash(key)
|
|
with _rate_lock:
|
|
bucket = [t for t in _rate_state.get(state_key, []) if t >= window_start]
|
|
if len(bucket) >= max(1, int(limit_per_min)):
|
|
_rate_state[state_key] = bucket
|
|
reset = max(1, int((bucket[0] + 60.0) - now)) if bucket else 60
|
|
return {
|
|
"allowed": False,
|
|
"limit": max(1, int(limit_per_min)),
|
|
"remaining": 0,
|
|
"reset": reset,
|
|
}
|
|
bucket.append(now)
|
|
_rate_state[state_key] = bucket
|
|
return {
|
|
"allowed": True,
|
|
"limit": max(1, int(limit_per_min)),
|
|
"remaining": max(0, int(limit_per_min) - len(bucket)),
|
|
"reset": max(1, int((bucket[0] + 60.0) - now)),
|
|
}
|
|
|
|
|
|
def _get_redis_rate_client():
|
|
global _redis_rate_client
|
|
if _redis_rate_client is not None:
|
|
return _redis_rate_client
|
|
with _redis_rate_lock:
|
|
if _redis_rate_client is None:
|
|
import redis
|
|
|
|
_redis_rate_client = redis.Redis.from_url(
|
|
cache_redis_url(),
|
|
socket_connect_timeout=0.25,
|
|
socket_timeout=0.25,
|
|
decode_responses=True,
|
|
)
|
|
return _redis_rate_client
|
|
|
|
|
|
def _rate_limit_one(key: str, limit_per_min: int) -> dict[str, int | bool]:
|
|
global _redis_rate_warned
|
|
limit = max(1, int(limit_per_min))
|
|
try:
|
|
if current_app.config.get("TESTING") or os.getenv("AGENT_RATE_LIMIT_BACKEND", "").lower() == "memory":
|
|
return _memory_rate_limit(key, limit)
|
|
except RuntimeError:
|
|
pass
|
|
|
|
minute = int(time.time() // 60)
|
|
redis_key = f"quantdinger:agent-rate:v1:{key}:{minute}"
|
|
try:
|
|
client = _get_redis_rate_client()
|
|
count, ttl = client.eval(
|
|
"""
|
|
local count = redis.call('INCR', KEYS[1])
|
|
if count == 1 then redis.call('EXPIRE', KEYS[1], ARGV[1]) end
|
|
local ttl = redis.call('TTL', KEYS[1])
|
|
return {count, ttl}
|
|
""",
|
|
1,
|
|
redis_key,
|
|
65,
|
|
)
|
|
count = int(count)
|
|
ttl = max(1, int(ttl))
|
|
return {
|
|
"allowed": count <= limit,
|
|
"limit": limit,
|
|
"remaining": max(0, limit - count),
|
|
"reset": ttl,
|
|
}
|
|
except Exception as exc:
|
|
if not _redis_rate_warned:
|
|
logger.warning("agent_auth: Redis rate limiter unavailable; using process-local fallback: %s", exc)
|
|
_redis_rate_warned = True
|
|
return _memory_rate_limit(key, limit)
|
|
|
|
|
|
def _check_rate_limit(token_id: int, user_id: int, limit_per_min: int) -> dict[str, int | bool]:
|
|
"""Enforce both per-token and aggregate tenant quotas."""
|
|
token_decision = _rate_limit_one(f"token:{int(token_id)}", limit_per_min)
|
|
if not bool(token_decision["allowed"]):
|
|
return token_decision
|
|
try:
|
|
tenant_limit = max(1, int(os.getenv("AGENT_TENANT_RATE_LIMIT_PER_MIN", "600")))
|
|
except Exception:
|
|
tenant_limit = 600
|
|
tenant_decision = _rate_limit_one(f"tenant:{int(user_id)}", tenant_limit)
|
|
if not bool(tenant_decision["allowed"]):
|
|
return tenant_decision
|
|
return token_decision
|
|
|
|
|
|
# ─────────────────────────── verification ───────────────────────────
|
|
|
|
def _extract_bearer() -> Optional[str]:
|
|
auth_header = request.headers.get("Authorization", "")
|
|
parts = auth_header.split()
|
|
if len(parts) == 2 and parts[0].lower() == "bearer":
|
|
return parts[1]
|
|
return None
|
|
|
|
|
|
def _lookup_token(raw_token: str) -> Optional[dict]:
|
|
token_hash = _hash_token(raw_token)
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"""
|
|
SELECT id, user_id, name, scopes, markets, instruments,
|
|
paper_only, rate_limit_per_min, max_order_notional,
|
|
max_daily_notional, status, expires_at
|
|
FROM qd_agent_tokens
|
|
WHERE token_hash = %s
|
|
""",
|
|
(token_hash,),
|
|
)
|
|
row = cur.fetchone()
|
|
cur.close()
|
|
return row
|
|
|
|
|
|
def _touch_token_last_used(token_id: int) -> None:
|
|
try:
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"UPDATE qd_agent_tokens SET last_used_at = NOW() WHERE id = %s",
|
|
(token_id,),
|
|
)
|
|
db.commit()
|
|
cur.close()
|
|
except Exception as exc:
|
|
logger.debug(f"agent_auth: failed to touch last_used_at: {exc}")
|
|
|
|
|
|
# ─────────────────────────── audit ───────────────────────────
|
|
|
|
_REDACT_KEYS = {
|
|
"password",
|
|
"secret",
|
|
"secretkey",
|
|
"secret_key",
|
|
"token",
|
|
"apikey",
|
|
"api_key",
|
|
"authorization",
|
|
"passphrase",
|
|
"privatekey",
|
|
"private_key",
|
|
"accesstoken",
|
|
"access_token",
|
|
"refreshtoken",
|
|
"refresh_token",
|
|
"bottoken",
|
|
"bot_token",
|
|
"webhooksecret",
|
|
"webhook_secret",
|
|
"signingsecret",
|
|
"signing_secret",
|
|
"clientsecret",
|
|
"client_secret",
|
|
}
|
|
|
|
|
|
def _redact(obj: Any, depth: int = 0) -> Any:
|
|
if depth > 3:
|
|
return "<truncated>"
|
|
if isinstance(obj, dict):
|
|
out = {}
|
|
for k, v in obj.items():
|
|
if str(k).replace("-", "_").lower() in _REDACT_KEYS:
|
|
out[k] = "<redacted>"
|
|
else:
|
|
out[k] = _redact(v, depth + 1)
|
|
return out
|
|
if isinstance(obj, list):
|
|
return [_redact(v, depth + 1) for v in obj[:20]]
|
|
if isinstance(obj, (str, int, float, bool)) or obj is None:
|
|
if isinstance(obj, str) and len(obj) > 500:
|
|
return obj[:500] + "..."
|
|
return obj
|
|
return str(type(obj).__name__)
|
|
|
|
|
|
def _audit(scope_class: str, status_code: int, response_summary: Any, duration_ms: int) -> None:
|
|
token_row = getattr(g, "agent_token", None) or {}
|
|
try:
|
|
req_summary: dict[str, Any] = {
|
|
"args": _redact(dict(request.args)),
|
|
}
|
|
if request.is_json:
|
|
try:
|
|
req_summary["json"] = _redact(request.get_json(silent=True) or {})
|
|
except Exception:
|
|
req_summary["json"] = "<unreadable>"
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"""
|
|
INSERT INTO qd_agent_audit
|
|
(user_id, agent_token_id, agent_name, route, method,
|
|
scope_class, status_code, idempotency_key,
|
|
request_summary, response_summary, duration_ms)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
""",
|
|
(
|
|
token_row.get("user_id") or 0,
|
|
token_row.get("id"),
|
|
token_row.get("name"),
|
|
request.path,
|
|
request.method,
|
|
scope_class,
|
|
int(status_code),
|
|
request.headers.get("Idempotency-Key"),
|
|
json.dumps(req_summary, default=str)[:8000],
|
|
json.dumps(_redact(response_summary), default=str)[:8000] if response_summary is not None else None,
|
|
int(duration_ms),
|
|
),
|
|
)
|
|
db.commit()
|
|
cur.close()
|
|
except Exception as exc:
|
|
logger.warning(f"agent_auth: audit insert failed: {exc}")
|
|
|
|
|
|
# ─────────────────────────── decorator ───────────────────────────
|
|
|
|
def _err(code: int, msg: str, details: Any = None, retriable: bool = False, status: int = 400):
|
|
body = {"code": code, "message": msg, "details": details, "retriable": retriable}
|
|
return jsonify(body), status
|
|
|
|
|
|
def _with_rate_headers(response, decision: dict[str, int | bool]):
|
|
response.headers["X-RateLimit-Limit"] = str(decision.get("limit", 0))
|
|
response.headers["X-RateLimit-Remaining"] = str(decision.get("remaining", 0))
|
|
response.headers["X-RateLimit-Reset"] = str(decision.get("reset", 0))
|
|
if response.status_code == 429:
|
|
response.headers["Retry-After"] = str(decision.get("reset", 1))
|
|
return response
|
|
|
|
|
|
def _request_fingerprint() -> str:
|
|
body = request.get_data(cache=True) or b""
|
|
raw = b"\n".join(
|
|
[
|
|
request.method.upper().encode("utf-8"),
|
|
request.path.encode("utf-8"),
|
|
request.query_string,
|
|
body,
|
|
]
|
|
)
|
|
return hashlib.sha256(raw).hexdigest()
|
|
|
|
|
|
def _reserve_idempotency(token_id: int, key: str) -> tuple[str, Optional[dict]]:
|
|
request_hash = _request_fingerprint()
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"""
|
|
INSERT INTO qd_agent_idempotency
|
|
(agent_token_id, method, route, idempotency_key, request_hash, status)
|
|
VALUES (%s, %s, %s, %s, %s, 'started')
|
|
ON CONFLICT (agent_token_id, method, route, idempotency_key) DO NOTHING
|
|
""",
|
|
(int(token_id), request.method.upper(), request.path, key, request_hash),
|
|
)
|
|
inserted = cur.rowcount > 0
|
|
if inserted:
|
|
db.commit()
|
|
cur.close()
|
|
return "reserved", None
|
|
cur.execute(
|
|
"""
|
|
SELECT request_hash, status, response_body, response_status, updated_at
|
|
FROM qd_agent_idempotency
|
|
WHERE agent_token_id = %s AND method = %s AND route = %s
|
|
AND idempotency_key = %s
|
|
""",
|
|
(int(token_id), request.method.upper(), request.path, key),
|
|
)
|
|
row = cur.fetchone()
|
|
if row and row.get("request_hash") == request_hash and row.get("status") == "started":
|
|
try:
|
|
stale_after = max(
|
|
60,
|
|
int(os.getenv("AGENT_IDEMPOTENCY_IN_PROGRESS_TTL_SEC", "900")),
|
|
)
|
|
except Exception:
|
|
stale_after = 900
|
|
cur.execute(
|
|
"""
|
|
UPDATE qd_agent_idempotency
|
|
SET updated_at = NOW()
|
|
WHERE agent_token_id = %s AND method = %s AND route = %s
|
|
AND idempotency_key = %s AND status = 'started'
|
|
AND updated_at < NOW() - (%s * INTERVAL '1 second')
|
|
""",
|
|
(
|
|
int(token_id),
|
|
request.method.upper(),
|
|
request.path,
|
|
key,
|
|
stale_after,
|
|
),
|
|
)
|
|
if cur.rowcount:
|
|
db.commit()
|
|
cur.close()
|
|
return "reserved", None
|
|
cur.close()
|
|
if not row:
|
|
return "in_progress", None
|
|
if row.get("request_hash") != request_hash:
|
|
return "mismatch", row
|
|
if row.get("status") == "completed":
|
|
return "completed", row
|
|
return "in_progress", row
|
|
|
|
|
|
def _complete_idempotency(token_id: int, key: str, response) -> None:
|
|
payload = response.get_json(silent=True)
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"""
|
|
UPDATE qd_agent_idempotency
|
|
SET status = 'completed', response_body = %s::jsonb,
|
|
response_status = %s, updated_at = NOW()
|
|
WHERE agent_token_id = %s AND method = %s AND route = %s
|
|
AND idempotency_key = %s
|
|
""",
|
|
(
|
|
json.dumps(payload, default=str),
|
|
int(response.status_code),
|
|
int(token_id),
|
|
request.method.upper(),
|
|
request.path,
|
|
key,
|
|
),
|
|
)
|
|
db.commit()
|
|
cur.close()
|
|
|
|
|
|
def agent_required(scope: str = SCOPE_R):
|
|
"""Flask decorator: enforce token auth + scope + rate limit + audit.
|
|
|
|
Sets `g.agent_token` (dict) and `g.agent_user_id` (int) for downstream code.
|
|
Logs every call (success or denial) into qd_agent_audit.
|
|
"""
|
|
if scope not in ALL_SCOPES:
|
|
raise ValueError(f"invalid scope: {scope}")
|
|
|
|
def decorator(fn: Callable):
|
|
@wraps(fn)
|
|
def wrapper(*args, **kwargs):
|
|
_ensure_schema()
|
|
t0 = time.time()
|
|
raw = _extract_bearer()
|
|
if not raw or not raw.startswith(TOKEN_PREFIX):
|
|
resp, code = _err(401, "Missing or malformed agent token", status=401)
|
|
_audit(scope, 401, {"reason": "missing_token"}, int((time.time() - t0) * 1000))
|
|
return resp, code
|
|
|
|
row = _lookup_token(raw)
|
|
if not row:
|
|
resp, code = _err(401, "Unknown agent token", status=401)
|
|
_audit(scope, 401, {"reason": "unknown_token"}, int((time.time() - t0) * 1000))
|
|
return resp, code
|
|
|
|
if row.get("status") != "active":
|
|
resp, code = _err(401, f"Token is {row.get('status')}", status=401)
|
|
_audit(scope, 401, {"reason": "inactive"}, int((time.time() - t0) * 1000))
|
|
return resp, code
|
|
|
|
expires_at = row.get("expires_at")
|
|
now = datetime.now(tz=expires_at.tzinfo) if isinstance(expires_at, datetime) else None
|
|
if expires_at and isinstance(expires_at, datetime) and now is not None and expires_at < now:
|
|
resp, code = _err(401, "Token expired", status=401)
|
|
_audit(scope, 401, {"reason": "expired"}, int((time.time() - t0) * 1000))
|
|
return resp, code
|
|
|
|
scopes = parse_scopes(row.get("scopes"))
|
|
if scope not in scopes:
|
|
g.agent_token = row
|
|
resp, code = _err(403, f"Token lacks required scope: {scope}", status=403)
|
|
_audit(scope, 403, {"granted": sorted(scopes)}, int((time.time() - t0) * 1000))
|
|
return resp, code
|
|
|
|
rate = _check_rate_limit(
|
|
row["id"],
|
|
row["user_id"],
|
|
int(row.get("rate_limit_per_min") or 60),
|
|
)
|
|
if not bool(rate["allowed"]):
|
|
g.agent_token = row
|
|
resp, code = _err(429, "Rate limit exceeded for this token", retriable=True, status=429)
|
|
_audit(scope, 429, {"limit_per_min": row.get("rate_limit_per_min")}, int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(make_response(resp, code), rate)
|
|
|
|
g.agent_token = row
|
|
g.agent_user_id = int(row["user_id"])
|
|
idempotency_key = (request.headers.get("Idempotency-Key") or "").strip()
|
|
needs_idempotency = request.method.upper() not in {"GET", "HEAD", "OPTIONS"} and scope in {
|
|
SCOPE_W, SCOPE_B, SCOPE_N, SCOPE_T,
|
|
}
|
|
if needs_idempotency:
|
|
if not idempotency_key:
|
|
response = make_response(*_err(
|
|
400,
|
|
"Idempotency-Key header is required for mutating agent calls",
|
|
status=400,
|
|
))
|
|
_audit(scope, 400, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
if len(idempotency_key) > 120:
|
|
response = make_response(*_err(400, "Idempotency-Key exceeds 120 characters", status=400))
|
|
_audit(scope, 400, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
try:
|
|
idem_state, idem_row = _reserve_idempotency(row["id"], idempotency_key)
|
|
except Exception as exc:
|
|
logger.error("agent_auth: idempotency reservation failed: %s", exc, exc_info=True)
|
|
response = make_response(*_err(
|
|
503,
|
|
"Idempotency service unavailable; request was not executed",
|
|
retriable=True,
|
|
status=503,
|
|
))
|
|
_audit(scope, 503, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
if idem_state == "mismatch":
|
|
response = make_response(*_err(
|
|
409,
|
|
"Idempotency-Key was already used with a different request",
|
|
status=409,
|
|
))
|
|
_audit(scope, 409, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
if idem_state == "in_progress":
|
|
response = make_response(*_err(
|
|
409,
|
|
"An identical request with this Idempotency-Key is still in progress",
|
|
retriable=True,
|
|
status=409,
|
|
))
|
|
_audit(scope, 409, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
if idem_state == "completed" and idem_row is not None:
|
|
response = make_response(
|
|
jsonify(idem_row.get("response_body")),
|
|
int(idem_row.get("response_status") or 200),
|
|
)
|
|
response.headers["Idempotent-Replayed"] = "true"
|
|
_audit(scope, response.status_code, response.get_json(silent=True), int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
|
|
try:
|
|
response = make_response(fn(*args, **kwargs))
|
|
except Exception as exc:
|
|
logger.error(f"agent route raised: {exc}", exc_info=True)
|
|
response = make_response(*_err(500, "Internal server error", details=str(exc), status=500))
|
|
|
|
status_code = int(response.status_code)
|
|
payload_summary: Any = response.get_json(silent=True)
|
|
if needs_idempotency and idempotency_key:
|
|
try:
|
|
_complete_idempotency(row["id"], idempotency_key, response)
|
|
except Exception as exc:
|
|
logger.error("agent_auth: failed to persist idempotent response: %s", exc, exc_info=True)
|
|
|
|
_touch_token_last_used(row["id"])
|
|
_audit(scope, status_code, payload_summary, int((time.time() - t0) * 1000))
|
|
return _with_rate_headers(response, rate)
|
|
|
|
return wrapper
|
|
|
|
return decorator
|
|
|
|
|
|
# ─────────────────────────── idempotency ───────────────────────────
|
|
|
|
@contextmanager
|
|
def with_idempotency(kind: str):
|
|
"""Context manager that yields an existing job dict if the same agent
|
|
already executed this kind+key, else yields None to indicate the caller
|
|
should perform the work and persist a new job row.
|
|
|
|
Use only on writeful (W/B/T) endpoints. Reads are naturally idempotent.
|
|
"""
|
|
token_row = getattr(g, "agent_token", None) or {}
|
|
key = request.headers.get("Idempotency-Key")
|
|
if not key or not token_row.get("id"):
|
|
yield None
|
|
return
|
|
|
|
try:
|
|
with get_db_connection() as db:
|
|
cur = db.cursor()
|
|
cur.execute(
|
|
"""
|
|
SELECT job_id, kind, status, request, result, error, created_at
|
|
FROM qd_agent_jobs
|
|
WHERE agent_token_id = %s AND kind = %s AND idempotency_key = %s
|
|
ORDER BY id DESC LIMIT 1
|
|
""",
|
|
(token_row["id"], kind, key),
|
|
)
|
|
existing = cur.fetchone()
|
|
cur.close()
|
|
except Exception as exc:
|
|
logger.warning(f"agent_auth: idempotency lookup failed: {exc}")
|
|
existing = None
|
|
|
|
yield existing
|
|
|
|
|
|
# ─────────────────────────── helpers for routes ───────────────────────────
|
|
|
|
def current_token() -> dict:
|
|
return getattr(g, "agent_token", {}) or {}
|
|
|
|
|
|
def current_user_id() -> int:
|
|
return int(getattr(g, "agent_user_id", 0) or 0)
|
|
|
|
|
|
def market_allowed(market: str) -> bool:
|
|
row = current_token()
|
|
return list_matches(market, parse_csv_list(row.get("markets"), default="*"))
|
|
|
|
|
|
def instrument_allowed(symbol: str) -> bool:
|
|
row = current_token()
|
|
return list_matches(symbol, parse_csv_list(row.get("instruments"), default="*"))
|
|
|
|
|
|
def paper_only() -> bool:
|
|
row = current_token()
|
|
return bool(row.get("paper_only", True))
|