feat: #85 T10 persistance etat (verrous stores, ratelimit SQLite) et cloture refonte

This commit is contained in:
2026-09-26 18:10:10 -04:00
parent 58312e64da
commit d6cca2b1af
18 changed files with 701 additions and 149 deletions
+33 -26
View File
@@ -119,25 +119,30 @@ def decode_token(token: str) -> dict | None:
_revoked_map: dict[str, int] = {}
_revoked_loaded = False
# ROADMAP #85 T10a — verrou autour du read-modify-write du store de
# révocation (perte de révocations en cas de logouts concurrents).
_revoked_lock = threading.RLock()
def _load_revoked():
"""Load revoked token JTIs from disk into memory (once)."""
global _revoked_loaded, _revoked_map
if _revoked_loaded:
return
if REVOKED_TOKENS_FILE.exists():
try:
data = json.loads(REVOKED_TOKENS_FILE.read_text())
# Drop entries whose underlying token has itself expired.
now = int(time.time())
_revoked_map = {
jti: int(exp) for jti, exp in data.items()
if int(exp) > now
}
except Exception as e:
logger.warning(f"Failed to load revoked tokens: {e}")
_revoked_map = {}
_revoked_loaded = True
with _revoked_lock:
if _revoked_loaded:
return
if REVOKED_TOKENS_FILE.exists():
try:
data = json.loads(REVOKED_TOKENS_FILE.read_text())
# Drop entries whose underlying token has itself expired.
now = int(time.time())
_revoked_map = {
jti: int(exp) for jti, exp in data.items()
if int(exp) > now
}
except Exception as e:
logger.warning(f"Failed to load revoked tokens: {e}")
_revoked_map = {}
_revoked_loaded = True
def _save_revoked():
@@ -154,24 +159,26 @@ def revoke_token(jti: str, expires_at: int | None = None):
``expires_at`` is the revoked token's own ``exp`` (unix seconds) — the
record is kept at least that long so a long-lived API token cannot
outlive its revocation. ``None`` means the token never expires (API/MCP
"sans fin") → the record is kept forever (capped at ~100 years, the JWT
"sans fin") → the record is kept forever (capped at ~100 years, the JWT
store's practical infinity). Default keeps 7 days (session tokens).
"""
_load_revoked()
now = int(time.time())
if expires_at is None:
until = now + 100 * 365 * 24 * 3600
else:
until = max(int(expires_at), now + REFRESH_TOKEN_EXPIRE_SECONDS)
_revoked_map[jti] = until
_save_revoked()
with _revoked_lock:
_load_revoked()
now = int(time.time())
if expires_at is None:
until = now + 100 * 365 * 24 * 3600
else:
until = max(int(expires_at), now + REFRESH_TOKEN_EXPIRE_SECONDS)
_revoked_map[jti] = until
_save_revoked()
logger.debug(f"Revoked token JTI: {jti[:8]}...")
def is_token_revoked(jti: str) -> bool:
"""Check if a token JTI has been revoked."""
_load_revoked()
return jti in _revoked_map
with _revoked_lock:
_load_revoked()
return jti in _revoked_map
# ---------------------------------------------------------------------------
+172 -1
View File
@@ -12,14 +12,24 @@ the per-account lockout in ``user_store.py``.
deployment, front this service with a shared store (Redis) or a single
worker. This limitation is intentional and documented (BUG-031).
Opt-in persistence (ROADMAP #85 T10b) : if ``OBSIGATE_RATELIMIT_DB`` points
to a SQLite file, counters are stored there instead (WAL mode, one short
connection per call — safe across threads, processes and restarts sharing
the same file). Semantics (windows, budgets, success reset) are identical
to the in-memory store, which remains the default when the variable is
unset.
Configuration via environment variables:
OBSIGATE_LOGIN_MAX_ATTEMPTS Max failures per IP (default: 10)
OBSIGATE_ACCOUNT_MAX_ATTEMPTS Max failures per account (default: 10)
OBSIGATE_LOGIN_WINDOW_SECONDS Lockout window in seconds (default: 900)
OBSIGATE_RATELIMIT_DB SQLite file for shared/persistent counters (default: unset = memory)
"""
import logging
import os
import sqlite3
import threading
import time
from collections import defaultdict
@@ -37,6 +47,127 @@ _last_cleanup = time.time()
CLEANUP_INTERVAL = 60 # seconds
def _db_path() -> str | None:
"""SQLite file for shared counters, or ``None`` for the in-memory store."""
path = os.environ.get("OBSIGATE_RATELIMIT_DB", "").strip()
return path or None
def _db_connect(path: str) -> sqlite3.Connection:
"""Open a short-lived connection (WAL + busy timeout for concurrent workers)."""
_db_ensure_schema(path)
conn = sqlite3.connect(path, timeout=10.0)
conn.execute("PRAGMA busy_timeout=10000")
return conn
_schema_ready: set[str] = set()
_schema_lock = threading.Lock()
def _db_ensure_schema(path: str) -> None:
"""Create the store schema once per file (DDL under a process-wide lock)."""
with _schema_lock:
if path in _schema_ready:
return
conn = sqlite3.connect(path, timeout=10.0)
try:
conn.execute("PRAGMA journal_mode=WAL")
conn.execute(
"CREATE TABLE IF NOT EXISTS attempts"
" (kind TEXT NOT NULL, key TEXT NOT NULL, ts REAL NOT NULL, success INTEGER NOT NULL)"
)
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_attempts_kind_key_ts"
" ON attempts (kind, key, ts)"
)
conn.commit()
finally:
conn.close()
_schema_ready.add(path)
def _db_write(fn, *args):
"""Run a write op, retrying once on lock contention (concurrent workers)."""
try:
return fn(*args)
except sqlite3.OperationalError as e:
if "locked" not in str(e).lower():
raise
time.sleep(0.05)
return fn(*args)
def _db_prune(conn: sqlite3.Connection, cutoff: float) -> None:
"""Drop expired entries (best-effort cap on disk growth)."""
conn.execute("DELETE FROM attempts WHERE ts <= ?", (cutoff,))
def _db_record(kind: str, key: str, success: bool) -> int:
"""Record one attempt in SQLite; return the live failure count."""
path = _db_path()
assert path is not None
now = time.time()
cutoff = now - WINDOW_SECONDS
def _write() -> int:
with _db_connect(path) as conn:
_db_prune(conn, cutoff)
if success:
# Mirror the in-memory reset: replace history with one success.
conn.execute("DELETE FROM attempts WHERE kind = ? AND key = ?", (kind, key))
conn.execute(
"INSERT INTO attempts (kind, key, ts, success) VALUES (?, ?, ?, ?)",
(kind, key, now, int(success)),
)
conn.commit()
(failures,) = conn.execute(
"SELECT COUNT(*) FROM attempts WHERE kind = ? AND key = ? AND ts > ? AND success = 0",
(kind, key, cutoff),
).fetchone()
return failures
return _db_write(_write)
def _db_failures(kind: str, key: str) -> int:
"""Live failure count in SQLite (expired entries never count)."""
path = _db_path()
assert path is not None
cutoff = time.time() - WINDOW_SECONDS
with _db_connect(path) as conn:
(failures,) = conn.execute(
"SELECT COUNT(*) FROM attempts WHERE kind = ? AND key = ? AND ts > ? AND success = 0",
(kind, key, cutoff),
).fetchone()
return failures
def _db_tracked(kind: str) -> int:
"""Number of distinct keys ever seen for one budget (SQLite)."""
path = _db_path()
assert path is not None
with _db_connect(path) as conn:
(n,) = conn.execute(
"SELECT COUNT(DISTINCT key) FROM attempts WHERE kind = ?", (kind,)
).fetchone()
return n
def _db_limited_count(kind: str, max_attempts: int) -> int:
"""Number of keys currently over budget (SQLite)."""
path = _db_path()
assert path is not None
cutoff = time.time() - WINDOW_SECONDS
with _db_connect(path) as conn:
rows = conn.execute(
"SELECT key, COUNT(*) FROM attempts"
" WHERE kind = ? AND ts > ? AND success = 0 GROUP BY key",
(kind, cutoff),
).fetchall()
return sum(1 for _, n in rows if n >= max_attempts)
def _prune(store: dict[str, list], cutoff: float) -> None:
"""Drop expired entries from one store in place."""
expired = []
@@ -66,6 +197,12 @@ def record_failure(ip: str) -> tuple[int, int]:
Returns:
(current_failure_count, remaining_attempts)
"""
if _db_path() is not None:
failures = _db_record("ip", ip, False)
remaining = max(0, MAX_ATTEMPTS - failures)
if failures >= MAX_ATTEMPTS:
logger.warning(f"IP {ip} rate-limited after {failures} failed logins")
return failures, remaining
_cleanup_expired()
_ip_attempts[ip].append((time.time(), False))
failures = sum(1 for _, success in _ip_attempts[ip] if not success)
@@ -77,12 +214,17 @@ def record_failure(ip: str) -> tuple[int, int]:
def record_success(ip: str):
"""Clear rate limit state for an IP after successful login."""
if _db_path() is not None:
_db_record("ip", ip, True)
return
_cleanup_expired()
_ip_attempts[ip] = [(time.time(), True)]
def is_rate_limited(ip: str) -> bool:
"""Check if an IP has exceeded the rate limit."""
if _db_path() is not None:
return _db_failures("ip", ip) >= MAX_ATTEMPTS
_cleanup_expired()
failures = sum(1 for _, success in _ip_attempts.get(ip, []) if not success)
return failures >= MAX_ATTEMPTS
@@ -94,8 +236,14 @@ def record_account_failure(account: str) -> tuple[int, int]:
Returns:
(current_failure_count, remaining_attempts)
"""
_cleanup_expired()
key = account.lower()
if _db_path() is not None:
failures = _db_record("account", key, False)
remaining = max(0, ACCOUNT_MAX_ATTEMPTS - failures)
if failures >= ACCOUNT_MAX_ATTEMPTS:
logger.warning(f"Account {account} rate-limited after {failures} failed attempts")
return failures, remaining
_cleanup_expired()
_account_attempts[key].append((time.time(), False))
failures = sum(1 for _, success in _account_attempts[key] if not success)
remaining = max(0, ACCOUNT_MAX_ATTEMPTS - failures)
@@ -106,12 +254,17 @@ def record_account_failure(account: str) -> tuple[int, int]:
def record_account_success(account: str):
"""Clear the per-account rate limit state after a successful login."""
if _db_path() is not None:
_db_record("account", account.lower(), True)
return
_cleanup_expired()
_account_attempts[account.lower()] = [(time.time(), True)]
def is_account_rate_limited(account: str) -> bool:
"""Check if an account has exceeded the per-account rate limit."""
if _db_path() is not None:
return _db_failures("account", account.lower()) >= ACCOUNT_MAX_ATTEMPTS
_cleanup_expired()
failures = sum(
1 for _, success in _account_attempts.get(account.lower(), []) if not success
@@ -121,6 +274,24 @@ def is_account_rate_limited(account: str) -> bool:
def get_status(ip: str | None = None) -> dict:
"""Get rate limit status for an IP (for diagnostics)."""
if _db_path() is not None:
if ip:
failures = _db_failures("ip", ip)
return {
"ip": ip,
"failures": failures,
"max": MAX_ATTEMPTS,
"limited": failures >= MAX_ATTEMPTS,
"window_seconds": WINDOW_SECONDS,
}
return {
"tracked_ips": _db_tracked("ip"),
"tracked_accounts": _db_tracked("account"),
"max_attempts": MAX_ATTEMPTS,
"account_max_attempts": ACCOUNT_MAX_ATTEMPTS,
"window_seconds": WINDOW_SECONDS,
"limited_ips": _db_limited_count("ip", MAX_ATTEMPTS),
}
_cleanup_expired()
if ip:
attempts = _ip_attempts.get(ip, [])
+48 -39
View File
@@ -10,6 +10,7 @@ No authentication required for public share views.
import json
import logging
import secrets
import threading
from datetime import datetime, timedelta, timezone
from pathlib import Path
@@ -17,6 +18,10 @@ logger = logging.getLogger("obsigate.share")
SHARES_FILE = Path("data/shares.json")
# ROADMAP #85 T10a — verrou autour des read-modify-write (perte de mises à
# jour en cas de créations/accès/révocations concurrents).
_lock = threading.RLock()
def _read() -> dict:
if not SHARES_FILE.exists():
@@ -41,26 +46,27 @@ def create_share(
expires_in_hours: int | None = None,
) -> dict:
"""Create a new share token for a document."""
data = _read()
token = secrets.token_hex(32) # 64-char hex token
with _lock:
data = _read()
token = secrets.token_hex(32) # 64-char hex token
expires_at = None
if expires_in_hours:
expires_at = (datetime.now(timezone.utc) + timedelta(hours=expires_in_hours)).isoformat()
expires_at = None
if expires_in_hours:
expires_at = (datetime.now(timezone.utc) + timedelta(hours=expires_in_hours)).isoformat()
share = {
"id": token,
"token": token,
"vault": vault,
"path": path,
"created_by": created_by,
"created_at": datetime.now(timezone.utc).isoformat(),
"expires_at": expires_at,
"access_count": 0,
"last_accessed": None,
}
data["shares"][token] = share
_write(data)
share = {
"id": token,
"token": token,
"vault": vault,
"path": path,
"created_by": created_by,
"created_at": datetime.now(timezone.utc).isoformat(),
"expires_at": expires_at,
"access_count": 0,
"last_accessed": None,
}
data["shares"][token] = share
_write(data)
logger.info(f"Created share for {vault}/{path} by {created_by}")
return share
@@ -80,22 +86,24 @@ def get_share_by_token(token: str) -> dict | None:
def record_access(token: str):
"""Increment access counter for a share."""
data = _read()
share = data["shares"].get(token)
if share:
share["access_count"] = share.get("access_count", 0) + 1
share["last_accessed"] = datetime.now(timezone.utc).isoformat()
_write(data)
with _lock:
data = _read()
share = data["shares"].get(token)
if share:
share["access_count"] = share.get("access_count", 0) + 1
share["last_accessed"] = datetime.now(timezone.utc).isoformat()
_write(data)
def revoke_share(share_id: str) -> bool:
"""Revoke (delete) a share by its token."""
data = _read()
if share_id in data["shares"]:
del data["shares"][share_id]
_write(data)
logger.info(f"Revoked share {share_id}")
return True
with _lock:
data = _read()
if share_id in data["shares"]:
del data["shares"][share_id]
_write(data)
logger.info(f"Revoked share {share_id}")
return True
return False
@@ -112,12 +120,13 @@ def list_shares(vault_filter: str | None = None) -> list:
def update_shares_after_rename(vault: str, old_path: str, new_path: str):
"""Update all shares when a file is renamed."""
data = _read()
updated = False
for sid, s in data["shares"].items():
if s.get("vault") == vault and s.get("path") == old_path:
s["path"] = new_path
updated = True
logger.info(f"Updated share {sid}: {vault}/{old_path} -> {new_path}")
if updated:
_write(data)
with _lock:
data = _read()
updated = False
for sid, s in data["shares"].items():
if s.get("vault") == vault and s.get("path") == old_path:
s["path"] = new_path
updated = True
logger.info(f"Updated share {sid}: {vault}/{old_path} -> {new_path}")
if updated:
_write(data)
+17 -11
View File
@@ -17,6 +17,7 @@ from __future__ import annotations
import json
import logging
import os
import threading
from pathlib import Path
logger = logging.getLogger("obsigate.tools.secrets")
@@ -34,6 +35,9 @@ TOOL_KEY_NAMES: tuple[str, ...] = (
_SECRET_MARKERS = ("API_KEY", "TOKEN")
# ROADMAP #85 T10a — verrou autour des read-modify-write du store de clés.
_lock = threading.RLock()
def _keys_file() -> Path:
base = os.environ.get("OBSIGATE_DATA_DIR", "data")
@@ -89,21 +93,23 @@ def set_tool_key(name: str, value: str) -> None:
if name not in TOOL_KEY_NAMES:
raise ValueError(f"Clé non prise en charge: {name}")
value = (value or "").strip()
keys = _read_keys()
if value:
keys[name] = value
else:
keys.pop(name, None)
_write_keys(keys)
with _lock:
keys = _read_keys()
if value:
keys[name] = value
else:
keys.pop(name, None)
_write_keys(keys)
def delete_tool_key(name: str) -> bool:
"""Remove one key from the store; return True when it existed."""
if name not in TOOL_KEY_NAMES:
raise ValueError(f"Clé non prise en charge: {name}")
keys = _read_keys()
if name in keys:
del keys[name]
_write_keys(keys)
return True
with _lock:
keys = _read_keys()
if name in keys:
del keys[name]
_write_keys(keys)
return True
return False
+54 -43
View File
@@ -26,6 +26,7 @@ import json
import logging
import os
import socket
import threading
import uuid
from datetime import datetime, timezone
from pathlib import Path
@@ -144,6 +145,12 @@ def _read_secrets() -> dict:
return {}
# ROADMAP #85 T10a — verrou autour des read-modify-write des deux stores
# (webhooks + secrets) : perte de mises à jour en cas de mutations
# concurrentes.
_lock = threading.RLock()
def _write_secrets(secrets: dict):
WEBHOOK_SECRETS_FILE.parent.mkdir(parents=True, exist_ok=True)
tmp = WEBHOOK_SECRETS_FILE.with_suffix(".tmp")
@@ -156,12 +163,13 @@ def _write_secrets(secrets: dict):
def _store_secret(wh_id: str, secret: str | None) -> None:
secrets = _read_secrets()
if secret:
secrets[wh_id] = secret
else:
secrets.pop(wh_id, None)
_write_secrets(secrets)
with _lock:
secrets = _read_secrets()
if secret:
secrets[wh_id] = secret
else:
secrets.pop(wh_id, None)
_write_secrets(secrets)
def _get_secret(wh: dict) -> str | None:
@@ -189,52 +197,55 @@ def get_webhooks() -> list:
def create_webhook(name: str, url: str, events: list[str], secret: str | None = None) -> dict:
validate_webhook_url(url)
webhooks = _read()
wh_id = str(uuid.uuid4())
wh = {
"id": wh_id,
"name": name,
"url": url,
"events": [e for e in events if e in VALID_EVENTS],
"enabled": True,
"created_at": datetime.now(timezone.utc).isoformat(),
"last_fired_at": None,
}
webhooks.append(wh)
_write(webhooks)
if secret:
_store_secret(wh_id, secret)
with _lock:
webhooks = _read()
wh_id = str(uuid.uuid4())
wh = {
"id": wh_id,
"name": name,
"url": url,
"events": [e for e in events if e in VALID_EVENTS],
"enabled": True,
"created_at": datetime.now(timezone.utc).isoformat(),
"last_fired_at": None,
}
webhooks.append(wh)
_write(webhooks)
if secret:
_store_secret(wh_id, secret)
logger.info(f"Created webhook '{name}' → {url}")
return _public_view(wh)
def update_webhook(wh_id: str, updates: dict) -> dict | None:
webhooks = _read()
for wh in webhooks:
if wh["id"] == wh_id:
if updates.get("url"):
validate_webhook_url(updates["url"])
if "secret" in updates:
_store_secret(wh_id, updates["secret"])
safe_updates = {
k: v for k, v in updates.items()
if k not in ("id", "secret")
}
wh.update(safe_updates)
_write(webhooks)
return _public_view(wh)
with _lock:
webhooks = _read()
for wh in webhooks:
if wh["id"] == wh_id:
if updates.get("url"):
validate_webhook_url(updates["url"])
if "secret" in updates:
_store_secret(wh_id, updates["secret"])
safe_updates = {
k: v for k, v in updates.items()
if k not in ("id", "secret")
}
wh.update(safe_updates)
_write(webhooks)
return _public_view(wh)
return None
def delete_webhook(wh_id: str) -> bool:
webhooks = _read()
new_list = [wh for wh in webhooks if wh["id"] != wh_id]
if len(new_list) == len(webhooks):
return False
_write(new_list)
secrets = _read_secrets()
if secrets.pop(wh_id, None) is not None:
_write_secrets(secrets)
with _lock:
webhooks = _read()
new_list = [wh for wh in webhooks if wh["id"] != wh_id]
if len(new_list) == len(webhooks):
return False
_write(new_list)
secrets = _read_secrets()
if secrets.pop(wh_id, None) is not None:
_write_secrets(secrets)
return True