Build reusable bot framework
This commit is contained in:
24
.dockerignore
Normal file
24
.dockerignore
Normal file
@@ -0,0 +1,24 @@
|
||||
.env
|
||||
.env.*
|
||||
config/.env
|
||||
config/.env.*
|
||||
|
||||
.git
|
||||
.github
|
||||
.gitignore
|
||||
__pycache__
|
||||
*.py[cod]
|
||||
.venv
|
||||
venv
|
||||
.coverage
|
||||
.coverage.*
|
||||
htmlcov
|
||||
.pytest_cache
|
||||
.ruff_cache
|
||||
build
|
||||
dist
|
||||
*.egg-info
|
||||
*.log
|
||||
*.pkl
|
||||
tests
|
||||
README.md
|
||||
42
.env.example
Normal file
42
.env.example
Normal file
@@ -0,0 +1,42 @@
|
||||
# Copy this file to .env and replace every placeholder before use.
|
||||
|
||||
# PostgreSQL
|
||||
DB_HOST=127.0.0.1
|
||||
DB_PORT=5432
|
||||
DB_NAME=app
|
||||
DB_USER=app
|
||||
DB_PASS=replace-with-a-database-password
|
||||
|
||||
# Authentication and internal service access
|
||||
JWT_SECRET=replace-with-a-long-random-jwt-secret
|
||||
BOT_API_KEY=replace-with-a-random-service-key-at-least-32-characters
|
||||
BOT_API_KEY_SCOPES=discord:session,outbox:claim,outbox:deliver
|
||||
|
||||
# Discord adapter
|
||||
DISCORD_BOT_TOKEN=replace-with-your-discord-bot-token
|
||||
DISCORD_ENROLLMENT_MODE=allowlist
|
||||
# Local default; Docker Compose overrides this inside the bot container.
|
||||
DISCORD_ALLOWLIST=123456789012345678
|
||||
API_URL=http://127.0.0.1:5000
|
||||
|
||||
# OpenAI-compatible model provider
|
||||
OPENROUTER_API_KEY=replace-with-your-provider-api-key
|
||||
OPENROUTER_BASE_URL=https://openrouter.ai/api/v1
|
||||
# AI_MODEL=provider/model-name
|
||||
|
||||
# Scheduler and durable outbox
|
||||
JOB_POLL_INTERVAL=5
|
||||
JOB_BATCH_SIZE=20
|
||||
JOB_LEASE_SECONDS=300
|
||||
OUTBOX_POLL_INTERVAL=5
|
||||
OUTBOX_BATCH_SIZE=20
|
||||
|
||||
# Reminder defaults
|
||||
DEFAULT_TIMEZONE=UTC
|
||||
|
||||
# Optional runtime settings
|
||||
LOG_LEVEL=INFO
|
||||
MAX_REQUEST_BYTES=1048576
|
||||
NOTIFICATION_TIMEOUT=10
|
||||
NTFY_BASE_URL=https://ntfy.sh
|
||||
# NTFY_TOKEN=
|
||||
71
.github/workflows/ci.yml
vendored
Normal file
71
.github/workflows/ci.yml
vendored
Normal file
@@ -0,0 +1,71 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
pull_request:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:16
|
||||
env:
|
||||
POSTGRES_DB: app
|
||||
POSTGRES_USER: app
|
||||
POSTGRES_PASSWORD: app
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U app -d app"
|
||||
--health-interval 5s
|
||||
--health-timeout 5s
|
||||
--health-retries 10
|
||||
env:
|
||||
DB_HOST: 127.0.0.1
|
||||
DB_PORT: 5432
|
||||
DB_NAME: app
|
||||
DB_USER: app
|
||||
DB_PASS: app
|
||||
JWT_SECRET: ci-only-jwt-secret
|
||||
BOT_API_KEY: ci-only-service-key-at-least-32-characters
|
||||
BOT_API_KEY_SCOPES: discord:session,outbox:claim,outbox:deliver
|
||||
DISCORD_ENROLLMENT_MODE: allowlist
|
||||
DISCORD_ALLOWLIST: "123456789"
|
||||
OPENROUTER_API_KEY: ci-only-provider-key
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
cache: pip
|
||||
cache-dependency-path: |
|
||||
requirements.txt
|
||||
requirements-dev.txt
|
||||
- name: Install dependencies
|
||||
run: python -m pip install -r requirements-dev.txt
|
||||
- name: Apply database migrations
|
||||
run: python -m core.migrations upgrade
|
||||
- name: Lint
|
||||
run: ruff check .
|
||||
- name: Test
|
||||
run: pytest --cov-fail-under=80
|
||||
|
||||
compose-smoke:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
ENV_FILE: .env.example
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Start migrated API and scheduler
|
||||
run: docker compose --env-file .env.example up --build --wait --wait-timeout 180 app scheduler
|
||||
- name: Check API health
|
||||
run: |
|
||||
curl --fail --show-error http://localhost:8080/health/live
|
||||
curl --fail --show-error http://localhost:8080/health/ready
|
||||
- name: Show service logs after failure
|
||||
if: failure()
|
||||
run: docker compose --env-file .env.example logs
|
||||
- name: Stop services
|
||||
if: always()
|
||||
run: docker compose --env-file .env.example down --volumes
|
||||
38
.gitignore
vendored
Normal file
38
.gitignore
vendored
Normal file
@@ -0,0 +1,38 @@
|
||||
# Local configuration and credentials
|
||||
.env
|
||||
.env.*
|
||||
!.env.example
|
||||
!**/.env.example
|
||||
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.pyd
|
||||
.venv/
|
||||
venv/
|
||||
|
||||
# Test and quality-tool output
|
||||
.coverage
|
||||
.coverage.*
|
||||
htmlcov/
|
||||
.pytest_cache/
|
||||
.ruff_cache/
|
||||
|
||||
# Local automation and agent state
|
||||
.agents/
|
||||
.codex/
|
||||
|
||||
# Packaging
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
|
||||
# Runtime state and logs
|
||||
*.log
|
||||
*.pkl
|
||||
|
||||
# Editors and operating systems
|
||||
.idea/
|
||||
.vscode/
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
15
Dockerfile
Normal file
15
Dockerfile
Normal file
@@ -0,0 +1,15 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY requirements.txt .
|
||||
RUN python -m pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
COPY . .
|
||||
|
||||
EXPOSE 5000
|
||||
|
||||
CMD ["gunicorn", "--bind", "0.0.0.0:5000", "--workers", "2", "--access-logfile", "-", "--error-logfile", "-", "api.main:app"]
|
||||
1
ai/__init__.py
Normal file
1
ai/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""LLM parsing support."""
|
||||
16
ai/ai_config.json
Normal file
16
ai/ai_config.json
Normal file
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"model": "qwen/qwen3-next-80b-a3b-thinking:nitro",
|
||||
"max_tokens": 2048,
|
||||
"json_mode": false,
|
||||
"prompts": {
|
||||
"command_parser": {
|
||||
"system": "You route user messages to one available command module. Return only a JSON object. Never invent a module name.",
|
||||
"user_template": "Available modules:\n{module_context}\n\nConversation context:\n{history_context}\n\nUser message: \"{user_input}\"\nCurrent UTC time: {current_time}\nUser timezone: {timezone}\n\nReturn exactly one JSON object with interaction_type, confidence from 0 to 1, and needs_clarification when the module is unclear."
|
||||
}
|
||||
},
|
||||
"validation": {
|
||||
"max_retries": 3,
|
||||
"timeout_seconds": 15,
|
||||
"confidence_threshold": 0.8
|
||||
}
|
||||
}
|
||||
318
ai/parser.py
Normal file
318
ai/parser.py
Normal file
@@ -0,0 +1,318 @@
|
||||
"""
|
||||
parser.py - LLM-powered JSON parsing with retry and validation
|
||||
|
||||
The framework first routes a message to a discovered feature module, then
|
||||
uses that module's focused prompt and Python validator to parse its action.
|
||||
Both synchronous and asynchronous entrypoints are kept for reusable modules.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
load_dotenv(Path(__file__).resolve().parents[1] / ".env", override=False)
|
||||
|
||||
CONFIG_PATH = os.environ.get(
|
||||
"AI_CONFIG_PATH", os.path.join(os.path.dirname(__file__), "ai_config.json")
|
||||
)
|
||||
|
||||
with open(CONFIG_PATH, "r", encoding="utf-8") as config_file:
|
||||
AI_CONFIG = json.load(config_file)
|
||||
|
||||
VALIDATORS = {}
|
||||
_sync_client = None
|
||||
_async_client = None
|
||||
|
||||
|
||||
def _get_client(async_client=False):
|
||||
global _sync_client, _async_client
|
||||
client_class = AsyncOpenAI if async_client else OpenAI
|
||||
current = _async_client if async_client else _sync_client
|
||||
if current is None:
|
||||
current = client_class(
|
||||
api_key=os.getenv("OPENROUTER_API_KEY"),
|
||||
base_url=os.getenv(
|
||||
"OPENROUTER_BASE_URL", "https://openrouter.ai/api/v1"
|
||||
),
|
||||
)
|
||||
if async_client:
|
||||
_async_client = current
|
||||
else:
|
||||
_sync_client = current
|
||||
return current
|
||||
|
||||
|
||||
def _extract_json_from_text(text):
|
||||
"""Decode the first complete JSON object, including nested objects."""
|
||||
if not isinstance(text, str):
|
||||
return None
|
||||
|
||||
stripped = text.strip()
|
||||
fence = chr(96) * 3
|
||||
if stripped.startswith(fence):
|
||||
stripped = re.sub(
|
||||
rf"^{re.escape(fence)}(?:json)?\s*",
|
||||
"",
|
||||
stripped,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
stripped = re.sub(rf"\s*{re.escape(fence)}$", "", stripped)
|
||||
|
||||
decoder = json.JSONDecoder()
|
||||
for index, character in enumerate(stripped):
|
||||
if character not in "[{":
|
||||
continue
|
||||
try:
|
||||
value, _ = decoder.raw_decode(stripped[index:])
|
||||
return value
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
def _render_template(template, values):
|
||||
"""Replace named placeholders without treating literal JSON braces as fields."""
|
||||
pattern = re.compile(r"\{([A-Za-z_][A-Za-z0-9_]*)\}")
|
||||
|
||||
def replace(match):
|
||||
key = match.group(1)
|
||||
return str(values[key]) if key in values else match.group(0)
|
||||
|
||||
return pattern.sub(replace, template)
|
||||
|
||||
|
||||
def _response_text(response):
|
||||
if not response.choices:
|
||||
return None
|
||||
message = response.choices[0].message
|
||||
if message.content:
|
||||
return message.content.strip()
|
||||
reasoning = getattr(message, "reasoning", None)
|
||||
return reasoning.strip() if reasoning else None
|
||||
|
||||
|
||||
def _request_args(system_prompt, user_prompt):
|
||||
args = {
|
||||
"model": os.getenv("AI_MODEL", AI_CONFIG["model"]),
|
||||
"max_tokens": AI_CONFIG.get("max_tokens", 2048),
|
||||
"timeout": AI_CONFIG["validation"].get("timeout_seconds", 15),
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
}
|
||||
if AI_CONFIG.get("json_mode", False):
|
||||
args["response_format"] = {"type": "json_object"}
|
||||
return args
|
||||
|
||||
|
||||
def _call_llm(system_prompt, user_prompt):
|
||||
"""Call an OpenAI-compatible API and return response text or None."""
|
||||
try:
|
||||
response = _get_client().chat.completions.create(
|
||||
**_request_args(system_prompt, user_prompt)
|
||||
)
|
||||
return _response_text(response)
|
||||
except Exception as error:
|
||||
logger.warning("LLM call failed: %s: %s", type(error).__name__, error)
|
||||
return None
|
||||
|
||||
|
||||
async def _call_llm_async(system_prompt, user_prompt):
|
||||
"""Asynchronously call an OpenAI-compatible API."""
|
||||
try:
|
||||
response = await _get_client(async_client=True).chat.completions.create(
|
||||
**_request_args(system_prompt, user_prompt)
|
||||
)
|
||||
return _response_text(response)
|
||||
except Exception as error:
|
||||
logger.warning("LLM call failed: %s: %s", type(error).__name__, error)
|
||||
return None
|
||||
|
||||
|
||||
def _history_context(history):
|
||||
if not history:
|
||||
return "No previous context"
|
||||
|
||||
history_lines = []
|
||||
for index, (message, result) in enumerate(history[-3:]):
|
||||
history_lines.append(f"{index + 1}. User: {message}")
|
||||
history_lines.append(f" Parsed: {json.dumps(result, default=str)}")
|
||||
return "\n".join(history_lines)
|
||||
|
||||
|
||||
def _build_prompt(user_input, prompt_config, history, errors, template_values):
|
||||
values = {
|
||||
"user_input": user_input,
|
||||
"history_context": _history_context(history),
|
||||
}
|
||||
values.update(template_values or {})
|
||||
user_prompt = _render_template(prompt_config["user_template"], values)
|
||||
if errors:
|
||||
user_prompt += (
|
||||
"\n\nThe previous response failed validation:\n- "
|
||||
+ "\n- ".join(str(error) for error in errors)
|
||||
+ "\nReturn a corrected JSON object."
|
||||
)
|
||||
return user_prompt
|
||||
|
||||
|
||||
def _get_prompt(interaction_type, prompt_override=None):
|
||||
if prompt_override:
|
||||
return prompt_override
|
||||
return AI_CONFIG.get("prompts", {}).get(interaction_type)
|
||||
|
||||
|
||||
def _validation_errors(parsed, validator):
|
||||
if not isinstance(parsed, dict):
|
||||
return ["Response must be a JSON object"]
|
||||
if validator:
|
||||
return list(validator(parsed) or [])
|
||||
return []
|
||||
|
||||
|
||||
def parse(
|
||||
user_input,
|
||||
interaction_type,
|
||||
retry_count=0,
|
||||
errors=None,
|
||||
history=None,
|
||||
prompt_override=None,
|
||||
validator=None,
|
||||
template_values=None,
|
||||
):
|
||||
"""Synchronously parse one prompt into a validated dictionary."""
|
||||
prompt_config = _get_prompt(interaction_type, prompt_override)
|
||||
if not prompt_config:
|
||||
return {"error": f"Unknown interaction type: {interaction_type}"}
|
||||
|
||||
validator = validator or VALIDATORS.get(interaction_type)
|
||||
max_attempts = AI_CONFIG["validation"].get("max_retries", 3)
|
||||
attempt = retry_count
|
||||
current_errors = errors
|
||||
|
||||
while attempt < max_attempts:
|
||||
user_prompt = _build_prompt(
|
||||
user_input, prompt_config, history, current_errors, template_values
|
||||
)
|
||||
response_text = _call_llm(prompt_config["system"], user_prompt)
|
||||
if not response_text:
|
||||
return {"error": "AI service unavailable", "user_input": user_input}
|
||||
|
||||
parsed = _extract_json_from_text(response_text)
|
||||
current_errors = (
|
||||
["Response was not valid JSON"]
|
||||
if parsed is None
|
||||
else _validation_errors(parsed, validator)
|
||||
)
|
||||
if not current_errors:
|
||||
return parsed
|
||||
attempt += 1
|
||||
|
||||
return {
|
||||
"error": f"Failed to parse after {max_attempts} attempts",
|
||||
"validation_errors": current_errors or [],
|
||||
"user_input": user_input,
|
||||
}
|
||||
|
||||
|
||||
async def parse_async(
|
||||
user_input,
|
||||
interaction_type,
|
||||
retry_count=0,
|
||||
errors=None,
|
||||
history=None,
|
||||
prompt_override=None,
|
||||
validator=None,
|
||||
template_values=None,
|
||||
):
|
||||
"""Asynchronously parse one prompt into a validated dictionary."""
|
||||
prompt_config = _get_prompt(interaction_type, prompt_override)
|
||||
if not prompt_config:
|
||||
return {"error": f"Unknown interaction type: {interaction_type}"}
|
||||
|
||||
validator = validator or VALIDATORS.get(interaction_type)
|
||||
max_attempts = AI_CONFIG["validation"].get("max_retries", 3)
|
||||
attempt = retry_count
|
||||
current_errors = errors
|
||||
|
||||
while attempt < max_attempts:
|
||||
user_prompt = _build_prompt(
|
||||
user_input, prompt_config, history, current_errors, template_values
|
||||
)
|
||||
response_text = await _call_llm_async(prompt_config["system"], user_prompt)
|
||||
if not response_text:
|
||||
return {"error": "AI service unavailable", "user_input": user_input}
|
||||
|
||||
parsed = _extract_json_from_text(response_text)
|
||||
current_errors = (
|
||||
["Response was not valid JSON"]
|
||||
if parsed is None
|
||||
else _validation_errors(parsed, validator)
|
||||
)
|
||||
if not current_errors:
|
||||
return parsed
|
||||
attempt += 1
|
||||
|
||||
return {
|
||||
"error": f"Failed to parse after {max_attempts} attempts",
|
||||
"validation_errors": current_errors or [],
|
||||
"user_input": user_input,
|
||||
}
|
||||
|
||||
|
||||
async def parse_command_async(user_input, module_registry, history=None, timezone_name="UTC"):
|
||||
"""Route a command, then parse it with the selected feature prompt."""
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
template_values = {
|
||||
"module_context": module_registry.router_context(),
|
||||
"current_time": now,
|
||||
"timezone": timezone_name,
|
||||
}
|
||||
routed = await parse_async(
|
||||
user_input,
|
||||
"command_parser",
|
||||
history=history,
|
||||
template_values=template_values,
|
||||
)
|
||||
if routed.get("error") or routed.get("needs_clarification"):
|
||||
return routed
|
||||
|
||||
threshold = AI_CONFIG["validation"].get("confidence_threshold", 0.8)
|
||||
confidence = routed.get("confidence")
|
||||
if isinstance(confidence, (int, float)) and confidence < threshold:
|
||||
return {
|
||||
"needs_clarification": "Could you rephrase that with a little more detail?",
|
||||
"confidence": confidence,
|
||||
}
|
||||
|
||||
interaction_type = routed.get("interaction_type")
|
||||
command = module_registry.get_command(interaction_type)
|
||||
if not command:
|
||||
return {"error": f"Unknown command type: {interaction_type}"}
|
||||
|
||||
parsed = await parse_async(
|
||||
user_input,
|
||||
interaction_type,
|
||||
history=history,
|
||||
prompt_override=command["prompt"],
|
||||
validator=command["validator"],
|
||||
template_values=template_values,
|
||||
)
|
||||
if isinstance(parsed, dict) and not parsed.get("error"):
|
||||
parsed["interaction_type"] = interaction_type
|
||||
return parsed
|
||||
|
||||
|
||||
def register_validator(interaction_type, validator_fn):
|
||||
"""Keep the original validator registration API for direct parser users."""
|
||||
VALIDATORS[interaction_type] = validator_fn
|
||||
1
api/__init__.py
Normal file
1
api/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Flask API package."""
|
||||
308
api/main.py
Normal file
308
api/main.py
Normal file
@@ -0,0 +1,308 @@
|
||||
"""Flask API and application factory for the reusable bot framework."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import flask
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from api.security import jsonObject, requireService, requireUser
|
||||
from core import api_keys as apiKeys
|
||||
from core import auth, identity, outbox, postgres, users
|
||||
from core.registry import discover_modules
|
||||
|
||||
|
||||
load_dotenv(Path(__file__).resolve().parents[1] / ".env", override=False)
|
||||
|
||||
logging.basicConfig(level=os.getenv("LOG_LEVEL", "INFO"))
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def createApp():
|
||||
app = flask.Flask(__name__)
|
||||
app.config["MAX_CONTENT_LENGTH"] = int(
|
||||
os.getenv("MAX_REQUEST_BYTES", 1024 * 1024)
|
||||
)
|
||||
app.config["SERVICE_KEY_READY"] = False
|
||||
|
||||
moduleRegistry = discover_modules()
|
||||
_registerCoreRoutes(app)
|
||||
for moduleName, registerRoutes in moduleRegistry.route_registrars:
|
||||
logger.info("Registering routes for module %s", moduleName)
|
||||
registerRoutes(app)
|
||||
|
||||
@app.before_request
|
||||
def bootstrapConfiguredServiceKey():
|
||||
if flask.request.path == "/health/live":
|
||||
return None
|
||||
if not app.config["SERVICE_KEY_READY"] and os.getenv("BOT_API_KEY"):
|
||||
apiKeys.bootstrapServiceApiKey()
|
||||
app.config["SERVICE_KEY_READY"] = True
|
||||
return None
|
||||
|
||||
@app.errorhandler(404)
|
||||
def notFound(_error):
|
||||
return flask.jsonify({"error": "not found"}), 404
|
||||
|
||||
@app.errorhandler(405)
|
||||
def methodNotAllowed(_error):
|
||||
return flask.jsonify({"error": "method not allowed"}), 405
|
||||
|
||||
@app.errorhandler(413)
|
||||
def requestTooLarge(_error):
|
||||
return flask.jsonify({"error": "request too large"}), 413
|
||||
|
||||
@app.errorhandler(Exception)
|
||||
def unexpectedError(error):
|
||||
logger.exception("Unhandled API error", exc_info=error)
|
||||
return flask.jsonify({"error": "internal server error"}), 500
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def _registerCoreRoutes(app):
|
||||
@app.route("/api/register", methods=["POST"])
|
||||
def api_register():
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
try:
|
||||
registered = users.registerUser(
|
||||
data.get("username"), data.get("password"), data
|
||||
)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
if not registered:
|
||||
return flask.jsonify({"error": "username taken"}), 409
|
||||
return flask.jsonify({"success": True}), 201
|
||||
|
||||
@app.route("/api/login", methods=["POST"])
|
||||
def api_login():
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
token = auth.getLoginToken(data.get("username"), data.get("password"))
|
||||
if not token:
|
||||
return flask.jsonify({"error": "invalid credentials"}), 401
|
||||
return flask.jsonify({"token": token}), 200
|
||||
|
||||
@app.route("/api/auth/discord/session", methods=["POST"])
|
||||
@requireService("discord:session")
|
||||
def api_discordSession():
|
||||
data = jsonObject()
|
||||
if data is None or not data.get("discord_id"):
|
||||
return flask.jsonify({"error": "discord_id required"}), 400
|
||||
try:
|
||||
user = identity.getOrCreateDiscordUser(
|
||||
data["discord_id"], displayName=data.get("display_name")
|
||||
)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
if not user:
|
||||
return flask.jsonify({"error": "Discord user is not allowed"}), 403
|
||||
token = auth.createLoginToken(
|
||||
str(user["id"]),
|
||||
name=data.get("display_name"),
|
||||
extraClaims={"provider": "discord"},
|
||||
)
|
||||
return flask.jsonify(
|
||||
{
|
||||
"token": token,
|
||||
"user_uuid": user["id"],
|
||||
"timezone": user.get("timezone") or "UTC",
|
||||
}
|
||||
), 200
|
||||
|
||||
@app.route("/api/keys", methods=["GET"])
|
||||
@requireUser(requireLogin=True)
|
||||
def api_listKeys():
|
||||
return flask.jsonify(
|
||||
{"keys": apiKeys.listUserApiKeys(flask.g.user_uuid)}
|
||||
), 200
|
||||
|
||||
@app.route("/api/keys", methods=["POST"])
|
||||
@requireUser(requireLogin=True)
|
||||
def api_createKey():
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
try:
|
||||
key = apiKeys.createUserApiKey(
|
||||
flask.g.user_uuid,
|
||||
data.get("name"),
|
||||
expiresAt=data.get("expires_at"),
|
||||
)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
return flask.jsonify(key), 201
|
||||
|
||||
@app.route("/api/keys/<keyID>", methods=["DELETE"])
|
||||
@requireUser(requireLogin=True)
|
||||
def api_revokeKey(keyID):
|
||||
if not apiKeys.revokeUserApiKey(flask.g.user_uuid, keyID):
|
||||
return flask.jsonify({"error": "API key not found"}), 404
|
||||
return flask.jsonify({"success": True}), 200
|
||||
|
||||
@app.route("/api/getUserUUID/<username>", methods=["GET"])
|
||||
@requireUser()
|
||||
def api_getUserUUID(username):
|
||||
userUUID = users.getUserUUID(username)
|
||||
if not userUUID:
|
||||
return flask.jsonify({"error": "user not found"}), 404
|
||||
if str(userUUID) != str(flask.g.user_uuid):
|
||||
return flask.jsonify({"error": "unauthorized"}), 403
|
||||
return flask.jsonify(userUUID), 200
|
||||
|
||||
@app.route("/api/user/<userUUID>", methods=["GET"])
|
||||
@requireUser()
|
||||
def api_getUser(userUUID):
|
||||
if str(userUUID) != str(flask.g.user_uuid):
|
||||
return flask.jsonify({"error": "unauthorized"}), 403
|
||||
user = users.getUser(userUUID)
|
||||
if not user:
|
||||
return flask.jsonify({"error": "user not found"}), 404
|
||||
user.pop("password_hashed", None)
|
||||
return flask.jsonify(user), 200
|
||||
|
||||
@app.route("/api/user/<userUUID>", methods=["PUT"])
|
||||
@requireUser()
|
||||
def api_updateUser(userUUID):
|
||||
if str(userUUID) != str(flask.g.user_uuid):
|
||||
return flask.jsonify({"error": "unauthorized"}), 403
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
try:
|
||||
updated = users.updateUser(userUUID, data)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
if not updated:
|
||||
return flask.jsonify({"error": "no valid fields to update"}), 400
|
||||
return flask.jsonify({"success": True}), 200
|
||||
|
||||
@app.route("/api/user/<userUUID>", methods=["DELETE"])
|
||||
@requireUser(requireLogin=True)
|
||||
def api_deleteUser(userUUID):
|
||||
if str(userUUID) != str(flask.g.user_uuid):
|
||||
return flask.jsonify({"error": "unauthorized"}), 403
|
||||
data = jsonObject()
|
||||
if data is None or not data.get("password"):
|
||||
return flask.jsonify(
|
||||
{"error": "password required for account deletion"}
|
||||
), 400
|
||||
if not auth.unregisterUser(userUUID, data["password"]):
|
||||
return flask.jsonify({"error": "invalid password"}), 401
|
||||
return flask.jsonify({"success": True}), 200
|
||||
|
||||
@app.route("/api/user/me/timezone", methods=["GET", "PUT"])
|
||||
@requireUser()
|
||||
def api_userTimezone():
|
||||
if flask.request.method == "GET":
|
||||
return flask.jsonify(
|
||||
{"timezone": users.getUserTimezone(flask.g.user_uuid)}
|
||||
), 200
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
try:
|
||||
timezoneName = users.normalizeTimezone(data.get("timezone"))
|
||||
users.updateUser(flask.g.user_uuid, {"timezone": timezoneName})
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
return flask.jsonify({"timezone": timezoneName}), 200
|
||||
|
||||
@app.route("/api/internal/outbox/claim", methods=["POST"])
|
||||
@requireService("outbox:claim")
|
||||
def api_claimOutbox():
|
||||
data = jsonObject() or {}
|
||||
workerID = data.get("worker_id")
|
||||
if not workerID:
|
||||
return flask.jsonify({"error": "worker_id required"}), 400
|
||||
try:
|
||||
claimed = outbox.claim_messages(
|
||||
workerID,
|
||||
channel=data.get("channel", "discord_dm"),
|
||||
limit=data.get("limit", 20),
|
||||
)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
|
||||
messages = []
|
||||
for message in claimed:
|
||||
discordIdentity = _discordIdentityForUser(message["user_uuid"])
|
||||
if not discordIdentity:
|
||||
outbox.retry_message(
|
||||
message["id"], workerID, "user has no Discord identity"
|
||||
)
|
||||
continue
|
||||
payload = message.get("payload") or {}
|
||||
messages.append(
|
||||
{
|
||||
"id": message["id"],
|
||||
"provider_user_id": discordIdentity["provider_user_id"],
|
||||
"content": payload.get("content", ""),
|
||||
"attempts": message["attempts"],
|
||||
"worker_id": workerID,
|
||||
}
|
||||
)
|
||||
return flask.jsonify({"messages": messages}), 200
|
||||
|
||||
@app.route("/api/internal/outbox/<messageID>/result", methods=["POST"])
|
||||
@requireService("outbox:deliver")
|
||||
def api_outboxResult(messageID):
|
||||
data = jsonObject()
|
||||
if data is None or data.get("status") not in {"sent", "retry"}:
|
||||
return flask.jsonify({"error": "valid status required"}), 400
|
||||
message = outbox.get_message(messageID)
|
||||
if not message or message.get("status") != "delivering":
|
||||
return flask.jsonify({"error": "leased message not found"}), 404
|
||||
workerID = message.get("leased_by")
|
||||
if not data.get("worker_id") or data["worker_id"] != workerID:
|
||||
return flask.jsonify({"error": "message lease is owned by another worker"}), 409
|
||||
if data["status"] == "sent":
|
||||
externalMessageID = data.get("external_message_id")
|
||||
if externalMessageID is not None:
|
||||
externalMessageID = str(externalMessageID)[:255]
|
||||
updated = outbox.mark_delivered(
|
||||
messageID,
|
||||
workerID,
|
||||
external_message_id=externalMessageID,
|
||||
)
|
||||
else:
|
||||
updated = outbox.retry_message(
|
||||
messageID, workerID, data.get("error", "delivery failed")
|
||||
)
|
||||
if not updated:
|
||||
return flask.jsonify({"error": "message lease expired"}), 409
|
||||
return flask.jsonify({"success": True, "message": updated}), 200
|
||||
|
||||
@app.route("/health/live", methods=["GET"])
|
||||
def healthLive():
|
||||
return flask.jsonify({"status": "ok"}), 200
|
||||
|
||||
@app.route("/health/ready", methods=["GET"])
|
||||
def healthReady():
|
||||
try:
|
||||
postgres.execute("SELECT 1 AS ready")
|
||||
except Exception:
|
||||
return flask.jsonify({"status": "not ready"}), 503
|
||||
return flask.jsonify({"status": "ready"}), 200
|
||||
|
||||
@app.route("/health", methods=["GET"])
|
||||
def healthCompatibility():
|
||||
return healthReady()
|
||||
|
||||
|
||||
def _discordIdentityForUser(userUUID):
|
||||
identities = identity.listUserIdentities(userUUID)
|
||||
return next(
|
||||
(item for item in identities if item.get("provider") == "discord"), None
|
||||
)
|
||||
|
||||
|
||||
app = createApp()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
app.run(host="0.0.0.0", port=5000)
|
||||
50
api/security.py
Normal file
50
api/security.py
Normal file
@@ -0,0 +1,50 @@
|
||||
"""Authentication decorators shared by API route modules."""
|
||||
|
||||
from functools import wraps
|
||||
|
||||
import flask
|
||||
|
||||
from core import auth
|
||||
|
||||
|
||||
def requireUser(requireLogin=False):
|
||||
def decorator(route):
|
||||
@wraps(route)
|
||||
def wrapped(*args, **kwargs):
|
||||
principal = auth.authenticateBearerToken(
|
||||
flask.request.headers.get("Authorization"),
|
||||
allowService=False,
|
||||
)
|
||||
if not auth.isUserPrincipal(principal, requireLogin=requireLogin):
|
||||
return flask.jsonify({"error": "unauthorized"}), 401
|
||||
flask.g.principal = principal
|
||||
flask.g.user_uuid = principal["user_uuid"]
|
||||
return route(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def requireService(scope):
|
||||
def decorator(route):
|
||||
@wraps(route)
|
||||
def wrapped(*args, **kwargs):
|
||||
principal = auth.authenticateBearerToken(
|
||||
flask.request.headers.get("Authorization"),
|
||||
requiredScopes=[scope],
|
||||
allowUser=False,
|
||||
)
|
||||
if not auth.hasServiceScope(principal, scope):
|
||||
return flask.jsonify({"error": "unauthorized"}), 401
|
||||
flask.g.principal = principal
|
||||
return route(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def jsonObject():
|
||||
data = flask.request.get_json(silent=True)
|
||||
return data if isinstance(data, dict) else None
|
||||
1
bot/__init__.py
Normal file
1
bot/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Messaging-platform adapters."""
|
||||
99
bot/api_client.py
Normal file
99
bot/api_client.py
Normal file
@@ -0,0 +1,99 @@
|
||||
"""Async API client used by the Discord adapter."""
|
||||
|
||||
import os
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class ApiClient:
|
||||
def __init__(self, discord_id, display_name):
|
||||
self.discord_id = str(discord_id)
|
||||
self.display_name = display_name
|
||||
self.api_url = os.getenv("API_URL", "http://127.0.0.1:5000").rstrip("/")
|
||||
self.service_key = os.getenv("BOT_API_KEY", "")
|
||||
self.token = None
|
||||
self.user_uuid = None
|
||||
self.timezone = os.getenv("DEFAULT_TIMEZONE", "UTC")
|
||||
self.http = httpx.AsyncClient(timeout=10.0)
|
||||
|
||||
async def authenticate(self):
|
||||
headers = {"Authorization": f"Bearer {self.service_key}"}
|
||||
try:
|
||||
response = await self.http.post(
|
||||
f"{self.api_url}/api/auth/discord/session",
|
||||
headers=headers,
|
||||
json={
|
||||
"discord_id": self.discord_id,
|
||||
"display_name": self.display_name,
|
||||
},
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return {"error": "API unavailable"}, 503
|
||||
data = _response_json(response)
|
||||
if response.status_code == 200:
|
||||
self.token = data.get("token")
|
||||
self.user_uuid = data.get("user_uuid")
|
||||
self.timezone = data.get("timezone") or self.timezone
|
||||
return data, response.status_code
|
||||
|
||||
async def request(self, method, endpoint, data=None, params=None):
|
||||
if not self.token:
|
||||
_, status = await self.authenticate()
|
||||
if status != 200:
|
||||
return {"error": "authentication failed"}, status
|
||||
|
||||
try:
|
||||
response = await self._request(
|
||||
method,
|
||||
endpoint,
|
||||
self.token,
|
||||
data=data,
|
||||
params=params,
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return {"error": "API unavailable"}, 503
|
||||
if response.status_code == 401:
|
||||
authResult, status = await self.authenticate()
|
||||
if status != 200:
|
||||
return authResult, status
|
||||
try:
|
||||
response = await self._request(
|
||||
method,
|
||||
endpoint,
|
||||
self.token,
|
||||
data=data,
|
||||
params=params,
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return {"error": "API unavailable"}, 503
|
||||
return _response_json(response), response.status_code
|
||||
|
||||
async def service_request(self, method, endpoint, data=None):
|
||||
try:
|
||||
response = await self._request(
|
||||
method, endpoint, self.service_key, data=data
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return {"error": "API unavailable"}, 503
|
||||
return _response_json(response), response.status_code
|
||||
|
||||
async def _request(self, method, endpoint, token, data=None, params=None):
|
||||
requestOptions = {
|
||||
"headers": {"Authorization": f"Bearer {token}"},
|
||||
"params": params,
|
||||
}
|
||||
if data is not None:
|
||||
requestOptions["json"] = data
|
||||
return await self.http.request(
|
||||
method.upper(), f"{self.api_url}{endpoint}", **requestOptions
|
||||
)
|
||||
|
||||
async def close(self):
|
||||
await self.http.aclose()
|
||||
|
||||
|
||||
def _response_json(response):
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return {}
|
||||
200
bot/bot.py
Normal file
200
bot/bot.py
Normal file
@@ -0,0 +1,200 @@
|
||||
"""
|
||||
bot.py - Discord DM adapter for the reusable framework
|
||||
|
||||
Discord authenticates to the API with a service key and exchanges each
|
||||
stable Discord identity for a short-lived user session. Feature modules see
|
||||
only CommandContext rather than discord.py internals.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import socket
|
||||
import uuid
|
||||
|
||||
import discord
|
||||
from discord.ext import tasks
|
||||
from dotenv import load_dotenv
|
||||
|
||||
import ai.parser as ai_parser
|
||||
from bot.api_client import ApiClient
|
||||
from bot.context import CommandContext
|
||||
from core.registry import discover_modules
|
||||
|
||||
|
||||
load_dotenv(Path(__file__).resolve().parents[1] / ".env", override=False)
|
||||
|
||||
logging.basicConfig(level=os.getenv("LOG_LEVEL", "INFO"))
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DISCORD_BOT_TOKEN = os.getenv("DISCORD_BOT_TOKEN")
|
||||
OUTBOX_POLL_INTERVAL = float(os.getenv("OUTBOX_POLL_INTERVAL", 5))
|
||||
OUTBOX_BATCH_SIZE = int(os.getenv("OUTBOX_BATCH_SIZE", 20))
|
||||
|
||||
module_registry = discover_modules()
|
||||
api_clients = {}
|
||||
message_history = {}
|
||||
user_locks = {}
|
||||
|
||||
intents = discord.Intents.default()
|
||||
intents.message_content = True
|
||||
client = discord.Client(intents=intents)
|
||||
service_client = ApiClient("service", "Discord bot")
|
||||
OUTBOX_WORKER_ID = (
|
||||
f"discord:{socket.gethostname()}:{os.getpid()}:{uuid.uuid4().hex}"
|
||||
)
|
||||
|
||||
|
||||
def getApiClient(message):
|
||||
discord_id = message.author.id
|
||||
if discord_id not in api_clients:
|
||||
display_name = getattr(message.author, "display_name", None) or str(
|
||||
message.author
|
||||
)
|
||||
api_clients[discord_id] = ApiClient(discord_id, display_name)
|
||||
return api_clients[discord_id]
|
||||
|
||||
|
||||
async def sendHelpMessage(message):
|
||||
lines = module_registry.help_lines()
|
||||
content = "\n".join(lines) if lines else "- No modules registered"
|
||||
await message.channel.send(
|
||||
f"**Available commands:**\n{content}\n\nJust talk naturally and I'll help."
|
||||
)
|
||||
|
||||
|
||||
async def authenticateMessage(message):
|
||||
api_client = getApiClient(message)
|
||||
if api_client.token:
|
||||
return api_client
|
||||
|
||||
result, status = await api_client.authenticate()
|
||||
if status == 200:
|
||||
return api_client
|
||||
if status == 403:
|
||||
await message.channel.send("This bot is not enabled for your Discord account.")
|
||||
elif status == 503:
|
||||
await message.channel.send("The bot is still starting. Please try again shortly.")
|
||||
else:
|
||||
detail = result.get("error", "authentication failed")
|
||||
logger.warning("Discord authentication failed: %s", detail)
|
||||
await message.channel.send("I couldn't start your session. Please try again.")
|
||||
return None
|
||||
|
||||
|
||||
async def routeCommand(message, api_client):
|
||||
normalized = message.content.strip().lower()
|
||||
if normalized in {"help", "?", "what can i say"}:
|
||||
await sendHelpMessage(message)
|
||||
return
|
||||
|
||||
discord_id = message.author.id
|
||||
async with message.channel.typing():
|
||||
history = message_history.get(discord_id, [])
|
||||
parsed = await ai_parser.parse_command_async(
|
||||
message.content,
|
||||
module_registry,
|
||||
history=history,
|
||||
timezone_name=api_client.timezone,
|
||||
)
|
||||
message_history.setdefault(discord_id, []).append((message.content, parsed))
|
||||
message_history[discord_id] = message_history[discord_id][-5:]
|
||||
|
||||
if parsed.get("needs_clarification"):
|
||||
await message.channel.send(parsed["needs_clarification"])
|
||||
return
|
||||
if parsed.get("error"):
|
||||
await message.channel.send(f"I had trouble understanding that: {parsed['error']}")
|
||||
return
|
||||
|
||||
interaction_type = parsed.get("interaction_type")
|
||||
command = module_registry.get_command(interaction_type)
|
||||
if not command:
|
||||
await message.channel.send(f"Unknown command type: {interaction_type}")
|
||||
return
|
||||
|
||||
context = CommandContext(message, api_client)
|
||||
try:
|
||||
await command["handler"](context, parsed)
|
||||
except Exception:
|
||||
logger.exception("Command handler failed: %s", interaction_type)
|
||||
await message.channel.send("That command failed unexpectedly. Please try again.")
|
||||
|
||||
|
||||
@client.event
|
||||
async def on_ready():
|
||||
logger.info("Bot logged in as %s", client.user)
|
||||
if not outboxLoop.is_running():
|
||||
outboxLoop.start()
|
||||
|
||||
|
||||
@client.event
|
||||
async def on_message(message):
|
||||
if message.author == client.user:
|
||||
return
|
||||
if not isinstance(message.channel, discord.DMChannel):
|
||||
return
|
||||
|
||||
user_lock = user_locks.setdefault(message.author.id, asyncio.Lock())
|
||||
async with user_lock:
|
||||
api_client = await authenticateMessage(message)
|
||||
if api_client:
|
||||
await routeCommand(message, api_client)
|
||||
|
||||
|
||||
@tasks.loop(seconds=OUTBOX_POLL_INTERVAL)
|
||||
async def outboxLoop():
|
||||
result, status = await service_client.service_request(
|
||||
"post",
|
||||
"/api/internal/outbox/claim",
|
||||
{
|
||||
"channel": "discord_dm",
|
||||
"worker_id": OUTBOX_WORKER_ID,
|
||||
"limit": OUTBOX_BATCH_SIZE,
|
||||
},
|
||||
)
|
||||
if status != 200:
|
||||
logger.warning("Outbox claim failed with status %s", status)
|
||||
return
|
||||
|
||||
for outbound in result.get("messages", []):
|
||||
await deliverOutboundMessage(outbound)
|
||||
|
||||
|
||||
async def deliverOutboundMessage(outbound):
|
||||
message_id = outbound["id"]
|
||||
try:
|
||||
user = await client.fetch_user(int(outbound["provider_user_id"]))
|
||||
sent = await user.send(outbound["content"])
|
||||
payload = {
|
||||
"status": "sent",
|
||||
"external_message_id": str(sent.id),
|
||||
"worker_id": outbound["worker_id"],
|
||||
}
|
||||
except Exception as error:
|
||||
logger.warning("Discord delivery failed for %s: %s", message_id, error)
|
||||
payload = {
|
||||
"status": "retry",
|
||||
"error": str(error)[:500],
|
||||
"worker_id": outbound["worker_id"],
|
||||
}
|
||||
|
||||
_, status = await service_client.service_request(
|
||||
"post", f"/api/internal/outbox/{message_id}/result", payload
|
||||
)
|
||||
if status != 200:
|
||||
logger.warning("Outbox result failed for %s with status %s", message_id, status)
|
||||
|
||||
|
||||
@outboxLoop.before_loop
|
||||
async def beforeOutboxLoop():
|
||||
await client.wait_until_ready()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if not DISCORD_BOT_TOKEN:
|
||||
raise RuntimeError("DISCORD_BOT_TOKEN is required")
|
||||
if not os.getenv("BOT_API_KEY"):
|
||||
raise RuntimeError("BOT_API_KEY is required")
|
||||
client.run(DISCORD_BOT_TOKEN)
|
||||
13
bot/context.py
Normal file
13
bot/context.py
Normal file
@@ -0,0 +1,13 @@
|
||||
"""Platform-neutral command context exposed to feature handlers."""
|
||||
|
||||
|
||||
class CommandContext:
|
||||
def __init__(self, message, api_client):
|
||||
self._message = message
|
||||
self.api = api_client
|
||||
self.user_uuid = api_client.user_uuid
|
||||
self.discord_user_id = str(message.author.id)
|
||||
self.timezone = api_client.timezone
|
||||
|
||||
async def reply(self, content):
|
||||
return await self._message.channel.send(content)
|
||||
71
config/migrations/0001_core_tables.sql
Normal file
71
config/migrations/0001_core_tables.sql
Normal file
@@ -0,0 +1,71 @@
|
||||
-- Baseline legacy tables and normalize their timestamps as UTC.
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id UUID PRIMARY KEY,
|
||||
username VARCHAR(255) UNIQUE NOT NULL,
|
||||
password_hashed BYTEA NOT NULL,
|
||||
timezone VARCHAR(64) NOT NULL DEFAULT 'UTC',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
|
||||
ALTER TABLE users
|
||||
ADD COLUMN IF NOT EXISTS timezone VARCHAR(64) NOT NULL DEFAULT 'UTC';
|
||||
|
||||
CREATE TABLE IF NOT EXISTS notifications (
|
||||
id UUID PRIMARY KEY,
|
||||
user_uuid UUID REFERENCES users(id) ON DELETE CASCADE UNIQUE,
|
||||
discord_webhook VARCHAR(500),
|
||||
discord_enabled BOOLEAN DEFAULT FALSE,
|
||||
ntfy_topic VARCHAR(255),
|
||||
ntfy_enabled BOOLEAN DEFAULT FALSE,
|
||||
last_message_sent TIMESTAMPTZ,
|
||||
current_notification_status VARCHAR(50) DEFAULT 'inactive',
|
||||
created_at TIMESTAMPTZ DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ DEFAULT NOW()
|
||||
);
|
||||
|
||||
DO $$
|
||||
BEGIN
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = 'public'
|
||||
AND table_name = 'users'
|
||||
AND column_name = 'created_at'
|
||||
AND data_type = 'timestamp without time zone'
|
||||
) THEN
|
||||
ALTER TABLE users ALTER COLUMN created_at TYPE TIMESTAMPTZ
|
||||
USING created_at AT TIME ZONE 'UTC';
|
||||
END IF;
|
||||
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = 'public'
|
||||
AND table_name = 'notifications'
|
||||
AND column_name = 'last_message_sent'
|
||||
AND data_type = 'timestamp without time zone'
|
||||
) THEN
|
||||
ALTER TABLE notifications ALTER COLUMN last_message_sent TYPE TIMESTAMPTZ
|
||||
USING last_message_sent AT TIME ZONE 'UTC';
|
||||
END IF;
|
||||
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = 'public'
|
||||
AND table_name = 'notifications'
|
||||
AND column_name = 'created_at'
|
||||
AND data_type = 'timestamp without time zone'
|
||||
) THEN
|
||||
ALTER TABLE notifications ALTER COLUMN created_at TYPE TIMESTAMPTZ
|
||||
USING created_at AT TIME ZONE 'UTC';
|
||||
END IF;
|
||||
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = 'public'
|
||||
AND table_name = 'notifications'
|
||||
AND column_name = 'updated_at'
|
||||
AND data_type = 'timestamp without time zone'
|
||||
) THEN
|
||||
ALTER TABLE notifications ALTER COLUMN updated_at TYPE TIMESTAMPTZ
|
||||
USING updated_at AT TIME ZONE 'UTC';
|
||||
END IF;
|
||||
END $$;
|
||||
35
config/migrations/0002_scheduled_jobs.sql
Normal file
35
config/migrations/0002_scheduled_jobs.sql
Normal file
@@ -0,0 +1,35 @@
|
||||
CREATE TABLE IF NOT EXISTS scheduled_jobs (
|
||||
id UUID PRIMARY KEY,
|
||||
user_uuid UUID REFERENCES users(id) ON DELETE CASCADE,
|
||||
job_type VARCHAR(100) NOT NULL,
|
||||
payload JSONB NOT NULL DEFAULT '{}'::JSONB,
|
||||
run_at TIMESTAMPTZ NOT NULL,
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'pending',
|
||||
attempts INTEGER NOT NULL DEFAULT 0,
|
||||
max_attempts INTEGER NOT NULL DEFAULT 3,
|
||||
leased_by VARCHAR(255),
|
||||
lease_until TIMESTAMPTZ,
|
||||
last_error TEXT,
|
||||
idempotency_key VARCHAR(255),
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
completed_at TIMESTAMPTZ,
|
||||
cancelled_at TIMESTAMPTZ,
|
||||
CONSTRAINT scheduled_jobs_status_check
|
||||
CHECK (status IN ('pending', 'running', 'completed', 'failed', 'cancelled')),
|
||||
CONSTRAINT scheduled_jobs_attempts_check
|
||||
CHECK (attempts >= 0 AND max_attempts > 0),
|
||||
CONSTRAINT scheduled_jobs_idempotency_unique
|
||||
UNIQUE (job_type, idempotency_key)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS scheduled_jobs_claim_idx
|
||||
ON scheduled_jobs (run_at, created_at)
|
||||
WHERE status = 'pending';
|
||||
|
||||
CREATE INDEX IF NOT EXISTS scheduled_jobs_expired_lease_idx
|
||||
ON scheduled_jobs (lease_until)
|
||||
WHERE status = 'running';
|
||||
|
||||
CREATE INDEX IF NOT EXISTS scheduled_jobs_user_idx
|
||||
ON scheduled_jobs (user_uuid, created_at DESC);
|
||||
33
config/migrations/0003_outbound_messages.sql
Normal file
33
config/migrations/0003_outbound_messages.sql
Normal file
@@ -0,0 +1,33 @@
|
||||
CREATE TABLE IF NOT EXISTS outbound_messages (
|
||||
id UUID PRIMARY KEY,
|
||||
user_uuid UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
channel VARCHAR(50) NOT NULL,
|
||||
payload JSONB NOT NULL,
|
||||
available_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'pending',
|
||||
attempts INTEGER NOT NULL DEFAULT 0,
|
||||
max_attempts INTEGER NOT NULL DEFAULT 3,
|
||||
leased_by VARCHAR(255),
|
||||
lease_until TIMESTAMPTZ,
|
||||
last_error TEXT,
|
||||
idempotency_key VARCHAR(255) NOT NULL UNIQUE,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
delivered_at TIMESTAMPTZ,
|
||||
cancelled_at TIMESTAMPTZ,
|
||||
CONSTRAINT outbound_messages_status_check
|
||||
CHECK (status IN ('pending', 'delivering', 'delivered', 'failed', 'cancelled')),
|
||||
CONSTRAINT outbound_messages_attempts_check
|
||||
CHECK (attempts >= 0 AND max_attempts > 0)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS outbound_messages_claim_idx
|
||||
ON outbound_messages (channel, available_at, created_at)
|
||||
WHERE status = 'pending';
|
||||
|
||||
CREATE INDEX IF NOT EXISTS outbound_messages_expired_lease_idx
|
||||
ON outbound_messages (lease_until)
|
||||
WHERE status = 'delivering';
|
||||
|
||||
CREATE INDEX IF NOT EXISTS outbound_messages_user_idx
|
||||
ON outbound_messages (user_uuid, created_at DESC);
|
||||
71
config/migrations/0004_auth_identity.sql
Normal file
71
config/migrations/0004_auth_identity.sql
Normal file
@@ -0,0 +1,71 @@
|
||||
-- Provider identities allow passwordless platform users while legacy credentials remain valid.
|
||||
ALTER TABLE users
|
||||
ALTER COLUMN username DROP NOT NULL,
|
||||
ALTER COLUMN password_hashed DROP NOT NULL;
|
||||
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (
|
||||
SELECT 1 FROM pg_constraint WHERE conname = 'users_credentials_pair_check'
|
||||
) THEN
|
||||
ALTER TABLE users ADD CONSTRAINT users_credentials_pair_check
|
||||
CHECK (
|
||||
(username IS NULL AND password_hashed IS NULL)
|
||||
OR (username IS NOT NULL AND password_hashed IS NOT NULL)
|
||||
);
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
CREATE TABLE IF NOT EXISTS provider_identities (
|
||||
id UUID PRIMARY KEY,
|
||||
user_uuid UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
provider VARCHAR(50) NOT NULL,
|
||||
provider_user_id VARCHAR(255) NOT NULL,
|
||||
display_name VARCHAR(255),
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT provider_identities_provider_user_unique
|
||||
UNIQUE (provider, provider_user_id),
|
||||
CONSTRAINT provider_identities_user_provider_unique
|
||||
UNIQUE (user_uuid, provider)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS provider_identities_user_idx
|
||||
ON provider_identities (user_uuid);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS api_keys (
|
||||
id UUID PRIMARY KEY,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
key_type VARCHAR(20) NOT NULL,
|
||||
user_uuid UUID REFERENCES users(id) ON DELETE CASCADE,
|
||||
service_name VARCHAR(255),
|
||||
key_prefix VARCHAR(20) NOT NULL,
|
||||
key_hash CHAR(64) NOT NULL UNIQUE,
|
||||
scopes JSONB NOT NULL DEFAULT '[]'::JSONB,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
expires_at TIMESTAMPTZ,
|
||||
revoked_at TIMESTAMPTZ,
|
||||
last_used_at TIMESTAMPTZ,
|
||||
CONSTRAINT api_keys_type_check
|
||||
CHECK (key_type IN ('user', 'service')),
|
||||
CONSTRAINT api_keys_scopes_check
|
||||
CHECK (jsonb_typeof(scopes) = 'array'),
|
||||
CONSTRAINT api_keys_owner_check
|
||||
CHECK (
|
||||
(key_type = 'user' AND user_uuid IS NOT NULL
|
||||
AND service_name IS NULL AND scopes = '[]'::JSONB)
|
||||
OR
|
||||
(key_type = 'service' AND user_uuid IS NULL
|
||||
AND service_name IS NOT NULL)
|
||||
)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS api_keys_prefix_idx
|
||||
ON api_keys (key_prefix);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS api_keys_user_idx
|
||||
ON api_keys (user_uuid, created_at DESC)
|
||||
WHERE key_type = 'user';
|
||||
|
||||
CREATE INDEX IF NOT EXISTS api_keys_service_idx
|
||||
ON api_keys (service_name, created_at DESC)
|
||||
WHERE key_type = 'service';
|
||||
2
config/migrations/0005_outbox_delivery_metadata.sql
Normal file
2
config/migrations/0005_outbox_delivery_metadata.sql
Normal file
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE outbound_messages
|
||||
ADD COLUMN IF NOT EXISTS external_message_id VARCHAR(255);
|
||||
2
config/migrations/0006_api_key_prefix.sql
Normal file
2
config/migrations/0006_api_key_prefix.sql
Normal file
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE api_keys
|
||||
ALTER COLUMN key_prefix TYPE VARCHAR(32);
|
||||
1
core/__init__.py
Normal file
1
core/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Framework services and persistence helpers."""
|
||||
351
core/api_keys.py
Normal file
351
core/api_keys.py
Normal file
@@ -0,0 +1,351 @@
|
||||
import datetime
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import uuid
|
||||
|
||||
import core.postgres as postgres
|
||||
import core.users as users
|
||||
|
||||
|
||||
USER_KEY_TYPE = "user"
|
||||
SERVICE_KEY_TYPE = "service"
|
||||
DEFAULT_SERVICE_SCOPES = (
|
||||
"discord:session",
|
||||
"outbox:claim",
|
||||
"outbox:deliver",
|
||||
)
|
||||
_TOKEN_PREFIX_LENGTH = 20
|
||||
_MIN_TOKEN_LENGTH = 32
|
||||
|
||||
|
||||
def _utcNow():
|
||||
return datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
|
||||
def _hashToken(token):
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _getPrefix(token):
|
||||
return token[:_TOKEN_PREFIX_LENGTH]
|
||||
|
||||
|
||||
def _normalizeScopes(scopes):
|
||||
if scopes is None:
|
||||
return []
|
||||
if isinstance(scopes, str):
|
||||
scopes = re.split(r"[\s,]+", scopes)
|
||||
if not isinstance(scopes, (list, tuple, set)):
|
||||
raise ValueError("scopes must be a list or comma-separated string")
|
||||
|
||||
normalized = []
|
||||
for scope in scopes:
|
||||
if not isinstance(scope, str) or not scope.strip():
|
||||
continue
|
||||
scope = scope.strip()
|
||||
if len(scope) > 100:
|
||||
raise ValueError("service scopes must be at most 100 characters")
|
||||
if scope not in normalized:
|
||||
normalized.append(scope)
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalizeExpiry(expiresAt):
|
||||
if expiresAt is None:
|
||||
return None
|
||||
if isinstance(expiresAt, str):
|
||||
try:
|
||||
expiresAt = datetime.datetime.fromisoformat(expiresAt.replace("Z", "+00:00"))
|
||||
except ValueError as error:
|
||||
raise ValueError("expires_at must be an ISO-8601 datetime") from error
|
||||
if not isinstance(expiresAt, datetime.datetime):
|
||||
raise ValueError("expires_at must be a datetime")
|
||||
if expiresAt.tzinfo is None:
|
||||
expiresAt = expiresAt.replace(tzinfo=datetime.timezone.utc)
|
||||
expiresAt = expiresAt.astimezone(datetime.timezone.utc)
|
||||
if expiresAt <= _utcNow():
|
||||
raise ValueError("expires_at must be in the future")
|
||||
return expiresAt
|
||||
|
||||
|
||||
def _validateSecret(token):
|
||||
if not isinstance(token, str) or len(token) < _MIN_TOKEN_LENGTH:
|
||||
raise ValueError(f"API keys must be at least {_MIN_TOKEN_LENGTH} characters")
|
||||
return token
|
||||
|
||||
|
||||
def _keyID(value):
|
||||
try:
|
||||
return str(uuid.UUID(str(value)))
|
||||
except (TypeError, ValueError, AttributeError):
|
||||
return None
|
||||
|
||||
|
||||
def _publicKey(record, includeSecret=None):
|
||||
if not record:
|
||||
return None
|
||||
public = dict(record)
|
||||
public.pop("key_hash", None)
|
||||
scopes = public.get("scopes")
|
||||
if isinstance(scopes, str):
|
||||
public["scopes"] = json.loads(scopes)
|
||||
if includeSecret is not None:
|
||||
public["key"] = includeSecret
|
||||
return public
|
||||
|
||||
|
||||
def createApiKey(
|
||||
name,
|
||||
keyType,
|
||||
userUUID=None,
|
||||
serviceName=None,
|
||||
scopes=None,
|
||||
expiresAt=None,
|
||||
secret=None,
|
||||
):
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("API key name is required")
|
||||
name = name.strip()
|
||||
if len(name) > 255:
|
||||
raise ValueError("API key name must be at most 255 characters")
|
||||
if keyType not in {USER_KEY_TYPE, SERVICE_KEY_TYPE}:
|
||||
raise ValueError("key_type must be user or service")
|
||||
|
||||
normalizedScopes = _normalizeScopes(scopes)
|
||||
if keyType == USER_KEY_TYPE:
|
||||
if not userUUID or not users.doesUserUUIDExist(userUUID):
|
||||
raise ValueError("user does not exist")
|
||||
if serviceName is not None:
|
||||
raise ValueError("user API keys cannot have a service name")
|
||||
if normalizedScopes:
|
||||
raise ValueError("user API keys cannot have service scopes")
|
||||
else:
|
||||
if userUUID is not None:
|
||||
raise ValueError("service API keys cannot have a user")
|
||||
if not isinstance(serviceName, str) or not serviceName.strip():
|
||||
raise ValueError("service name is required")
|
||||
serviceName = serviceName.strip()
|
||||
if len(serviceName) > 255:
|
||||
raise ValueError("service name must be at most 255 characters")
|
||||
|
||||
expiresAt = _normalizeExpiry(expiresAt)
|
||||
if secret is None:
|
||||
secret = f"llmbot_{keyType}_{secrets.token_urlsafe(32)}"
|
||||
secret = _validateSecret(secret)
|
||||
|
||||
keyData = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"name": name,
|
||||
"key_type": keyType,
|
||||
"user_uuid": str(userUUID) if userUUID else None,
|
||||
"service_name": serviceName,
|
||||
"key_prefix": _getPrefix(secret),
|
||||
"key_hash": _hashToken(secret),
|
||||
"scopes": json.dumps(normalizedScopes),
|
||||
"expires_at": expiresAt,
|
||||
}
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO api_keys (
|
||||
id, name, key_type, user_uuid, service_name,
|
||||
key_prefix, key_hash, scopes, expires_at
|
||||
) VALUES (
|
||||
%(id)s, %(name)s, %(key_type)s, %(user_uuid)s, %(service_name)s,
|
||||
%(key_prefix)s, %(key_hash)s, %(scopes)s::jsonb, %(expires_at)s
|
||||
)
|
||||
RETURNING *
|
||||
""",
|
||||
keyData,
|
||||
)
|
||||
record = dict(cursor.fetchone())
|
||||
return _publicKey(record, includeSecret=secret)
|
||||
|
||||
|
||||
def createUserApiKey(userUUID, name, expiresAt=None):
|
||||
return createApiKey(
|
||||
name,
|
||||
USER_KEY_TYPE,
|
||||
userUUID=userUUID,
|
||||
expiresAt=expiresAt,
|
||||
)
|
||||
|
||||
|
||||
def createServiceApiKey(serviceName, name, scopes, expiresAt=None):
|
||||
return createApiKey(
|
||||
name,
|
||||
SERVICE_KEY_TYPE,
|
||||
serviceName=serviceName,
|
||||
scopes=scopes,
|
||||
expiresAt=expiresAt,
|
||||
)
|
||||
|
||||
|
||||
def getApiKey(keyID):
|
||||
keyID = _keyID(keyID)
|
||||
if not keyID:
|
||||
return None
|
||||
return _publicKey(postgres.select_one("api_keys", {"id": keyID}))
|
||||
|
||||
|
||||
def listApiKeys(userUUID=None, serviceName=None, includeRevoked=False):
|
||||
if userUUID is not None and serviceName is not None:
|
||||
raise ValueError("filter by either user or service, not both")
|
||||
|
||||
clauses = []
|
||||
params = {}
|
||||
if userUUID is not None:
|
||||
clauses.append("user_uuid = %(user_uuid)s")
|
||||
params["user_uuid"] = str(userUUID)
|
||||
if serviceName is not None:
|
||||
clauses.append("service_name = %(service_name)s")
|
||||
params["service_name"] = serviceName
|
||||
if not includeRevoked:
|
||||
clauses.append("revoked_at IS NULL")
|
||||
|
||||
query = "SELECT * FROM api_keys"
|
||||
if clauses:
|
||||
query += " WHERE " + " AND ".join(clauses)
|
||||
query += " ORDER BY created_at DESC"
|
||||
return [_publicKey(record) for record in postgres.execute(query, params)]
|
||||
|
||||
|
||||
def listUserApiKeys(userUUID, includeRevoked=False):
|
||||
return listApiKeys(userUUID=userUUID, includeRevoked=includeRevoked)
|
||||
|
||||
|
||||
def revokeApiKey(keyID, userUUID=None, serviceName=None):
|
||||
keyID = _keyID(keyID)
|
||||
if not keyID:
|
||||
return False
|
||||
clauses = ["id = %(id)s", "revoked_at IS NULL"]
|
||||
params = {"id": keyID}
|
||||
if userUUID is not None:
|
||||
clauses.append("user_uuid = %(user_uuid)s")
|
||||
params["user_uuid"] = str(userUUID)
|
||||
if serviceName is not None:
|
||||
clauses.append("service_name = %(service_name)s")
|
||||
params["service_name"] = serviceName
|
||||
|
||||
records = postgres.execute(
|
||||
"UPDATE api_keys SET revoked_at = CURRENT_TIMESTAMP "
|
||||
f"WHERE {' AND '.join(clauses)} RETURNING id",
|
||||
params,
|
||||
)
|
||||
return bool(records)
|
||||
|
||||
|
||||
def revokeUserApiKey(userUUID, keyID):
|
||||
return revokeApiKey(keyID, userUUID=userUUID)
|
||||
|
||||
|
||||
def authenticateApiKey(token, requiredScopes=None):
|
||||
if not isinstance(token, str) or len(token) < _MIN_TOKEN_LENGTH:
|
||||
return None
|
||||
|
||||
records = postgres.select(
|
||||
"api_keys",
|
||||
{
|
||||
"key_prefix": _getPrefix(token),
|
||||
"key_hash": _hashToken(token),
|
||||
"revoked_at": None,
|
||||
},
|
||||
)
|
||||
if not records:
|
||||
return None
|
||||
|
||||
record = records[0]
|
||||
if not hmac.compare_digest(record["key_hash"], _hashToken(token)):
|
||||
return None
|
||||
expiresAt = record.get("expires_at")
|
||||
if expiresAt is not None:
|
||||
if expiresAt.tzinfo is None:
|
||||
expiresAt = expiresAt.replace(tzinfo=datetime.timezone.utc)
|
||||
if expiresAt <= _utcNow():
|
||||
return None
|
||||
|
||||
scopes = record.get("scopes") or []
|
||||
if isinstance(scopes, str):
|
||||
scopes = json.loads(scopes)
|
||||
requiredScopes = _normalizeScopes(requiredScopes)
|
||||
if requiredScopes and "*" not in scopes:
|
||||
if not set(requiredScopes).issubset(set(scopes)):
|
||||
return None
|
||||
|
||||
postgres.update(
|
||||
"api_keys",
|
||||
{"last_used_at": _utcNow()},
|
||||
{"id": record["id"]},
|
||||
)
|
||||
return {
|
||||
"type": record["key_type"],
|
||||
"authentication": "api_key",
|
||||
"user_uuid": record.get("user_uuid"),
|
||||
"service_name": record.get("service_name"),
|
||||
"api_key_id": record["id"],
|
||||
"scopes": scopes,
|
||||
"can_manage_api_keys": False,
|
||||
}
|
||||
|
||||
|
||||
def bootstrapServiceApiKey():
|
||||
secret = os.getenv("BOT_API_KEY")
|
||||
if not secret:
|
||||
return None
|
||||
_validateSecret(secret)
|
||||
|
||||
serviceName = "discord-bot"
|
||||
keyName = "BOT_API_KEY"
|
||||
scopes = _normalizeScopes(os.getenv("BOT_API_KEY_SCOPES"))
|
||||
if not scopes:
|
||||
scopes = list(DEFAULT_SERVICE_SCOPES)
|
||||
tokenHash = _hashToken(secret)
|
||||
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE api_keys
|
||||
SET revoked_at = CURRENT_TIMESTAMP
|
||||
WHERE key_type = 'service'
|
||||
AND service_name = %(service_name)s
|
||||
AND name = %(name)s
|
||||
AND key_hash != %(key_hash)s
|
||||
AND revoked_at IS NULL
|
||||
""",
|
||||
{
|
||||
"service_name": serviceName,
|
||||
"name": keyName,
|
||||
"key_hash": tokenHash,
|
||||
},
|
||||
)
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO api_keys (
|
||||
id, name, key_type, service_name,
|
||||
key_prefix, key_hash, scopes
|
||||
) VALUES (
|
||||
%(id)s, %(name)s, 'service', %(service_name)s,
|
||||
%(key_prefix)s, %(key_hash)s, %(scopes)s::jsonb
|
||||
)
|
||||
ON CONFLICT (key_hash) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
service_name = EXCLUDED.service_name,
|
||||
scopes = EXCLUDED.scopes,
|
||||
expires_at = NULL,
|
||||
revoked_at = NULL
|
||||
RETURNING *
|
||||
""",
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"name": keyName,
|
||||
"service_name": serviceName,
|
||||
"key_prefix": _getPrefix(secret),
|
||||
"key_hash": tokenHash,
|
||||
"scopes": json.dumps(scopes),
|
||||
},
|
||||
)
|
||||
return _publicKey(dict(cursor.fetchone()))
|
||||
180
core/auth.py
Normal file
180
core/auth.py
Normal file
@@ -0,0 +1,180 @@
|
||||
import datetime
|
||||
import os
|
||||
|
||||
import bcrypt
|
||||
import jwt
|
||||
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
|
||||
|
||||
import core.api_keys as apiKeys
|
||||
import core.postgres as postgres
|
||||
import core.users as users
|
||||
|
||||
|
||||
DEFAULT_TOKEN_LIFETIME = datetime.timedelta(hours=1)
|
||||
|
||||
|
||||
def decodeLoginToken(loginToken):
|
||||
secret = os.getenv("JWT_SECRET")
|
||||
if not secret or not isinstance(loginToken, str):
|
||||
return None
|
||||
try:
|
||||
payload = jwt.decode(loginToken, secret, algorithms=["HS256"])
|
||||
except (ExpiredSignatureError, InvalidTokenError):
|
||||
return None
|
||||
|
||||
userUUID = payload.get("sub")
|
||||
if not isinstance(userUUID, str) or not users.doesUserUUIDExist(userUUID):
|
||||
return None
|
||||
return payload
|
||||
|
||||
|
||||
def createLoginToken(userUUID, name=None, expiresIn=None, extraClaims=None):
|
||||
secret = os.getenv("JWT_SECRET")
|
||||
if not secret:
|
||||
raise RuntimeError("JWT_SECRET is required")
|
||||
if not userUUID or not users.doesUserUUIDExist(userUUID):
|
||||
raise ValueError("user does not exist")
|
||||
|
||||
now = datetime.datetime.now(datetime.timezone.utc)
|
||||
if expiresIn is None:
|
||||
expiresIn = DEFAULT_TOKEN_LIFETIME
|
||||
elif isinstance(expiresIn, (int, float)):
|
||||
expiresIn = datetime.timedelta(seconds=expiresIn)
|
||||
if not isinstance(expiresIn, datetime.timedelta) or expiresIn.total_seconds() <= 0:
|
||||
raise ValueError("token lifetime must be positive")
|
||||
|
||||
payload = {
|
||||
"sub": str(userUUID),
|
||||
"name": name if name is not None else users.getUserFirstName(userUUID),
|
||||
"iat": now,
|
||||
"exp": now + expiresIn,
|
||||
}
|
||||
if extraClaims is not None:
|
||||
if not isinstance(extraClaims, dict):
|
||||
raise ValueError("extra claims must be an object")
|
||||
protectedClaims = {"sub", "iat", "exp"}
|
||||
payload.update(
|
||||
{
|
||||
key: value
|
||||
for key, value in extraClaims.items()
|
||||
if key not in protectedClaims
|
||||
}
|
||||
)
|
||||
return jwt.encode(payload, secret, algorithm="HS256")
|
||||
|
||||
|
||||
def verifyLoginToken(login_token, username=False, userUUID=False):
|
||||
if username:
|
||||
userUUID = users.getUserUUID(username)
|
||||
|
||||
if userUUID:
|
||||
decodedToken = decodeLoginToken(login_token)
|
||||
if decodedToken and decodedToken.get("sub") == str(userUUID):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def getUserpasswordHash(userUUID):
|
||||
user = postgres.select_one("users", {"id": userUUID})
|
||||
if user:
|
||||
pw_hash = user.get("password_hashed")
|
||||
if isinstance(pw_hash, memoryview):
|
||||
return bytes(pw_hash)
|
||||
return pw_hash
|
||||
return None
|
||||
|
||||
|
||||
def getLoginToken(username, password):
|
||||
if not isinstance(username, str) or not isinstance(password, str):
|
||||
return False
|
||||
userUUID = users.getUserUUID(username)
|
||||
if userUUID:
|
||||
formattedPass = password.encode("utf-8")
|
||||
if not formattedPass or len(formattedPass) > users.MAX_PASSWORD_BYTES:
|
||||
return False
|
||||
usersHashedPassword = getUserpasswordHash(userUUID)
|
||||
if not usersHashedPassword:
|
||||
return False
|
||||
try:
|
||||
if bcrypt.checkpw(formattedPass, usersHashedPassword):
|
||||
return createLoginToken(userUUID)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def unregisterUser(userUUID, password):
|
||||
pw_hash = getUserpasswordHash(userUUID)
|
||||
if not pw_hash or not isinstance(password, str):
|
||||
return False
|
||||
formattedPassword = password.encode("utf-8")
|
||||
if not formattedPassword or len(formattedPassword) > users.MAX_PASSWORD_BYTES:
|
||||
return False
|
||||
try:
|
||||
if bcrypt.checkpw(formattedPassword, pw_hash):
|
||||
return users.deleteUser(userUUID)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def getBearerToken(authorization):
|
||||
if not isinstance(authorization, str):
|
||||
return None
|
||||
scheme, separator, token = authorization.strip().partition(" ")
|
||||
if not separator or scheme.lower() != "bearer" or not token.strip():
|
||||
return None
|
||||
return token.strip()
|
||||
|
||||
|
||||
def authenticateBearerToken(
|
||||
authorization,
|
||||
requiredScopes=None,
|
||||
allowUser=True,
|
||||
allowService=True,
|
||||
):
|
||||
token = getBearerToken(authorization)
|
||||
if not token:
|
||||
return None
|
||||
|
||||
if token.count(".") == 2:
|
||||
payload = decodeLoginToken(token)
|
||||
if payload:
|
||||
if not allowUser or requiredScopes:
|
||||
return None
|
||||
return {
|
||||
"type": "user",
|
||||
"authentication": "jwt",
|
||||
"user_uuid": payload["sub"],
|
||||
"service_name": None,
|
||||
"api_key_id": None,
|
||||
"scopes": [],
|
||||
"can_manage_api_keys": True,
|
||||
"claims": payload,
|
||||
}
|
||||
|
||||
principal = apiKeys.authenticateApiKey(token, requiredScopes=requiredScopes)
|
||||
if not principal:
|
||||
return None
|
||||
if principal["type"] == apiKeys.USER_KEY_TYPE and not allowUser:
|
||||
return None
|
||||
if principal["type"] == apiKeys.SERVICE_KEY_TYPE and not allowService:
|
||||
return None
|
||||
return principal
|
||||
|
||||
|
||||
def isUserPrincipal(principal, userUUID=None, requireLogin=False):
|
||||
if not principal or principal.get("type") != "user":
|
||||
return False
|
||||
if userUUID is not None and str(principal.get("user_uuid")) != str(userUUID):
|
||||
return False
|
||||
if requireLogin and principal.get("authentication") != "jwt":
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def hasServiceScope(principal, scope):
|
||||
if not principal or principal.get("type") != "service":
|
||||
return False
|
||||
scopes = principal.get("scopes") or []
|
||||
return "*" in scopes or scope in scopes
|
||||
224
core/identity.py
Normal file
224
core/identity.py
Normal file
@@ -0,0 +1,224 @@
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
|
||||
import core.postgres as postgres
|
||||
import core.users as users
|
||||
|
||||
|
||||
DISCORD_PROVIDER = "discord"
|
||||
DEFAULT_ENROLLMENT_MODE = "allowlist"
|
||||
_PROVIDER_PATTERN = re.compile(r"^[a-z0-9][a-z0-9_.-]{0,49}$")
|
||||
|
||||
|
||||
def _normalizeProvider(provider):
|
||||
provider = provider.strip().lower() if isinstance(provider, str) else ""
|
||||
if not _PROVIDER_PATTERN.fullmatch(provider):
|
||||
raise ValueError("provider must be a valid provider name")
|
||||
return provider
|
||||
|
||||
|
||||
def _normalizeProviderUserID(providerUserID):
|
||||
if providerUserID is None:
|
||||
raise ValueError("provider user ID is required")
|
||||
providerUserID = str(providerUserID).strip()
|
||||
if not providerUserID or len(providerUserID) > 255:
|
||||
raise ValueError("provider user ID must be between 1 and 255 characters")
|
||||
return providerUserID
|
||||
|
||||
|
||||
def _normalizeDisplayName(displayName):
|
||||
if displayName is None:
|
||||
return None
|
||||
if not isinstance(displayName, str):
|
||||
raise ValueError("display name must be a string")
|
||||
displayName = displayName.strip()
|
||||
if not displayName:
|
||||
return None
|
||||
if len(displayName) > 255:
|
||||
raise ValueError("display name must be at most 255 characters")
|
||||
return displayName
|
||||
|
||||
|
||||
def getDiscordEnrollmentMode():
|
||||
mode = os.getenv("DISCORD_ENROLLMENT_MODE", DEFAULT_ENROLLMENT_MODE)
|
||||
mode = mode.strip().lower()
|
||||
return "open" if mode == "open" else DEFAULT_ENROLLMENT_MODE
|
||||
|
||||
|
||||
def getDiscordAllowlist():
|
||||
configured = os.getenv("DISCORD_ALLOWLIST")
|
||||
if configured is None:
|
||||
configured = os.getenv("DISCORD_ALLOWED_USER_IDS", "")
|
||||
return {
|
||||
value
|
||||
for value in re.split(r"[\s,]+", configured)
|
||||
if value
|
||||
}
|
||||
|
||||
|
||||
def isDiscordEnrollmentAllowed(discordID):
|
||||
discordID = _normalizeProviderUserID(discordID)
|
||||
if getDiscordEnrollmentMode() == "open":
|
||||
return True
|
||||
return discordID in getDiscordAllowlist()
|
||||
|
||||
|
||||
def getProviderIdentity(provider, providerUserID):
|
||||
provider = _normalizeProvider(provider)
|
||||
providerUserID = _normalizeProviderUserID(providerUserID)
|
||||
return postgres.select_one(
|
||||
"provider_identities",
|
||||
{"provider": provider, "provider_user_id": providerUserID},
|
||||
)
|
||||
|
||||
|
||||
def getUserForIdentity(provider, providerUserID):
|
||||
identity = getProviderIdentity(provider, providerUserID)
|
||||
if not identity:
|
||||
return None
|
||||
return users.getUser(identity["user_uuid"])
|
||||
|
||||
|
||||
def listUserIdentities(userUUID):
|
||||
return postgres.select("provider_identities", {"user_uuid": userUUID})
|
||||
|
||||
|
||||
def linkProviderIdentity(userUUID, provider, providerUserID, displayName=None):
|
||||
if not users.doesUserUUIDExist(userUUID):
|
||||
raise ValueError("user does not exist")
|
||||
provider = _normalizeProvider(provider)
|
||||
providerUserID = _normalizeProviderUserID(providerUserID)
|
||||
displayName = _normalizeDisplayName(displayName)
|
||||
|
||||
existing = getProviderIdentity(provider, providerUserID)
|
||||
if existing:
|
||||
if str(existing["user_uuid"]) != str(userUUID):
|
||||
raise ValueError("provider identity is already linked to another user")
|
||||
if existing.get("display_name") != displayName:
|
||||
updated = postgres.update(
|
||||
"provider_identities",
|
||||
{"display_name": displayName},
|
||||
{"id": existing["id"]},
|
||||
)
|
||||
return updated[0]
|
||||
return existing
|
||||
|
||||
return postgres.insert(
|
||||
"provider_identities",
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_uuid": str(userUUID),
|
||||
"provider": provider,
|
||||
"provider_user_id": providerUserID,
|
||||
"display_name": displayName,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def unlinkProviderIdentity(userUUID, provider, providerUserID):
|
||||
provider = _normalizeProvider(provider)
|
||||
providerUserID = _normalizeProviderUserID(providerUserID)
|
||||
deleted = postgres.delete(
|
||||
"provider_identities",
|
||||
{
|
||||
"user_uuid": userUUID,
|
||||
"provider": provider,
|
||||
"provider_user_id": providerUserID,
|
||||
},
|
||||
)
|
||||
return bool(deleted)
|
||||
|
||||
|
||||
def linkDiscordIdentity(userUUID, discordID, displayName=None):
|
||||
return linkProviderIdentity(
|
||||
userUUID,
|
||||
DISCORD_PROVIDER,
|
||||
discordID,
|
||||
displayName=displayName,
|
||||
)
|
||||
|
||||
|
||||
def linkDiscordUser(discordID, userUUID=None, username=None, displayName=None):
|
||||
if bool(userUUID) == bool(username):
|
||||
raise ValueError("provide either user UUID or username")
|
||||
if username:
|
||||
userUUID = users.getUserUUID(username)
|
||||
if not userUUID or not users.doesUserUUIDExist(userUUID):
|
||||
raise ValueError("user does not exist")
|
||||
return linkDiscordIdentity(userUUID, discordID, displayName=displayName)
|
||||
|
||||
|
||||
def getDiscordUser(discordID):
|
||||
return getUserForIdentity(DISCORD_PROVIDER, discordID)
|
||||
|
||||
|
||||
def getOrCreateDiscordUser(discordID, displayName=None, timezoneName=None):
|
||||
discordID = _normalizeProviderUserID(discordID)
|
||||
displayName = _normalizeDisplayName(displayName)
|
||||
if timezoneName is None:
|
||||
timezoneName = users.getDefaultTimezone()
|
||||
else:
|
||||
timezoneName = users.normalizeTimezone(timezoneName)
|
||||
|
||||
existing = getProviderIdentity(DISCORD_PROVIDER, discordID)
|
||||
if existing:
|
||||
if displayName is not None and existing.get("display_name") != displayName:
|
||||
postgres.update(
|
||||
"provider_identities",
|
||||
{"display_name": displayName},
|
||||
{"id": existing["id"]},
|
||||
)
|
||||
return users.getUser(existing["user_uuid"])
|
||||
if not isDiscordEnrollmentAllowed(discordID):
|
||||
return None
|
||||
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"SELECT pg_advisory_xact_lock(hashtext(%(lock_key)s))",
|
||||
{"lock_key": f"discord:{discordID}"},
|
||||
)
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT users.*
|
||||
FROM provider_identities
|
||||
JOIN users ON users.id = provider_identities.user_uuid
|
||||
WHERE provider = %(provider)s AND provider_user_id = %(provider_user_id)s
|
||||
""",
|
||||
{
|
||||
"provider": DISCORD_PROVIDER,
|
||||
"provider_user_id": discordID,
|
||||
},
|
||||
)
|
||||
existingUser = cursor.fetchone()
|
||||
if existingUser:
|
||||
return dict(existingUser)
|
||||
|
||||
userUUID = str(uuid.uuid4())
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO users (id, username, password_hashed, timezone)
|
||||
VALUES (%(id)s, NULL, NULL, %(timezone)s)
|
||||
RETURNING *
|
||||
""",
|
||||
{"id": userUUID, "timezone": timezoneName},
|
||||
)
|
||||
userRecord = dict(cursor.fetchone())
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO provider_identities (
|
||||
id, user_uuid, provider, provider_user_id, display_name
|
||||
) VALUES (
|
||||
%(id)s, %(user_uuid)s, %(provider)s,
|
||||
%(provider_user_id)s, %(display_name)s
|
||||
)
|
||||
""",
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_uuid": userUUID,
|
||||
"provider": DISCORD_PROVIDER,
|
||||
"provider_user_id": discordID,
|
||||
"display_name": displayName,
|
||||
},
|
||||
)
|
||||
return userRecord
|
||||
343
core/jobs.py
Normal file
343
core/jobs.py
Normal file
@@ -0,0 +1,343 @@
|
||||
"""PostgreSQL-backed scheduled job operations."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timezone
|
||||
import uuid
|
||||
|
||||
from psycopg2.extras import Json
|
||||
|
||||
from core import postgres
|
||||
|
||||
|
||||
DEFAULT_MAX_ATTEMPTS = 3
|
||||
DEFAULT_LEASE_SECONDS = 300
|
||||
DEFAULT_RETRY_SECONDS = 30
|
||||
MAX_RETRY_SECONDS = 900
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _using_cursor(cursor=None):
|
||||
if cursor is not None:
|
||||
yield cursor
|
||||
return
|
||||
with postgres.get_cursor() as owned_cursor:
|
||||
yield owned_cursor
|
||||
|
||||
|
||||
def _timestamp(value, field="timestamp"):
|
||||
if isinstance(value, str):
|
||||
value = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
if not isinstance(value, datetime):
|
||||
raise ValueError(f"{field} must be a datetime or ISO-8601 string")
|
||||
if value.tzinfo is None:
|
||||
value = value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _positive(value, field, maximum=None):
|
||||
try:
|
||||
value = int(value)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise ValueError(f"{field} must be a whole number") from error
|
||||
if value < 1 or (maximum is not None and value > maximum):
|
||||
suffix = f" and at most {maximum}" if maximum is not None else ""
|
||||
raise ValueError(f"{field} must be at least 1{suffix}")
|
||||
return value
|
||||
|
||||
|
||||
def _row(cursor):
|
||||
record = cursor.fetchone()
|
||||
return dict(record) if record else None
|
||||
|
||||
|
||||
def create_job(
|
||||
job_type,
|
||||
payload,
|
||||
run_at,
|
||||
user_uuid=None,
|
||||
max_attempts=DEFAULT_MAX_ATTEMPTS,
|
||||
idempotency_key=None,
|
||||
job_id=None,
|
||||
cursor=None,
|
||||
):
|
||||
"""Create a job, returning the existing row for a repeated idempotency key."""
|
||||
if not isinstance(job_type, str) or not job_type.strip():
|
||||
raise ValueError("job_type is required")
|
||||
max_attempts = _positive(max_attempts, "max_attempts", 100)
|
||||
values = {
|
||||
"id": str(job_id or uuid.uuid4()),
|
||||
"user_uuid": user_uuid,
|
||||
"job_type": job_type.strip(),
|
||||
"payload": Json(payload if payload is not None else {}),
|
||||
"run_at": _timestamp(run_at, "run_at"),
|
||||
"max_attempts": max_attempts,
|
||||
"idempotency_key": idempotency_key,
|
||||
}
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
INSERT INTO scheduled_jobs (
|
||||
id, user_uuid, job_type, payload, run_at,
|
||||
max_attempts, idempotency_key
|
||||
) VALUES (
|
||||
%(id)s, %(user_uuid)s, %(job_type)s, %(payload)s, %(run_at)s,
|
||||
%(max_attempts)s, %(idempotency_key)s
|
||||
)
|
||||
ON CONFLICT (job_type, idempotency_key)
|
||||
DO UPDATE SET idempotency_key = EXCLUDED.idempotency_key
|
||||
WHERE scheduled_jobs.user_uuid IS NOT DISTINCT FROM EXCLUDED.user_uuid
|
||||
RETURNING *
|
||||
""",
|
||||
values,
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def get_job(job_id, cursor=None):
|
||||
"""Return one job by UUID."""
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute("SELECT * FROM scheduled_jobs WHERE id = %s", (job_id,))
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def list_jobs(user_uuid=None, status=None, job_type=None, limit=100, cursor=None):
|
||||
"""List jobs newest first, optionally filtered by owner, status, or type."""
|
||||
limit = _positive(limit, "limit", 500)
|
||||
clauses = []
|
||||
params = []
|
||||
if user_uuid is not None:
|
||||
clauses.append("user_uuid = %s")
|
||||
params.append(user_uuid)
|
||||
if status is not None:
|
||||
clauses.append("status = %s")
|
||||
params.append(status)
|
||||
if job_type is not None:
|
||||
clauses.append("job_type = %s")
|
||||
params.append(job_type)
|
||||
where = f"WHERE {' AND '.join(clauses)}" if clauses else ""
|
||||
params.append(limit)
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
f"""
|
||||
SELECT * FROM scheduled_jobs
|
||||
{where}
|
||||
ORDER BY created_at DESC
|
||||
LIMIT %s
|
||||
""",
|
||||
params,
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
|
||||
|
||||
def claim_due_jobs(
|
||||
worker_id,
|
||||
limit=10,
|
||||
lease_seconds=DEFAULT_LEASE_SECONDS,
|
||||
job_types=None,
|
||||
cursor=None,
|
||||
):
|
||||
"""Atomically lease due jobs using row locks that skip other workers."""
|
||||
if not worker_id:
|
||||
raise ValueError("worker_id is required")
|
||||
limit = _positive(limit, "limit", 100)
|
||||
lease_seconds = _positive(lease_seconds, "lease_seconds")
|
||||
if isinstance(job_types, str):
|
||||
job_types = [job_types]
|
||||
elif job_types is not None:
|
||||
job_types = list(job_types)
|
||||
if not job_types:
|
||||
return []
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
WITH exhausted AS (
|
||||
SELECT id
|
||||
FROM scheduled_jobs
|
||||
WHERE status = 'running'
|
||||
AND (lease_until IS NULL OR lease_until <= NOW())
|
||||
AND attempts >= max_attempts
|
||||
FOR UPDATE SKIP LOCKED
|
||||
)
|
||||
UPDATE scheduled_jobs AS job
|
||||
SET status = 'failed',
|
||||
lease_until = NULL,
|
||||
leased_by = NULL,
|
||||
last_error = COALESCE(last_error, 'lease expired'),
|
||||
updated_at = NOW()
|
||||
FROM exhausted
|
||||
WHERE job.id = exhausted.id
|
||||
"""
|
||||
)
|
||||
active_cursor.execute(
|
||||
"""
|
||||
WITH candidates AS (
|
||||
SELECT id
|
||||
FROM scheduled_jobs
|
||||
WHERE run_at <= NOW()
|
||||
AND attempts < max_attempts
|
||||
AND (
|
||||
status = 'pending'
|
||||
OR (
|
||||
status = 'running'
|
||||
AND (lease_until IS NULL OR lease_until <= NOW())
|
||||
)
|
||||
)
|
||||
AND (%(job_types)s IS NULL OR job_type = ANY(%(job_types)s))
|
||||
ORDER BY run_at, created_at
|
||||
FOR UPDATE SKIP LOCKED
|
||||
LIMIT %(limit)s
|
||||
)
|
||||
UPDATE scheduled_jobs AS job
|
||||
SET status = 'running',
|
||||
attempts = job.attempts + 1,
|
||||
leased_by = %(worker_id)s,
|
||||
lease_until = NOW() + (%(lease_seconds)s * INTERVAL '1 second'),
|
||||
updated_at = NOW()
|
||||
FROM candidates
|
||||
WHERE job.id = candidates.id
|
||||
RETURNING job.*
|
||||
""",
|
||||
{
|
||||
"job_types": job_types,
|
||||
"limit": limit,
|
||||
"worker_id": worker_id,
|
||||
"lease_seconds": lease_seconds,
|
||||
},
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
|
||||
|
||||
def renew_job_lease(job_id, worker_id, lease_seconds=DEFAULT_LEASE_SECONDS, cursor=None):
|
||||
"""Extend a lease only while it is owned by the requesting worker."""
|
||||
lease_seconds = _positive(lease_seconds, "lease_seconds")
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE scheduled_jobs
|
||||
SET lease_until = NOW() + (%s * INTERVAL '1 second'), updated_at = NOW()
|
||||
WHERE id = %s AND status = 'running' AND leased_by = %s
|
||||
RETURNING *
|
||||
""",
|
||||
(lease_seconds, job_id, worker_id),
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def complete_job(job_id, worker_id, cursor=None):
|
||||
"""Mark a job complete when its lease is still owned by the worker."""
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE scheduled_jobs
|
||||
SET status = 'completed', completed_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL, last_error = NULL
|
||||
WHERE id = %s AND status = 'running' AND leased_by = %s
|
||||
RETURNING *
|
||||
""",
|
||||
(job_id, worker_id),
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def fail_job(
|
||||
job_id,
|
||||
worker_id,
|
||||
error,
|
||||
retry_seconds=DEFAULT_RETRY_SECONDS,
|
||||
max_retry_seconds=MAX_RETRY_SECONDS,
|
||||
cursor=None,
|
||||
):
|
||||
"""Fail or reschedule an owned job using bounded exponential backoff."""
|
||||
retry_seconds = _positive(retry_seconds, "retry_seconds")
|
||||
max_retry_seconds = _positive(max_retry_seconds, "max_retry_seconds")
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
SELECT attempts, max_attempts
|
||||
FROM scheduled_jobs
|
||||
WHERE id = %s AND status = 'running' AND leased_by = %s
|
||||
FOR UPDATE
|
||||
""",
|
||||
(job_id, worker_id),
|
||||
)
|
||||
current = active_cursor.fetchone()
|
||||
if not current:
|
||||
return None
|
||||
|
||||
exhausted = current["attempts"] >= current["max_attempts"]
|
||||
delay = min(
|
||||
max_retry_seconds,
|
||||
retry_seconds * (2 ** min(max(current["attempts"] - 1, 0), 30)),
|
||||
)
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE scheduled_jobs
|
||||
SET status = %(status)s,
|
||||
run_at = CASE
|
||||
WHEN %(exhausted)s THEN run_at
|
||||
ELSE NOW() + (%(delay)s * INTERVAL '1 second')
|
||||
END,
|
||||
leased_by = NULL,
|
||||
lease_until = NULL,
|
||||
last_error = %(error)s,
|
||||
updated_at = NOW()
|
||||
WHERE id = %(job_id)s
|
||||
RETURNING *
|
||||
""",
|
||||
{
|
||||
"status": "failed" if exhausted else "pending",
|
||||
"exhausted": exhausted,
|
||||
"delay": delay,
|
||||
"error": str(error)[:4000],
|
||||
"job_id": job_id,
|
||||
},
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def cancel_job(job_id, user_uuid=None, cursor=None):
|
||||
"""Cancel one unfinished job, optionally enforcing its owner."""
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE scheduled_jobs
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE id = %(job_id)s
|
||||
AND status IN ('pending', 'running')
|
||||
AND (%(user_uuid)s IS NULL OR user_uuid = %(user_uuid)s)
|
||||
RETURNING *
|
||||
""",
|
||||
{"job_id": job_id, "user_uuid": user_uuid},
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def cancel_jobs(
|
||||
user_uuid=None, job_type=None, idempotency_key=None, cursor=None
|
||||
):
|
||||
"""Cancel matching unfinished jobs; at least one filter is required."""
|
||||
filters = {
|
||||
"user_uuid": user_uuid,
|
||||
"job_type": job_type,
|
||||
"idempotency_key": idempotency_key,
|
||||
}
|
||||
clauses = [f"{name} = %({name})s" for name, value in filters.items() if value is not None]
|
||||
if not clauses:
|
||||
raise ValueError("at least one cancellation filter is required")
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
f"""
|
||||
UPDATE scheduled_jobs
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE status IN ('pending', 'running')
|
||||
AND {' AND '.join(clauses)}
|
||||
RETURNING *
|
||||
""",
|
||||
filters,
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
76
core/manage.py
Normal file
76
core/manage.py
Normal file
@@ -0,0 +1,76 @@
|
||||
"""Small administrative command-line tools for framework operators."""
|
||||
|
||||
import argparse
|
||||
import uuid
|
||||
|
||||
import psycopg2
|
||||
|
||||
from core import identity
|
||||
|
||||
|
||||
def _discordID(value):
|
||||
value = value.strip()
|
||||
if not value or not value.isdigit():
|
||||
raise argparse.ArgumentTypeError("Discord ID must contain only digits")
|
||||
return value
|
||||
|
||||
|
||||
def _userUUID(value):
|
||||
try:
|
||||
return str(uuid.UUID(value))
|
||||
except (AttributeError, TypeError, ValueError) as error:
|
||||
raise argparse.ArgumentTypeError("user UUID must be a valid UUID") from error
|
||||
|
||||
|
||||
def _username(value):
|
||||
value = value.strip()
|
||||
if not value:
|
||||
raise argparse.ArgumentTypeError("username cannot be empty")
|
||||
return value
|
||||
|
||||
|
||||
def _linkDiscord(args):
|
||||
try:
|
||||
linked = identity.linkDiscordUser(
|
||||
args.discord_id,
|
||||
userUUID=args.user_uuid,
|
||||
username=args.username,
|
||||
)
|
||||
except ValueError as error:
|
||||
args.command_parser.error(str(error))
|
||||
except psycopg2.Error:
|
||||
args.command_parser.exit(
|
||||
1,
|
||||
"link-discord failed: database operation failed\n",
|
||||
)
|
||||
|
||||
print(
|
||||
"linked Discord user "
|
||||
f"{linked['provider_user_id']} to user {linked['user_uuid']}"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def buildParser():
|
||||
parser = argparse.ArgumentParser(description="Manage framework data")
|
||||
commands = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
linkParser = commands.add_parser(
|
||||
"link-discord",
|
||||
help="link an existing user to a Discord account",
|
||||
)
|
||||
linkParser.add_argument("--discord-id", required=True, type=_discordID)
|
||||
userSelector = linkParser.add_mutually_exclusive_group(required=True)
|
||||
userSelector.add_argument("--user-uuid", type=_userUUID)
|
||||
userSelector.add_argument("--username", type=_username)
|
||||
linkParser.set_defaults(handler=_linkDiscord, command_parser=linkParser)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
args = buildParser().parse_args(argv)
|
||||
return args.handler(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
202
core/migrations/__init__.py
Normal file
202
core/migrations/__init__.py
Normal file
@@ -0,0 +1,202 @@
|
||||
"""Versioned PostgreSQL migrations for the framework and feature modules."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
import re
|
||||
|
||||
import psycopg2.extras
|
||||
|
||||
from core import postgres
|
||||
|
||||
|
||||
MIGRATION_NAME = re.compile(r"^(?P<version>\d+)(?:[_-].*)?\.sql$")
|
||||
LOCK_NAME = "llm-bot-framework:schema-migrations"
|
||||
|
||||
|
||||
class MigrationError(RuntimeError):
|
||||
"""Raised when migration history is inconsistent or cannot be applied."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Migration:
|
||||
namespace: str
|
||||
version: int
|
||||
name: str
|
||||
path: Path
|
||||
checksum: str
|
||||
sql: str
|
||||
|
||||
|
||||
def _project_root(project_root=None):
|
||||
if project_root is not None:
|
||||
return Path(project_root).resolve()
|
||||
return Path(__file__).resolve().parents[2]
|
||||
|
||||
|
||||
def _read_namespace(namespace, directory):
|
||||
migrations = []
|
||||
if not directory.is_dir():
|
||||
return migrations
|
||||
|
||||
for path in sorted(directory.glob("*.sql")):
|
||||
match = MIGRATION_NAME.match(path.name)
|
||||
if not match:
|
||||
raise MigrationError(f"Invalid migration filename: {path}")
|
||||
sql = path.read_text(encoding="utf-8")
|
||||
migrations.append(
|
||||
Migration(
|
||||
namespace=namespace,
|
||||
version=int(match.group("version")),
|
||||
name=path.name,
|
||||
path=path,
|
||||
checksum=sha256(sql.encode("utf-8")).hexdigest(),
|
||||
sql=sql,
|
||||
)
|
||||
)
|
||||
return migrations
|
||||
|
||||
|
||||
def discover_migrations(project_root=None):
|
||||
"""Return core and feature migrations in deterministic application order."""
|
||||
root = _project_root(project_root)
|
||||
migrations = _read_namespace("core", root / "config" / "migrations")
|
||||
|
||||
modules_root = root / "modules"
|
||||
if modules_root.is_dir():
|
||||
for module_path in sorted(modules_root.iterdir(), key=lambda path: path.name):
|
||||
if module_path.is_dir() and not module_path.name.startswith("_"):
|
||||
migrations.extend(
|
||||
_read_namespace(module_path.name, module_path / "migrations")
|
||||
)
|
||||
|
||||
seen = set()
|
||||
for migration in migrations:
|
||||
key = (migration.namespace, migration.version)
|
||||
if key in seen:
|
||||
raise MigrationError(
|
||||
f"Duplicate migration {migration.namespace}:{migration.version}"
|
||||
)
|
||||
seen.add(key)
|
||||
|
||||
return sorted(
|
||||
migrations,
|
||||
key=lambda item: (
|
||||
item.namespace != "core",
|
||||
item.namespace,
|
||||
item.version,
|
||||
item.name,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _ensure_history_table(cursor):
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
namespace VARCHAR(255) NOT NULL,
|
||||
version INTEGER NOT NULL,
|
||||
checksum CHAR(64) NOT NULL,
|
||||
applied_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
PRIMARY KEY (namespace, version)
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _applied_migrations(cursor):
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT namespace, version, checksum, applied_at
|
||||
FROM schema_migrations
|
||||
ORDER BY namespace, version
|
||||
"""
|
||||
)
|
||||
return {
|
||||
(row["namespace"], row["version"]): dict(row) for row in cursor.fetchall()
|
||||
}
|
||||
|
||||
|
||||
def upgrade(project_root=None):
|
||||
"""Apply pending migrations transactionally under a PostgreSQL advisory lock."""
|
||||
migrations = discover_migrations(project_root)
|
||||
applied_now = []
|
||||
|
||||
with postgres.get_connection() as connection:
|
||||
with connection.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cursor:
|
||||
cursor.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (LOCK_NAME,))
|
||||
_ensure_history_table(cursor)
|
||||
applied = _applied_migrations(cursor)
|
||||
|
||||
for migration in migrations:
|
||||
key = (migration.namespace, migration.version)
|
||||
previous = applied.get(key)
|
||||
if previous:
|
||||
if previous["checksum"] != migration.checksum:
|
||||
raise MigrationError(
|
||||
"Applied migration checksum changed: "
|
||||
f"{migration.namespace}:{migration.version}"
|
||||
)
|
||||
continue
|
||||
|
||||
cursor.execute(migration.sql)
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO schema_migrations (namespace, version, checksum)
|
||||
VALUES (%s, %s, %s)
|
||||
""",
|
||||
(migration.namespace, migration.version, migration.checksum),
|
||||
)
|
||||
applied_now.append(migration)
|
||||
|
||||
return applied_now
|
||||
|
||||
|
||||
def migration_status(project_root=None):
|
||||
"""Return applied, pending, changed, and source-missing migration records."""
|
||||
migrations = discover_migrations(project_root)
|
||||
with postgres.get_connection() as connection:
|
||||
with connection.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cursor:
|
||||
cursor.execute(
|
||||
"SELECT to_regclass('public.schema_migrations') AS table_name"
|
||||
)
|
||||
if cursor.fetchone()["table_name"] is None:
|
||||
applied = {}
|
||||
else:
|
||||
applied = _applied_migrations(cursor)
|
||||
|
||||
status = []
|
||||
source_keys = set()
|
||||
for migration in migrations:
|
||||
key = (migration.namespace, migration.version)
|
||||
source_keys.add(key)
|
||||
record = applied.get(key)
|
||||
state = "pending"
|
||||
applied_at = None
|
||||
if record:
|
||||
state = "applied" if record["checksum"] == migration.checksum else "changed"
|
||||
applied_at = record["applied_at"]
|
||||
status.append(
|
||||
{
|
||||
"namespace": migration.namespace,
|
||||
"version": migration.version,
|
||||
"name": migration.name,
|
||||
"checksum": migration.checksum,
|
||||
"applied_at": applied_at,
|
||||
"state": state,
|
||||
}
|
||||
)
|
||||
|
||||
for key, record in sorted(applied.items()):
|
||||
if key not in source_keys:
|
||||
status.append(
|
||||
{
|
||||
"namespace": key[0],
|
||||
"version": key[1],
|
||||
"name": None,
|
||||
"checksum": record["checksum"],
|
||||
"applied_at": record["applied_at"],
|
||||
"state": "missing",
|
||||
}
|
||||
)
|
||||
return status
|
||||
37
core/migrations/__main__.py
Normal file
37
core/migrations/__main__.py
Normal file
@@ -0,0 +1,37 @@
|
||||
"""Command-line entry point for database migrations."""
|
||||
|
||||
import argparse
|
||||
|
||||
from core.migrations import MigrationError, migration_status, upgrade
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
parser = argparse.ArgumentParser(description="Manage database migrations")
|
||||
parser.add_argument("command", choices=("upgrade", "status"))
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
try:
|
||||
if args.command == "upgrade":
|
||||
applied = upgrade()
|
||||
for migration in applied:
|
||||
print(f"applied {migration.namespace}:{migration.version} {migration.name}")
|
||||
if not applied:
|
||||
print("database is up to date")
|
||||
return 0
|
||||
|
||||
records = migration_status()
|
||||
if not records:
|
||||
print("no migrations found")
|
||||
for record in records:
|
||||
name = record["name"] or "<source missing>"
|
||||
print(
|
||||
f"{record['state']:<8} "
|
||||
f"{record['namespace']}:{record['version']} {name}"
|
||||
)
|
||||
return 1 if any(row["state"] == "changed" for row in records) else 0
|
||||
except MigrationError as error:
|
||||
parser.exit(1, f"migration error: {error}\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
105
core/notifications.py
Normal file
105
core/notifications.py
Normal file
@@ -0,0 +1,105 @@
|
||||
"""Optional Discord-webhook and ntfy notification helpers."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from urllib.parse import quote, urlparse
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
|
||||
import core.postgres as postgres
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
REQUEST_TIMEOUT = float(os.getenv("NOTIFICATION_TIMEOUT", 10))
|
||||
|
||||
|
||||
def _sendToEnabledChannels(notif_settings, message):
|
||||
"""Send to enabled channels and return True when at least one succeeds."""
|
||||
if not isinstance(notif_settings, dict):
|
||||
return False
|
||||
sent = False
|
||||
if notif_settings.get("discord_enabled") and notif_settings.get(
|
||||
"discord_webhook"
|
||||
):
|
||||
sent = discord.send(notif_settings["discord_webhook"], message) or sent
|
||||
if notif_settings.get("ntfy_enabled") and notif_settings.get("ntfy_topic"):
|
||||
sent = ntfy.send(notif_settings["ntfy_topic"], message) or sent
|
||||
return sent
|
||||
|
||||
|
||||
def getNotificationSettings(userUUID):
|
||||
return postgres.select_one("notifications", {"user_uuid": userUUID}) or False
|
||||
|
||||
|
||||
def setNotificationSettings(userUUID, data_dict):
|
||||
if not isinstance(data_dict, dict):
|
||||
return False
|
||||
allowed = {
|
||||
"discord_webhook",
|
||||
"discord_enabled",
|
||||
"ntfy_topic",
|
||||
"ntfy_enabled",
|
||||
}
|
||||
updates = {key: value for key, value in data_dict.items() if key in allowed}
|
||||
if not updates:
|
||||
return False
|
||||
if updates.get("discord_webhook"):
|
||||
_validateDiscordWebhook(updates["discord_webhook"])
|
||||
|
||||
existing = postgres.select_one("notifications", {"user_uuid": userUUID})
|
||||
if existing:
|
||||
postgres.update("notifications", updates, {"user_uuid": userUUID})
|
||||
else:
|
||||
updates["id"] = str(uuid.uuid4())
|
||||
updates["user_uuid"] = userUUID
|
||||
postgres.insert("notifications", updates)
|
||||
return True
|
||||
|
||||
|
||||
def _validateDiscordWebhook(webhookURL):
|
||||
parsed = urlparse(webhookURL)
|
||||
allowedHosts = {"discord.com", "canary.discord.com", "ptb.discord.com"}
|
||||
if parsed.scheme != "https" or parsed.hostname not in allowedHosts:
|
||||
raise ValueError("Discord webhook must use an official HTTPS Discord host")
|
||||
if not parsed.path.startswith("/api/webhooks/"):
|
||||
raise ValueError("Invalid Discord webhook path")
|
||||
return webhookURL
|
||||
|
||||
|
||||
class discord:
|
||||
@staticmethod
|
||||
def send(webhook_url, message):
|
||||
try:
|
||||
_validateDiscordWebhook(webhook_url)
|
||||
response = requests.post(
|
||||
webhook_url,
|
||||
json={"content": str(message)},
|
||||
timeout=REQUEST_TIMEOUT,
|
||||
)
|
||||
return response.status_code in {200, 204}
|
||||
except (ValueError, requests.RequestException) as error:
|
||||
logger.warning("Discord webhook delivery failed: %s", error)
|
||||
return False
|
||||
|
||||
|
||||
class ntfy:
|
||||
@staticmethod
|
||||
def send(topic, message):
|
||||
if not isinstance(topic, str) or not topic.strip():
|
||||
return False
|
||||
baseURL = os.getenv("NTFY_BASE_URL", "https://ntfy.sh").rstrip("/")
|
||||
headers = {}
|
||||
if os.getenv("NTFY_TOKEN"):
|
||||
headers["Authorization"] = f"Bearer {os.environ['NTFY_TOKEN']}"
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{baseURL}/{quote(topic.strip(), safe='')}",
|
||||
data=str(message).encode("utf-8"),
|
||||
headers=headers,
|
||||
timeout=REQUEST_TIMEOUT,
|
||||
)
|
||||
return 200 <= response.status_code < 300
|
||||
except requests.RequestException as error:
|
||||
logger.warning("ntfy delivery failed: %s", error)
|
||||
return False
|
||||
357
core/outbox.py
Normal file
357
core/outbox.py
Normal file
@@ -0,0 +1,357 @@
|
||||
"""Durable, lease-based outbound message operations."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timezone
|
||||
import uuid
|
||||
|
||||
from psycopg2.extras import Json
|
||||
|
||||
from core import postgres
|
||||
|
||||
|
||||
DEFAULT_MAX_ATTEMPTS = 3
|
||||
DEFAULT_LEASE_SECONDS = 300
|
||||
DEFAULT_RETRY_SECONDS = 30
|
||||
MAX_RETRY_SECONDS = 900
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _using_cursor(cursor=None):
|
||||
if cursor is not None:
|
||||
yield cursor
|
||||
return
|
||||
with postgres.get_cursor() as owned_cursor:
|
||||
yield owned_cursor
|
||||
|
||||
|
||||
def _timestamp(value, field="timestamp"):
|
||||
if isinstance(value, str):
|
||||
value = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
if not isinstance(value, datetime):
|
||||
raise ValueError(f"{field} must be a datetime or ISO-8601 string")
|
||||
if value.tzinfo is None:
|
||||
value = value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _positive(value, field, maximum=None):
|
||||
try:
|
||||
value = int(value)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise ValueError(f"{field} must be a whole number") from error
|
||||
if value < 1 or (maximum is not None and value > maximum):
|
||||
suffix = f" and at most {maximum}" if maximum is not None else ""
|
||||
raise ValueError(f"{field} must be at least 1{suffix}")
|
||||
return value
|
||||
|
||||
|
||||
def _row(cursor):
|
||||
record = cursor.fetchone()
|
||||
return dict(record) if record else None
|
||||
|
||||
|
||||
def enqueue_message(
|
||||
user_uuid,
|
||||
channel,
|
||||
payload,
|
||||
idempotency_key,
|
||||
available_at=None,
|
||||
max_attempts=DEFAULT_MAX_ATTEMPTS,
|
||||
message_id=None,
|
||||
cursor=None,
|
||||
):
|
||||
"""Queue a message, returning the existing row for a repeated key."""
|
||||
if not user_uuid:
|
||||
raise ValueError("user_uuid is required")
|
||||
if not isinstance(channel, str) or not channel.strip():
|
||||
raise ValueError("channel is required")
|
||||
if not idempotency_key:
|
||||
raise ValueError("idempotency_key is required")
|
||||
if payload is None:
|
||||
raise ValueError("payload is required")
|
||||
max_attempts = _positive(max_attempts, "max_attempts", 100)
|
||||
available_at = available_at or datetime.now(timezone.utc)
|
||||
values = {
|
||||
"id": str(message_id or uuid.uuid4()),
|
||||
"user_uuid": user_uuid,
|
||||
"channel": channel.strip(),
|
||||
"payload": Json(payload),
|
||||
"idempotency_key": idempotency_key,
|
||||
"available_at": _timestamp(available_at, "available_at"),
|
||||
"max_attempts": max_attempts,
|
||||
}
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
INSERT INTO outbound_messages (
|
||||
id, user_uuid, channel, payload, idempotency_key,
|
||||
available_at, max_attempts
|
||||
) VALUES (
|
||||
%(id)s, %(user_uuid)s, %(channel)s, %(payload)s,
|
||||
%(idempotency_key)s, %(available_at)s, %(max_attempts)s
|
||||
)
|
||||
ON CONFLICT (idempotency_key)
|
||||
DO UPDATE SET idempotency_key = EXCLUDED.idempotency_key
|
||||
WHERE outbound_messages.user_uuid = EXCLUDED.user_uuid
|
||||
AND outbound_messages.channel = EXCLUDED.channel
|
||||
RETURNING *
|
||||
""",
|
||||
values,
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def get_message(message_id, cursor=None):
|
||||
"""Return one outbound message by UUID."""
|
||||
try:
|
||||
message_id = str(uuid.UUID(str(message_id)))
|
||||
except (TypeError, ValueError, AttributeError):
|
||||
return None
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"SELECT * FROM outbound_messages WHERE id = %s", (message_id,)
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def list_messages(user_uuid=None, status=None, channel=None, limit=100, cursor=None):
|
||||
"""List outbound messages newest first with optional filters."""
|
||||
limit = _positive(limit, "limit", 500)
|
||||
clauses = []
|
||||
params = []
|
||||
if user_uuid is not None:
|
||||
clauses.append("user_uuid = %s")
|
||||
params.append(user_uuid)
|
||||
if status is not None:
|
||||
clauses.append("status = %s")
|
||||
params.append(status)
|
||||
if channel is not None:
|
||||
clauses.append("channel = %s")
|
||||
params.append(channel)
|
||||
where = f"WHERE {' AND '.join(clauses)}" if clauses else ""
|
||||
params.append(limit)
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
f"""
|
||||
SELECT * FROM outbound_messages
|
||||
{where}
|
||||
ORDER BY created_at DESC
|
||||
LIMIT %s
|
||||
""",
|
||||
params,
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
|
||||
|
||||
def claim_messages(
|
||||
worker_id,
|
||||
channel=None,
|
||||
limit=10,
|
||||
lease_seconds=DEFAULT_LEASE_SECONDS,
|
||||
cursor=None,
|
||||
):
|
||||
"""Atomically lease deliverable messages while skipping other workers."""
|
||||
if not worker_id:
|
||||
raise ValueError("worker_id is required")
|
||||
limit = _positive(limit, "limit", 100)
|
||||
lease_seconds = _positive(lease_seconds, "lease_seconds")
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
WITH exhausted AS (
|
||||
SELECT id
|
||||
FROM outbound_messages
|
||||
WHERE status = 'delivering'
|
||||
AND (lease_until IS NULL OR lease_until <= NOW())
|
||||
AND attempts >= max_attempts
|
||||
FOR UPDATE SKIP LOCKED
|
||||
)
|
||||
UPDATE outbound_messages AS message
|
||||
SET status = 'failed',
|
||||
lease_until = NULL,
|
||||
leased_by = NULL,
|
||||
last_error = COALESCE(last_error, 'lease expired'),
|
||||
updated_at = NOW()
|
||||
FROM exhausted
|
||||
WHERE message.id = exhausted.id
|
||||
"""
|
||||
)
|
||||
active_cursor.execute(
|
||||
"""
|
||||
WITH candidates AS (
|
||||
SELECT id
|
||||
FROM outbound_messages
|
||||
WHERE available_at <= NOW()
|
||||
AND attempts < max_attempts
|
||||
AND (
|
||||
status = 'pending'
|
||||
OR (
|
||||
status = 'delivering'
|
||||
AND (lease_until IS NULL OR lease_until <= NOW())
|
||||
)
|
||||
)
|
||||
AND (%(channel)s IS NULL OR channel = %(channel)s)
|
||||
ORDER BY available_at, created_at
|
||||
FOR UPDATE SKIP LOCKED
|
||||
LIMIT %(limit)s
|
||||
)
|
||||
UPDATE outbound_messages AS message
|
||||
SET status = 'delivering',
|
||||
attempts = message.attempts + 1,
|
||||
leased_by = %(worker_id)s,
|
||||
lease_until = NOW() + (%(lease_seconds)s * INTERVAL '1 second'),
|
||||
updated_at = NOW()
|
||||
FROM candidates
|
||||
WHERE message.id = candidates.id
|
||||
RETURNING message.*
|
||||
""",
|
||||
{
|
||||
"channel": channel,
|
||||
"limit": limit,
|
||||
"worker_id": worker_id,
|
||||
"lease_seconds": lease_seconds,
|
||||
},
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
|
||||
|
||||
def renew_message_lease(
|
||||
message_id, worker_id, lease_seconds=DEFAULT_LEASE_SECONDS, cursor=None
|
||||
):
|
||||
"""Extend a delivery lease owned by the requesting worker."""
|
||||
lease_seconds = _positive(lease_seconds, "lease_seconds")
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE outbound_messages
|
||||
SET lease_until = NOW() + (%s * INTERVAL '1 second'), updated_at = NOW()
|
||||
WHERE id = %s AND status = 'delivering' AND leased_by = %s
|
||||
RETURNING *
|
||||
""",
|
||||
(lease_seconds, message_id, worker_id),
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def mark_delivered(
|
||||
message_id, worker_id, external_message_id=None, cursor=None
|
||||
):
|
||||
"""Mark a message delivered when its lease is owned by the worker."""
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE outbound_messages
|
||||
SET status = 'delivered', delivered_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL, last_error = NULL,
|
||||
external_message_id = %s
|
||||
WHERE id = %s AND status = 'delivering' AND leased_by = %s
|
||||
RETURNING *
|
||||
""",
|
||||
(external_message_id, message_id, worker_id),
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def retry_message(
|
||||
message_id,
|
||||
worker_id,
|
||||
error,
|
||||
retry_seconds=DEFAULT_RETRY_SECONDS,
|
||||
max_retry_seconds=MAX_RETRY_SECONDS,
|
||||
cursor=None,
|
||||
):
|
||||
"""Fail or reschedule an owned delivery with exponential backoff."""
|
||||
retry_seconds = _positive(retry_seconds, "retry_seconds")
|
||||
max_retry_seconds = _positive(max_retry_seconds, "max_retry_seconds")
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
SELECT attempts, max_attempts
|
||||
FROM outbound_messages
|
||||
WHERE id = %s AND status = 'delivering' AND leased_by = %s
|
||||
FOR UPDATE
|
||||
""",
|
||||
(message_id, worker_id),
|
||||
)
|
||||
current = active_cursor.fetchone()
|
||||
if not current:
|
||||
return None
|
||||
|
||||
exhausted = current["attempts"] >= current["max_attempts"]
|
||||
delay = min(
|
||||
max_retry_seconds,
|
||||
retry_seconds * (2 ** min(max(current["attempts"] - 1, 0), 30)),
|
||||
)
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE outbound_messages
|
||||
SET status = %(status)s,
|
||||
available_at = CASE
|
||||
WHEN %(exhausted)s THEN available_at
|
||||
ELSE NOW() + (%(delay)s * INTERVAL '1 second')
|
||||
END,
|
||||
leased_by = NULL,
|
||||
lease_until = NULL,
|
||||
last_error = %(error)s,
|
||||
updated_at = NOW()
|
||||
WHERE id = %(message_id)s
|
||||
RETURNING *
|
||||
""",
|
||||
{
|
||||
"status": "failed" if exhausted else "pending",
|
||||
"exhausted": exhausted,
|
||||
"delay": delay,
|
||||
"error": str(error)[:4000],
|
||||
"message_id": message_id,
|
||||
},
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def cancel_message(message_id, user_uuid=None, cursor=None):
|
||||
"""Cancel one undelivered message, optionally enforcing its owner."""
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
"""
|
||||
UPDATE outbound_messages
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE id = %(message_id)s
|
||||
AND status IN ('pending', 'delivering')
|
||||
AND (%(user_uuid)s IS NULL OR user_uuid = %(user_uuid)s)
|
||||
RETURNING *
|
||||
""",
|
||||
{"message_id": message_id, "user_uuid": user_uuid},
|
||||
)
|
||||
return _row(active_cursor)
|
||||
|
||||
|
||||
def cancel_messages(
|
||||
user_uuid=None, channel=None, idempotency_key=None, cursor=None
|
||||
):
|
||||
"""Cancel matching undelivered messages; at least one filter is required."""
|
||||
filters = {
|
||||
"user_uuid": user_uuid,
|
||||
"channel": channel,
|
||||
"idempotency_key": idempotency_key,
|
||||
}
|
||||
clauses = [f"{name} = %({name})s" for name, value in filters.items() if value is not None]
|
||||
if not clauses:
|
||||
raise ValueError("at least one cancellation filter is required")
|
||||
|
||||
with _using_cursor(cursor) as active_cursor:
|
||||
active_cursor.execute(
|
||||
f"""
|
||||
UPDATE outbound_messages
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE status IN ('pending', 'delivering')
|
||||
AND {' AND '.join(clauses)}
|
||||
RETURNING *
|
||||
""",
|
||||
filters,
|
||||
)
|
||||
return [dict(record) for record in active_cursor.fetchall()]
|
||||
289
core/postgres.py
Normal file
289
core/postgres.py
Normal file
@@ -0,0 +1,289 @@
|
||||
"""
|
||||
postgres.py - Small parameterized PostgreSQL CRUD layer
|
||||
|
||||
Connection configuration is read from DB_HOST, DB_PORT, DB_NAME, DB_USER,
|
||||
and DB_PASS. Raw SQL remains available through execute() for domain services.
|
||||
"""
|
||||
|
||||
from contextlib import contextmanager
|
||||
import os
|
||||
from pathlib import Path
|
||||
import re
|
||||
|
||||
import psycopg2
|
||||
import psycopg2.extras
|
||||
from dotenv import load_dotenv
|
||||
|
||||
|
||||
load_dotenv(Path(__file__).resolve().parents[1] / ".env", override=False)
|
||||
|
||||
IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
|
||||
|
||||
def _get_config():
|
||||
return {
|
||||
"host": os.environ.get("DB_HOST", "localhost"),
|
||||
"port": int(os.environ.get("DB_PORT", 5432)),
|
||||
"dbname": os.environ.get("DB_NAME", "app"),
|
||||
"user": os.environ.get("DB_USER", "app"),
|
||||
"password": os.environ.get("DB_PASS", ""),
|
||||
}
|
||||
|
||||
|
||||
def _safe_id(name):
|
||||
if not isinstance(name, str) or not IDENTIFIER.fullmatch(name):
|
||||
raise ValueError(f"Invalid SQL identifier: {name}")
|
||||
return f'"{name}"'
|
||||
|
||||
|
||||
def _build_where(where, prefix=""):
|
||||
if not isinstance(where, dict):
|
||||
raise ValueError("where must be a dictionary")
|
||||
clauses = []
|
||||
params = {}
|
||||
for index, (column, value) in enumerate(where.items()):
|
||||
paramName = f"{prefix}{column}_{index}"
|
||||
safeColumn = _safe_id(column)
|
||||
|
||||
if isinstance(value, tuple) and len(value) == 2:
|
||||
operator, operand = value
|
||||
operator = str(operator).upper()
|
||||
allowed = {"=", "!=", "<", ">", "<=", ">=", "LIKE", "ILIKE", "IN"}
|
||||
if operator not in allowed:
|
||||
raise ValueError(f"Unsupported operator: {operator}")
|
||||
if operator == "IN":
|
||||
values = list(operand)
|
||||
if not values:
|
||||
clauses.append("FALSE")
|
||||
continue
|
||||
placeholders = []
|
||||
for itemIndex, item in enumerate(values):
|
||||
itemName = f"{paramName}_{itemIndex}"
|
||||
placeholders.append(f"%({itemName})s")
|
||||
params[itemName] = item
|
||||
clauses.append(f"{safeColumn} IN ({', '.join(placeholders)})")
|
||||
else:
|
||||
clauses.append(f"{safeColumn} {operator} %({paramName})s")
|
||||
params[paramName] = operand
|
||||
elif value is None:
|
||||
clauses.append(f"{safeColumn} IS NULL")
|
||||
else:
|
||||
clauses.append(f"{safeColumn} = %({paramName})s")
|
||||
params[paramName] = value
|
||||
return " AND ".join(clauses), params
|
||||
|
||||
|
||||
def _order_clause(orderBy):
|
||||
if not orderBy:
|
||||
return ""
|
||||
items = orderBy if isinstance(orderBy, (list, tuple)) else str(orderBy).split(",")
|
||||
safeItems = []
|
||||
for item in items:
|
||||
if isinstance(item, (list, tuple)):
|
||||
if len(item) != 2:
|
||||
raise ValueError("order tuple must contain column and direction")
|
||||
column, direction = item
|
||||
else:
|
||||
parts = str(item).strip().split()
|
||||
if not parts or len(parts) > 2:
|
||||
raise ValueError(f"Invalid order expression: {item}")
|
||||
column = parts[0]
|
||||
direction = parts[1] if len(parts) == 2 else "ASC"
|
||||
direction = str(direction).upper()
|
||||
if direction not in {"ASC", "DESC"}:
|
||||
raise ValueError(f"Invalid order direction: {direction}")
|
||||
safeItems.append(f"{_safe_id(column)} {direction}")
|
||||
return ", ".join(safeItems)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def get_connection():
|
||||
connection = psycopg2.connect(**_get_config())
|
||||
try:
|
||||
yield connection
|
||||
connection.commit()
|
||||
except Exception:
|
||||
connection.rollback()
|
||||
raise
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def get_cursor(dict_cursor=True):
|
||||
with get_connection() as connection:
|
||||
factory = psycopg2.extras.RealDictCursor if dict_cursor else None
|
||||
cursor = connection.cursor(cursor_factory=factory)
|
||||
try:
|
||||
yield cursor
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
|
||||
def insert(table, data):
|
||||
if not isinstance(data, dict) or not data:
|
||||
raise ValueError("insert data cannot be empty")
|
||||
columns = list(data.keys())
|
||||
query = f"""
|
||||
INSERT INTO {_safe_id(table)} ({', '.join(_safe_id(col) for col in columns)})
|
||||
VALUES ({', '.join(f'%({col})s' for col in columns)})
|
||||
RETURNING *
|
||||
"""
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, data)
|
||||
record = cursor.fetchone()
|
||||
return dict(record) if record else None
|
||||
|
||||
|
||||
def select(table, where=None, order_by=None, limit=None, offset=None):
|
||||
query = f"SELECT * FROM {_safe_id(table)}"
|
||||
params = {}
|
||||
if where:
|
||||
clauses, params = _build_where(where)
|
||||
query += f" WHERE {clauses}"
|
||||
orderClause = _order_clause(order_by)
|
||||
if orderClause:
|
||||
query += f" ORDER BY {orderClause}"
|
||||
if limit is not None:
|
||||
limit = int(limit)
|
||||
if limit < 0:
|
||||
raise ValueError("limit cannot be negative")
|
||||
query += " LIMIT %(query_limit)s"
|
||||
params["query_limit"] = limit
|
||||
if offset is not None:
|
||||
offset = int(offset)
|
||||
if offset < 0:
|
||||
raise ValueError("offset cannot be negative")
|
||||
query += " OFFSET %(query_offset)s"
|
||||
params["query_offset"] = offset
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, params)
|
||||
return [dict(record) for record in cursor.fetchall()]
|
||||
|
||||
|
||||
def select_one(table, where):
|
||||
records = select(table, where=where, limit=1)
|
||||
return records[0] if records else None
|
||||
|
||||
|
||||
def update(table, data, where):
|
||||
if not isinstance(data, dict) or not data:
|
||||
raise ValueError("update data cannot be empty")
|
||||
if not isinstance(where, dict) or not where:
|
||||
raise ValueError("update requires a non-empty where clause")
|
||||
setClause = ", ".join(f"{_safe_id(col)} = %(set_{col})s" for col in data)
|
||||
params = {f"set_{col}": value for col, value in data.items()}
|
||||
whereClause, whereParams = _build_where(where, prefix="where_")
|
||||
params.update(whereParams)
|
||||
query = f"""
|
||||
UPDATE {_safe_id(table)} SET {setClause}
|
||||
WHERE {whereClause}
|
||||
RETURNING *
|
||||
"""
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, params)
|
||||
return [dict(record) for record in cursor.fetchall()]
|
||||
|
||||
|
||||
def delete(table, where):
|
||||
if not isinstance(where, dict) or not where:
|
||||
raise ValueError("delete requires a non-empty where clause")
|
||||
whereClause, params = _build_where(where)
|
||||
query = f"DELETE FROM {_safe_id(table)} WHERE {whereClause} RETURNING *"
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, params)
|
||||
return [dict(record) for record in cursor.fetchall()]
|
||||
|
||||
|
||||
def count(table, where=None):
|
||||
query = f"SELECT COUNT(*) AS count FROM {_safe_id(table)}"
|
||||
params = {}
|
||||
if where:
|
||||
clauses, params = _build_where(where)
|
||||
query += f" WHERE {clauses}"
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, params)
|
||||
return cursor.fetchone()["count"]
|
||||
|
||||
|
||||
def exists(table, where):
|
||||
return count(table, where) > 0
|
||||
|
||||
|
||||
def upsert(table, data, conflict_columns):
|
||||
if not isinstance(data, dict) or not data:
|
||||
raise ValueError("upsert data cannot be empty")
|
||||
if not conflict_columns:
|
||||
raise ValueError("conflict columns are required")
|
||||
columns = list(data.keys())
|
||||
updates = [column for column in columns if column not in conflict_columns]
|
||||
action = "DO NOTHING"
|
||||
if updates:
|
||||
action = "DO UPDATE SET " + ", ".join(
|
||||
f"{_safe_id(column)} = EXCLUDED.{_safe_id(column)}" for column in updates
|
||||
)
|
||||
query = f"""
|
||||
INSERT INTO {_safe_id(table)} ({', '.join(_safe_id(col) for col in columns)})
|
||||
VALUES ({', '.join(f'%({col})s' for col in columns)})
|
||||
ON CONFLICT ({', '.join(_safe_id(col) for col in conflict_columns)})
|
||||
{action}
|
||||
RETURNING *
|
||||
"""
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, data)
|
||||
record = cursor.fetchone()
|
||||
return dict(record) if record else None
|
||||
|
||||
|
||||
def insert_many(table, rows):
|
||||
if not rows:
|
||||
return 0
|
||||
columns = list(rows[0].keys())
|
||||
if any(list(row.keys()) != columns for row in rows):
|
||||
raise ValueError("all rows must use the same columns in the same order")
|
||||
query = f"""
|
||||
INSERT INTO {_safe_id(table)} ({', '.join(_safe_id(col) for col in columns)})
|
||||
VALUES %s
|
||||
"""
|
||||
template = f"({', '.join(f'%({column})s' for column in columns)})"
|
||||
with get_cursor() as cursor:
|
||||
psycopg2.extras.execute_values(
|
||||
cursor, query, rows, template=template, page_size=100
|
||||
)
|
||||
return cursor.rowcount
|
||||
|
||||
|
||||
def execute(query, params=None):
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(query, params or {})
|
||||
if cursor.description:
|
||||
return [dict(record) for record in cursor.fetchall()]
|
||||
return cursor.rowcount
|
||||
|
||||
|
||||
def table_exists(table):
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT EXISTS (
|
||||
SELECT FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_name = %(table)s
|
||||
)
|
||||
""",
|
||||
{"table": table},
|
||||
)
|
||||
return cursor.fetchone()["exists"]
|
||||
|
||||
|
||||
def get_columns(table):
|
||||
with get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT column_name, data_type, is_nullable, column_default
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = 'public' AND table_name = %(table)s
|
||||
ORDER BY ordinal_position
|
||||
""",
|
||||
{"table": table},
|
||||
)
|
||||
return [dict(record) for record in cursor.fetchall()]
|
||||
171
core/registry.py
Normal file
171
core/registry.py
Normal file
@@ -0,0 +1,171 @@
|
||||
"""
|
||||
registry.py - Discovery and registration for framework feature modules
|
||||
|
||||
Each package directly under ``modules`` may expose ``register(registry)``.
|
||||
The same registry is loaded by the API, Discord bot, and scheduler so a
|
||||
feature can keep its routes, commands, prompts, and jobs together.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import pkgutil
|
||||
|
||||
|
||||
class FrameworkRegistry:
|
||||
def __init__(self):
|
||||
self.modules = {}
|
||||
self.commands = {}
|
||||
self.route_registrars = []
|
||||
self.job_handlers = {}
|
||||
self._loading_module = None
|
||||
|
||||
def clear(self):
|
||||
self.modules.clear()
|
||||
self.commands.clear()
|
||||
self.route_registrars.clear()
|
||||
self.job_handlers.clear()
|
||||
self._loading_module = None
|
||||
|
||||
def begin_module(self, name, package):
|
||||
if name in self.modules:
|
||||
raise ValueError(f"Duplicate module name: {name}")
|
||||
self.modules[name] = {
|
||||
"name": name,
|
||||
"package": package,
|
||||
"description": "",
|
||||
}
|
||||
self._loading_module = name
|
||||
|
||||
def finish_module(self):
|
||||
self._loading_module = None
|
||||
|
||||
def describe(self, description):
|
||||
module_name = self._require_loading_module()
|
||||
self.modules[module_name]["description"] = description.strip()
|
||||
|
||||
def register_command(
|
||||
self,
|
||||
interaction_type,
|
||||
handler,
|
||||
prompt,
|
||||
validator=None,
|
||||
help_text=None,
|
||||
description="",
|
||||
):
|
||||
module_name = self._require_loading_module()
|
||||
if interaction_type in self.commands:
|
||||
raise ValueError(f"Duplicate command type: {interaction_type}")
|
||||
if not callable(handler):
|
||||
raise TypeError(f"Handler for {interaction_type} must be callable")
|
||||
if validator is not None and not callable(validator):
|
||||
raise TypeError(f"Validator for {interaction_type} must be callable")
|
||||
if not isinstance(prompt, dict) or not prompt.get("system") or not prompt.get(
|
||||
"user_template"
|
||||
):
|
||||
raise ValueError(
|
||||
f"Command {interaction_type} must provide system and user_template prompts"
|
||||
)
|
||||
|
||||
self.commands[interaction_type] = {
|
||||
"module": module_name,
|
||||
"handler": handler,
|
||||
"prompt": prompt,
|
||||
"validator": validator,
|
||||
"help_text": list(help_text or []),
|
||||
"description": description.strip(),
|
||||
}
|
||||
|
||||
def register_routes(self, registrar):
|
||||
module_name = self._require_loading_module()
|
||||
if not callable(registrar):
|
||||
raise TypeError(f"Route registrar for {module_name} must be callable")
|
||||
self.route_registrars.append((module_name, registrar))
|
||||
|
||||
def register_job(self, job_type, handler):
|
||||
module_name = self._require_loading_module()
|
||||
if job_type in self.job_handlers:
|
||||
raise ValueError(f"Duplicate job type: {job_type}")
|
||||
if not callable(handler):
|
||||
raise TypeError(f"Job handler for {job_type} must be callable")
|
||||
self.job_handlers[job_type] = {
|
||||
"module": module_name,
|
||||
"handler": handler,
|
||||
}
|
||||
|
||||
def get_command(self, interaction_type):
|
||||
return self.commands.get(interaction_type)
|
||||
|
||||
def get_job_handler(self, job_type):
|
||||
registration = self.job_handlers.get(job_type)
|
||||
return registration["handler"] if registration else None
|
||||
|
||||
def list_commands(self):
|
||||
return list(self.commands.keys())
|
||||
|
||||
def router_context(self):
|
||||
lines = []
|
||||
for name, command in sorted(self.commands.items()):
|
||||
description = command["description"] or "No description provided"
|
||||
lines.append(f"- {name}: {description}")
|
||||
return "\n".join(lines) if lines else "No modules are available"
|
||||
|
||||
def help_lines(self):
|
||||
lines = []
|
||||
for name, command in sorted(self.commands.items()):
|
||||
if command["help_text"]:
|
||||
lines.extend(f"- {item}" for item in command["help_text"])
|
||||
else:
|
||||
lines.append(f"- {name}: {command['description']}")
|
||||
return lines
|
||||
|
||||
def _require_loading_module(self):
|
||||
if not self._loading_module:
|
||||
raise RuntimeError("Registration must happen inside a module register() call")
|
||||
return self._loading_module
|
||||
|
||||
|
||||
registry = FrameworkRegistry()
|
||||
_loaded = False
|
||||
|
||||
|
||||
def discover_modules(force=False):
|
||||
"""Discover and register feature packages exactly once per process."""
|
||||
global _loaded
|
||||
if _loaded and not force:
|
||||
return registry
|
||||
|
||||
_loaded = False
|
||||
|
||||
package = importlib.import_module("modules")
|
||||
discovered = sorted(
|
||||
item.name
|
||||
for item in pkgutil.iter_modules(package.__path__)
|
||||
if item.ispkg and not item.name.startswith("_")
|
||||
)
|
||||
|
||||
registry.clear()
|
||||
try:
|
||||
for module_name in discovered:
|
||||
qualified_name = f"modules.{module_name}"
|
||||
feature_module = importlib.import_module(qualified_name)
|
||||
register_fn = getattr(feature_module, "register", None)
|
||||
if not callable(register_fn):
|
||||
raise RuntimeError(f"{qualified_name} must expose register(registry)")
|
||||
|
||||
registry.begin_module(module_name, qualified_name)
|
||||
try:
|
||||
register_fn(registry)
|
||||
finally:
|
||||
registry.finish_module()
|
||||
except Exception:
|
||||
registry.clear()
|
||||
raise
|
||||
|
||||
_loaded = True
|
||||
return registry
|
||||
|
||||
|
||||
def reset_registry():
|
||||
"""Reset discovery state for tests."""
|
||||
global _loaded
|
||||
registry.clear()
|
||||
_loaded = False
|
||||
205
core/users.py
Normal file
205
core/users.py
Normal file
@@ -0,0 +1,205 @@
|
||||
import os
|
||||
import uuid
|
||||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
import bcrypt
|
||||
from psycopg2.errors import UniqueViolation
|
||||
|
||||
import core.postgres as postgres
|
||||
|
||||
|
||||
MIN_PASSWORD_LENGTH = 8
|
||||
MAX_PASSWORD_BYTES = 72
|
||||
REGISTRATION_FIELDS = {"timezone"}
|
||||
USER_UPDATE_FIELDS = {"timezone"}
|
||||
|
||||
|
||||
def getUser(userUUID):
|
||||
return postgres.select_one("users", {"id": userUUID})
|
||||
|
||||
|
||||
def getUserUUID(username):
|
||||
if not username:
|
||||
return False
|
||||
userRecord = postgres.select_one("users", {"username": username})
|
||||
if userRecord:
|
||||
return userRecord["id"]
|
||||
return False
|
||||
|
||||
|
||||
def getUserFirstName(userUUID):
|
||||
userRecord = getUser(userUUID)
|
||||
if userRecord:
|
||||
return userRecord.get("username")
|
||||
return None
|
||||
|
||||
|
||||
def getUserTimezone(userUUID):
|
||||
userRecord = getUser(userUUID)
|
||||
if userRecord:
|
||||
return userRecord.get("timezone") or "UTC"
|
||||
return None
|
||||
|
||||
|
||||
def isUsernameAvailable(username):
|
||||
if not username:
|
||||
return False
|
||||
return not postgres.exists("users", {"username": username})
|
||||
|
||||
|
||||
def doesUserUUIDExist(userUUID):
|
||||
if not userUUID:
|
||||
return False
|
||||
return postgres.exists("users", {"id": userUUID})
|
||||
|
||||
|
||||
def isValidTimezone(timezoneName):
|
||||
if not isinstance(timezoneName, str) or not timezoneName.strip():
|
||||
return False
|
||||
try:
|
||||
ZoneInfo(timezoneName.strip())
|
||||
return True
|
||||
except (ZoneInfoNotFoundError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def normalizeTimezone(timezoneName):
|
||||
timezoneName = timezoneName.strip() if isinstance(timezoneName, str) else ""
|
||||
if not isValidTimezone(timezoneName):
|
||||
raise ValueError("timezone must be a valid IANA timezone name")
|
||||
return timezoneName
|
||||
|
||||
|
||||
def getDefaultTimezone():
|
||||
return normalizeTimezone(os.getenv("DEFAULT_TIMEZONE", "UTC"))
|
||||
|
||||
|
||||
def validatePassword(password):
|
||||
errors = []
|
||||
if not isinstance(password, str) or not password:
|
||||
return False, ["password"]
|
||||
|
||||
encodedPassword = password.encode("utf-8")
|
||||
if len(password) < MIN_PASSWORD_LENGTH:
|
||||
errors.append(f"password must be at least {MIN_PASSWORD_LENGTH} characters")
|
||||
if len(encodedPassword) > MAX_PASSWORD_BYTES:
|
||||
errors.append(f"password must be at most {MAX_PASSWORD_BYTES} UTF-8 bytes")
|
||||
return not errors, errors
|
||||
|
||||
|
||||
def registerUser(username, password, data=None):
|
||||
if not isinstance(username, str) or not username.strip():
|
||||
raise ValueError("username is required")
|
||||
username = username.strip()
|
||||
if len(username) > 255:
|
||||
raise ValueError("username must be at most 255 characters")
|
||||
|
||||
isValid, errors = validatePassword(password)
|
||||
if not isValid:
|
||||
raise ValueError(f"Invalid password: {', '.join(errors)}")
|
||||
if not isUsernameAvailable(username):
|
||||
return False
|
||||
|
||||
userData = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"username": username,
|
||||
"password_hashed": bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()),
|
||||
}
|
||||
if data is not None:
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("registration data must be an object")
|
||||
for key in REGISTRATION_FIELDS:
|
||||
if key in data:
|
||||
userData[key] = data[key]
|
||||
|
||||
try:
|
||||
createUser(userData)
|
||||
except UniqueViolation:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def updateUser(userUUID, data_dict):
|
||||
if not getUser(userUUID) or not isinstance(data_dict, dict):
|
||||
return False
|
||||
|
||||
updates = {key: data_dict[key] for key in USER_UPDATE_FIELDS if key in data_dict}
|
||||
if "timezone" in updates:
|
||||
updates["timezone"] = normalizeTimezone(updates["timezone"])
|
||||
if not updates:
|
||||
return False
|
||||
|
||||
postgres.update("users", updates, {"id": userUUID})
|
||||
return True
|
||||
|
||||
|
||||
def changePassword(userUUID, new_password):
|
||||
if not getUser(userUUID):
|
||||
return False
|
||||
isValid, errors = validatePassword(new_password)
|
||||
if not isValid:
|
||||
raise ValueError(f"Invalid password: {', '.join(errors)}")
|
||||
|
||||
hashedPassword = bcrypt.hashpw(new_password.encode("utf-8"), bcrypt.gensalt())
|
||||
postgres.update("users", {"password_hashed": hashedPassword}, {"id": userUUID})
|
||||
return True
|
||||
|
||||
|
||||
def deleteUser(userUUID):
|
||||
if not getUser(userUUID):
|
||||
return False
|
||||
postgres.delete("users", {"id": userUUID})
|
||||
return True
|
||||
|
||||
|
||||
def createUser(data_dict):
|
||||
if not isinstance(data_dict, dict):
|
||||
raise ValueError("user data must be an object")
|
||||
|
||||
allowedFields = {"id", "username", "password_hashed", "timezone", "created_at"}
|
||||
userData = {key: value for key, value in data_dict.items() if key in allowedFields}
|
||||
if "timezone" not in userData:
|
||||
userData["timezone"] = getDefaultTimezone()
|
||||
|
||||
isValid, errors = validateUser(userData)
|
||||
if not isValid:
|
||||
raise ValueError(f"Invalid user data: {', '.join(errors)}")
|
||||
|
||||
return postgres.insert("users", userData)
|
||||
|
||||
|
||||
def validateUser(user):
|
||||
errors = []
|
||||
if not isinstance(user, dict):
|
||||
return False, ["user"]
|
||||
|
||||
userUUID = user.get("id")
|
||||
if not userUUID:
|
||||
errors.append("id")
|
||||
else:
|
||||
try:
|
||||
uuid.UUID(str(userUUID))
|
||||
except (TypeError, ValueError, AttributeError):
|
||||
errors.append("id must be a valid UUID")
|
||||
|
||||
username = user.get("username")
|
||||
passwordHash = user.get("password_hashed")
|
||||
if username is None and passwordHash is not None:
|
||||
errors.append("username")
|
||||
elif username is not None and passwordHash is None:
|
||||
errors.append("password_hashed")
|
||||
elif username is not None:
|
||||
if not isinstance(username, str) or not username.strip():
|
||||
errors.append("username")
|
||||
elif len(username) > 255:
|
||||
errors.append("username must be at most 255 characters")
|
||||
if isinstance(passwordHash, memoryview):
|
||||
passwordHash = bytes(passwordHash)
|
||||
if not isinstance(passwordHash, bytes) or not passwordHash:
|
||||
errors.append("password_hashed")
|
||||
|
||||
timezoneName = user.get("timezone", "UTC")
|
||||
if not isValidTimezone(timezoneName):
|
||||
errors.append("timezone must be a valid IANA timezone name")
|
||||
|
||||
return not errors, errors
|
||||
81
docker-compose.yml
Normal file
81
docker-compose.yml
Normal file
@@ -0,0 +1,81 @@
|
||||
services:
|
||||
db:
|
||||
image: postgres:16
|
||||
restart: unless-stopped
|
||||
env_file:
|
||||
- ${ENV_FILE:-.env}
|
||||
environment:
|
||||
POSTGRES_DB: ${DB_NAME:-app}
|
||||
POSTGRES_USER: ${DB_USER:-app}
|
||||
POSTGRES_PASSWORD: ${DB_PASS}
|
||||
ports:
|
||||
- "5432:5432"
|
||||
volumes:
|
||||
- pgdata:/var/lib/postgresql/data
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "pg_isready -U $${POSTGRES_USER} -d $${POSTGRES_DB}"]
|
||||
interval: 5s
|
||||
timeout: 5s
|
||||
retries: 10
|
||||
start_period: 5s
|
||||
|
||||
migrate:
|
||||
build: .
|
||||
command: ["python", "-m", "core.migrations", "upgrade"]
|
||||
environment:
|
||||
DB_HOST: db
|
||||
env_file:
|
||||
- ${ENV_FILE:-.env}
|
||||
depends_on:
|
||||
db:
|
||||
condition: service_healthy
|
||||
restart: "no"
|
||||
|
||||
app:
|
||||
build: .
|
||||
restart: unless-stopped
|
||||
init: true
|
||||
ports:
|
||||
- "8080:5000"
|
||||
environment:
|
||||
DB_HOST: db
|
||||
env_file:
|
||||
- ${ENV_FILE:-.env}
|
||||
depends_on:
|
||||
migrate:
|
||||
condition: service_completed_successfully
|
||||
healthcheck:
|
||||
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:5000/health/ready', timeout=3).read()"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 10
|
||||
start_period: 10s
|
||||
|
||||
scheduler:
|
||||
build: .
|
||||
command: ["python", "-m", "scheduler.daemon"]
|
||||
environment:
|
||||
DB_HOST: db
|
||||
restart: unless-stopped
|
||||
init: true
|
||||
env_file:
|
||||
- ${ENV_FILE:-.env}
|
||||
depends_on:
|
||||
migrate:
|
||||
condition: service_completed_successfully
|
||||
|
||||
bot:
|
||||
build: .
|
||||
command: ["python", "-m", "bot.bot"]
|
||||
environment:
|
||||
API_URL: http://app:5000
|
||||
restart: unless-stopped
|
||||
init: true
|
||||
env_file:
|
||||
- ${ENV_FILE:-.env}
|
||||
depends_on:
|
||||
app:
|
||||
condition: service_healthy
|
||||
|
||||
volumes:
|
||||
pgdata:
|
||||
1
modules/__init__.py
Normal file
1
modules/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Feature packages discovered by core.registry."""
|
||||
25
modules/reminders/__init__.py
Normal file
25
modules/reminders/__init__.py
Normal file
@@ -0,0 +1,25 @@
|
||||
"""Reference reminders feature module."""
|
||||
|
||||
from modules.reminders.commands import handleReminder, validateReminder
|
||||
from modules.reminders.prompts import REMINDER_PROMPT
|
||||
from modules.reminders.routes import registerRoutes
|
||||
from modules.reminders.service import JOB_TYPE, runReminderJob
|
||||
|
||||
|
||||
def register(registry):
|
||||
registry.describe("Create, list, cancel, and deliver scheduled reminders")
|
||||
registry.register_command(
|
||||
"reminder",
|
||||
handleReminder,
|
||||
prompt=REMINDER_PROMPT,
|
||||
validator=validateReminder,
|
||||
description="Create, list, cancel, or configure reminders",
|
||||
help_text=[
|
||||
"remind me tomorrow at 9 AM to call the dentist",
|
||||
"list my reminders",
|
||||
"cancel a reminder by its ID",
|
||||
"set my timezone to America/Chicago",
|
||||
],
|
||||
)
|
||||
registry.register_routes(registerRoutes)
|
||||
registry.register_job(JOB_TYPE, runReminderJob)
|
||||
109
modules/reminders/commands.py
Normal file
109
modules/reminders/commands.py
Normal file
@@ -0,0 +1,109 @@
|
||||
"""Discord-facing reminder command handler and parser validation."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from core import users
|
||||
from modules.reminders import service
|
||||
|
||||
|
||||
VALID_ACTIONS = {"create", "list", "cancel", "set_timezone"}
|
||||
|
||||
|
||||
def validateReminder(data):
|
||||
if not isinstance(data, dict):
|
||||
return ["Response must be a JSON object"]
|
||||
if data.get("needs_clarification"):
|
||||
return []
|
||||
|
||||
errors = []
|
||||
action = data.get("action")
|
||||
if action not in VALID_ACTIONS:
|
||||
errors.append("action must be create, list, cancel, or set_timezone")
|
||||
return errors
|
||||
|
||||
if action == "create":
|
||||
if not isinstance(data.get("message"), str) or not data["message"].strip():
|
||||
errors.append("create requires a reminder message")
|
||||
try:
|
||||
runAt = datetime.fromisoformat(
|
||||
str(data.get("run_at", "")).replace("Z", "+00:00")
|
||||
)
|
||||
if runAt.tzinfo is None:
|
||||
errors.append("run_at must include a timezone offset")
|
||||
elif runAt.astimezone(timezone.utc) <= datetime.now(timezone.utc):
|
||||
errors.append("run_at must be in the future")
|
||||
except ValueError:
|
||||
errors.append("create requires an ISO-8601 run_at")
|
||||
try:
|
||||
service.normalizeRecurrence(data.get("recurrence"))
|
||||
except ValueError as error:
|
||||
errors.append(str(error))
|
||||
elif action == "cancel" and not data.get("reminder_id"):
|
||||
errors.append("cancel requires reminder_id from the reminder list")
|
||||
elif action == "set_timezone" and not users.isValidTimezone(data.get("timezone")):
|
||||
errors.append("set_timezone requires a valid IANA timezone")
|
||||
return errors
|
||||
|
||||
|
||||
async def handleReminder(context, parsed):
|
||||
action = parsed["action"]
|
||||
if action == "create":
|
||||
result, status = await context.api.request(
|
||||
"post",
|
||||
"/api/reminders",
|
||||
{
|
||||
"message": parsed["message"],
|
||||
"run_at": parsed["run_at"],
|
||||
"recurrence": parsed.get("recurrence"),
|
||||
},
|
||||
)
|
||||
if status == 201:
|
||||
recurrence = " (recurring)" if result.get("recurrence") else ""
|
||||
await context.reply(
|
||||
f"Reminder set for **{result['next_run_at']}**{recurrence}: "
|
||||
f"{result['message']}"
|
||||
)
|
||||
else:
|
||||
await context.reply(_errorMessage(result, "I couldn't create that reminder."))
|
||||
return
|
||||
|
||||
if action == "list":
|
||||
result, status = await context.api.request("get", "/api/reminders")
|
||||
if status != 200:
|
||||
await context.reply(_errorMessage(result, "I couldn't list reminders."))
|
||||
return
|
||||
reminders = result.get("reminders", [])
|
||||
if not reminders:
|
||||
await context.reply("You have no active reminders.")
|
||||
return
|
||||
lines = [
|
||||
f"- `{item['id']}` — {item['next_run_at']}: {item['message']}"
|
||||
for item in reminders
|
||||
]
|
||||
await context.reply("**Active reminders:**\n" + "\n".join(lines))
|
||||
return
|
||||
|
||||
if action == "cancel":
|
||||
result, status = await context.api.request(
|
||||
"delete", f"/api/reminders/{parsed['reminder_id']}"
|
||||
)
|
||||
if status == 200:
|
||||
await context.reply(f"Cancelled reminder: {result['message']}")
|
||||
else:
|
||||
await context.reply(_errorMessage(result, "I couldn't cancel that reminder."))
|
||||
return
|
||||
|
||||
result, status = await context.api.request(
|
||||
"put", "/api/user/me/timezone", {"timezone": parsed["timezone"]}
|
||||
)
|
||||
if status == 200:
|
||||
context.api.timezone = result["timezone"]
|
||||
context.timezone = result["timezone"]
|
||||
await context.reply(f"Your timezone is now **{result['timezone']}**.")
|
||||
else:
|
||||
await context.reply(_errorMessage(result, "I couldn't update your timezone."))
|
||||
|
||||
|
||||
def _errorMessage(result, fallback):
|
||||
detail = result.get("error") if isinstance(result, dict) else None
|
||||
return f"{fallback} {detail}" if detail else fallback
|
||||
21
modules/reminders/migrations/0001_reminders.sql
Normal file
21
modules/reminders/migrations/0001_reminders.sql
Normal file
@@ -0,0 +1,21 @@
|
||||
CREATE TABLE IF NOT EXISTS reminders (
|
||||
id UUID PRIMARY KEY,
|
||||
user_uuid UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
message TEXT NOT NULL,
|
||||
timezone VARCHAR(64) NOT NULL,
|
||||
recurrence JSONB,
|
||||
next_run_at TIMESTAMPTZ NOT NULL,
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'active',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
completed_at TIMESTAMPTZ,
|
||||
cancelled_at TIMESTAMPTZ,
|
||||
CONSTRAINT reminders_status_check
|
||||
CHECK (status IN ('active', 'completed', 'cancelled')),
|
||||
CONSTRAINT reminders_recurrence_check
|
||||
CHECK (recurrence IS NULL OR jsonb_typeof(recurrence) = 'object')
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS reminders_user_active_idx
|
||||
ON reminders (user_uuid, next_run_at)
|
||||
WHERE status = 'active';
|
||||
18
modules/reminders/prompts.py
Normal file
18
modules/reminders/prompts.py
Normal file
@@ -0,0 +1,18 @@
|
||||
REMINDER_PROMPT = {
|
||||
"system": (
|
||||
"You parse reminder commands. Return only a JSON object and do not "
|
||||
"invent missing dates, messages, timezones, or reminder IDs."
|
||||
),
|
||||
"user_template": (
|
||||
"User timezone: {timezone}\n"
|
||||
"Current UTC time: {current_time}\n"
|
||||
"Conversation context:\n{history_context}\n\n"
|
||||
"User message: \"{user_input}\"\n\n"
|
||||
"Return an action: create, list, cancel, or set_timezone. "
|
||||
"For create include message, an ISO-8601 run_at with UTC offset, and "
|
||||
"recurrence as null or an object with frequency daily/weekly and "
|
||||
"optional interval. For cancel include reminder_id. For set_timezone "
|
||||
"include an IANA timezone. If required information is missing, include "
|
||||
"needs_clarification instead."
|
||||
),
|
||||
}
|
||||
71
modules/reminders/routes.py
Normal file
71
modules/reminders/routes.py
Normal file
@@ -0,0 +1,71 @@
|
||||
"""Authenticated reminder API routes."""
|
||||
|
||||
from datetime import datetime
|
||||
import flask
|
||||
import uuid
|
||||
|
||||
from api.security import jsonObject, requireUser
|
||||
from modules.reminders import service
|
||||
|
||||
|
||||
def registerRoutes(app):
|
||||
@app.route("/api/reminders", methods=["GET"])
|
||||
@requireUser()
|
||||
def api_listReminders():
|
||||
includeFinished = flask.request.args.get("include_finished", "false").lower()
|
||||
includeFinished = includeFinished in {"1", "true", "yes"}
|
||||
reminders = service.listReminders(
|
||||
flask.g.user_uuid,
|
||||
includeFinished=includeFinished,
|
||||
limit=flask.request.args.get("limit", 50),
|
||||
)
|
||||
return flask.jsonify(
|
||||
{"reminders": [_serializeReminder(item) for item in reminders]}
|
||||
), 200
|
||||
|
||||
@app.route("/api/reminders", methods=["POST"])
|
||||
@requireUser()
|
||||
def api_createReminder():
|
||||
data = jsonObject()
|
||||
if data is None:
|
||||
return flask.jsonify({"error": "JSON object required"}), 400
|
||||
try:
|
||||
reminder = service.createReminder(
|
||||
flask.g.user_uuid,
|
||||
data.get("message"),
|
||||
data.get("run_at"),
|
||||
data.get("timezone") or _userTimezone(),
|
||||
recurrence=data.get("recurrence"),
|
||||
)
|
||||
except ValueError as error:
|
||||
return flask.jsonify({"error": str(error)}), 400
|
||||
return flask.jsonify(_serializeReminder(reminder)), 201
|
||||
|
||||
@app.route("/api/reminders/<reminderID>", methods=["DELETE"])
|
||||
@requireUser()
|
||||
def api_cancelReminder(reminderID):
|
||||
try:
|
||||
reminder = service.cancelReminder(flask.g.user_uuid, reminderID)
|
||||
except (TypeError, ValueError):
|
||||
reminder = None
|
||||
if not reminder:
|
||||
return flask.jsonify({"error": "active reminder not found"}), 404
|
||||
return flask.jsonify(_serializeReminder(reminder)), 200
|
||||
|
||||
|
||||
def _userTimezone():
|
||||
from core import users
|
||||
|
||||
return users.getUserTimezone(flask.g.user_uuid) or "UTC"
|
||||
|
||||
|
||||
def _serializeReminder(reminder):
|
||||
serialized = {}
|
||||
for key, value in reminder.items():
|
||||
if isinstance(value, datetime):
|
||||
serialized[key] = value.isoformat()
|
||||
elif isinstance(value, uuid.UUID):
|
||||
serialized[key] = str(value)
|
||||
else:
|
||||
serialized[key] = value
|
||||
return serialized
|
||||
246
modules/reminders/service.py
Normal file
246
modules/reminders/service.py
Normal file
@@ -0,0 +1,246 @@
|
||||
"""Database and scheduling operations for the reminders feature."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import uuid
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from psycopg2.extras import Json
|
||||
|
||||
from core import jobs, outbox, postgres, users
|
||||
|
||||
|
||||
JOB_TYPE = "reminders.deliver"
|
||||
CHANNEL = "discord_dm"
|
||||
|
||||
|
||||
def _asUtc(value, field="run_at"):
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError as error:
|
||||
raise ValueError(f"{field} must be an ISO-8601 datetime") from error
|
||||
if not isinstance(value, datetime) or value.tzinfo is None:
|
||||
raise ValueError(f"{field} must include a timezone offset")
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def normalizeRecurrence(recurrence):
|
||||
if recurrence in (None, False, "none", "once"):
|
||||
return None
|
||||
if isinstance(recurrence, str):
|
||||
recurrence = {"frequency": recurrence}
|
||||
if not isinstance(recurrence, dict):
|
||||
raise ValueError("recurrence must be null or an object")
|
||||
|
||||
frequency = str(recurrence.get("frequency", "")).strip().lower()
|
||||
if frequency not in {"daily", "weekly"}:
|
||||
raise ValueError("recurrence frequency must be daily or weekly")
|
||||
try:
|
||||
interval = int(recurrence.get("interval", 1))
|
||||
except (TypeError, ValueError) as error:
|
||||
raise ValueError("recurrence interval must be a number") from error
|
||||
if interval < 1 or interval > 365:
|
||||
raise ValueError("recurrence interval must be between 1 and 365")
|
||||
return {"frequency": frequency, "interval": interval}
|
||||
|
||||
|
||||
def _jobKey(reminderID, scheduledFor):
|
||||
return f"reminder:{reminderID}:{scheduledFor.isoformat()}"
|
||||
|
||||
|
||||
def createReminder(userUUID, message, runAt, timezoneName, recurrence=None):
|
||||
if not isinstance(message, str) or not message.strip():
|
||||
raise ValueError("reminder message is required")
|
||||
message = message.strip()
|
||||
if len(message) > 1800:
|
||||
raise ValueError("reminder message must be at most 1800 characters")
|
||||
timezoneName = users.normalizeTimezone(timezoneName)
|
||||
runAt = _asUtc(runAt)
|
||||
if runAt <= datetime.now(timezone.utc):
|
||||
raise ValueError("reminder time must be in the future")
|
||||
recurrence = normalizeRecurrence(recurrence)
|
||||
reminderID = str(uuid.uuid4())
|
||||
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO reminders (
|
||||
id, user_uuid, message, timezone, recurrence, next_run_at
|
||||
) VALUES (
|
||||
%(id)s, %(user_uuid)s, %(message)s, %(timezone)s,
|
||||
%(recurrence)s, %(next_run_at)s
|
||||
)
|
||||
RETURNING *
|
||||
""",
|
||||
{
|
||||
"id": reminderID,
|
||||
"user_uuid": userUUID,
|
||||
"message": message,
|
||||
"timezone": timezoneName,
|
||||
"recurrence": Json(recurrence) if recurrence else None,
|
||||
"next_run_at": runAt,
|
||||
},
|
||||
)
|
||||
reminder = dict(cursor.fetchone())
|
||||
jobs.create_job(
|
||||
JOB_TYPE,
|
||||
{"reminder_id": reminderID, "scheduled_for": runAt.isoformat()},
|
||||
runAt,
|
||||
user_uuid=userUUID,
|
||||
idempotency_key=_jobKey(reminderID, runAt),
|
||||
cursor=cursor,
|
||||
)
|
||||
return reminder
|
||||
|
||||
|
||||
def listReminders(userUUID, includeFinished=False, limit=50):
|
||||
try:
|
||||
limit = min(max(int(limit), 1), 100)
|
||||
except (TypeError, ValueError):
|
||||
limit = 50
|
||||
statusClause = "" if includeFinished else "AND status = 'active'"
|
||||
return postgres.execute(
|
||||
f"""
|
||||
SELECT * FROM reminders
|
||||
WHERE user_uuid = %(user_uuid)s {statusClause}
|
||||
ORDER BY next_run_at, created_at
|
||||
LIMIT %(limit)s
|
||||
""",
|
||||
{"user_uuid": userUUID, "limit": limit},
|
||||
)
|
||||
|
||||
|
||||
def getReminder(userUUID, reminderID, cursor=None, forUpdate=False):
|
||||
try:
|
||||
uuid.UUID(str(reminderID))
|
||||
except (TypeError, ValueError, AttributeError):
|
||||
return None
|
||||
lock = " FOR UPDATE" if forUpdate else ""
|
||||
query = (
|
||||
"SELECT * FROM reminders WHERE id = %(id)s AND user_uuid = %(user_uuid)s"
|
||||
+ lock
|
||||
)
|
||||
if cursor is not None:
|
||||
cursor.execute(query, {"id": reminderID, "user_uuid": userUUID})
|
||||
record = cursor.fetchone()
|
||||
return dict(record) if record else None
|
||||
rows = postgres.execute(query, {"id": reminderID, "user_uuid": userUUID})
|
||||
return rows[0] if rows else None
|
||||
|
||||
|
||||
def cancelReminder(userUUID, reminderID):
|
||||
with postgres.get_cursor() as cursor:
|
||||
reminder = getReminder(userUUID, reminderID, cursor=cursor, forUpdate=True)
|
||||
if not reminder or reminder["status"] != "active":
|
||||
return None
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE reminders
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW()
|
||||
WHERE id = %(id)s
|
||||
RETURNING *
|
||||
""",
|
||||
{"id": reminderID},
|
||||
)
|
||||
cancelled = dict(cursor.fetchone())
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE scheduled_jobs
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE user_uuid = %(user_uuid)s
|
||||
AND job_type = %(job_type)s
|
||||
AND payload->>'reminder_id' = %(reminder_id)s
|
||||
AND status IN ('pending', 'running')
|
||||
""",
|
||||
{
|
||||
"user_uuid": userUUID,
|
||||
"job_type": JOB_TYPE,
|
||||
"reminder_id": reminderID,
|
||||
},
|
||||
)
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE outbound_messages
|
||||
SET status = 'cancelled', cancelled_at = NOW(), updated_at = NOW(),
|
||||
leased_by = NULL, lease_until = NULL
|
||||
WHERE user_uuid = %(user_uuid)s
|
||||
AND payload->>'reminder_id' = %(reminder_id)s
|
||||
AND status = 'pending'
|
||||
""",
|
||||
{"user_uuid": userUUID, "reminder_id": reminderID},
|
||||
)
|
||||
return cancelled
|
||||
|
||||
|
||||
def _nextRun(scheduledFor, recurrence, timezoneName, now=None):
|
||||
now = now or datetime.now(timezone.utc)
|
||||
localRun = _asUtc(scheduledFor, "scheduled_for").astimezone(
|
||||
ZoneInfo(timezoneName)
|
||||
)
|
||||
interval = recurrence.get("interval", 1)
|
||||
days = interval if recurrence["frequency"] == "daily" else interval * 7
|
||||
nextLocal = localRun + timedelta(days=days)
|
||||
while nextLocal.astimezone(timezone.utc) <= now:
|
||||
nextLocal += timedelta(days=days)
|
||||
return nextLocal.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def runReminderJob(job, workerID):
|
||||
payload = job.get("payload") or {}
|
||||
reminderID = payload.get("reminder_id")
|
||||
scheduledFor = _asUtc(payload.get("scheduled_for"), "scheduled_for")
|
||||
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"SELECT * FROM reminders WHERE id = %s FOR UPDATE", (reminderID,)
|
||||
)
|
||||
record = cursor.fetchone()
|
||||
reminder = dict(record) if record else None
|
||||
if not reminder or reminder["status"] != "active":
|
||||
return jobs.complete_job(job["id"], workerID, cursor=cursor)
|
||||
|
||||
outbox.enqueue_message(
|
||||
reminder["user_uuid"],
|
||||
CHANNEL,
|
||||
{
|
||||
"content": f"Reminder: {reminder['message']}",
|
||||
"reminder_id": str(reminder["id"]),
|
||||
"scheduled_for": scheduledFor.isoformat(),
|
||||
},
|
||||
idempotency_key=_jobKey(reminder["id"], scheduledFor),
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
recurrence = reminder.get("recurrence")
|
||||
if recurrence:
|
||||
nextRun = _nextRun(scheduledFor, recurrence, reminder["timezone"])
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE reminders
|
||||
SET next_run_at = %s, updated_at = NOW()
|
||||
WHERE id = %s
|
||||
""",
|
||||
(nextRun, reminder["id"]),
|
||||
)
|
||||
jobs.create_job(
|
||||
JOB_TYPE,
|
||||
{
|
||||
"reminder_id": str(reminder["id"]),
|
||||
"scheduled_for": nextRun.isoformat(),
|
||||
},
|
||||
nextRun,
|
||||
user_uuid=reminder["user_uuid"],
|
||||
idempotency_key=_jobKey(reminder["id"], nextRun),
|
||||
cursor=cursor,
|
||||
)
|
||||
else:
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE reminders
|
||||
SET status = 'completed', completed_at = NOW(), updated_at = NOW()
|
||||
WHERE id = %s
|
||||
""",
|
||||
(reminder["id"],),
|
||||
)
|
||||
return jobs.complete_job(job["id"], workerID, cursor=cursor)
|
||||
27
pyproject.toml
Normal file
27
pyproject.toml
Normal file
@@ -0,0 +1,27 @@
|
||||
[tool.pytest.ini_options]
|
||||
addopts = "-ra --strict-markers --cov --cov-report=term-missing"
|
||||
asyncio_mode = "auto"
|
||||
pythonpath = ["."]
|
||||
testpaths = ["tests"]
|
||||
|
||||
[tool.coverage.run]
|
||||
branch = true
|
||||
relative_files = true
|
||||
source = ["ai", "api", "bot", "core", "modules", "scheduler"]
|
||||
omit = [
|
||||
"core/manage.py",
|
||||
"core/migrations/__main__.py",
|
||||
]
|
||||
|
||||
[tool.coverage.report]
|
||||
show_missing = true
|
||||
skip_covered = true
|
||||
|
||||
[tool.ruff]
|
||||
target-version = "py311"
|
||||
line-length = 88
|
||||
extend-exclude = [".venv"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
# Keep linting focused on syntax errors, undefined names, and invalid control flow.
|
||||
select = ["E9", "F63", "F7", "F82"]
|
||||
6
requirements-dev.txt
Normal file
6
requirements-dev.txt
Normal file
@@ -0,0 +1,6 @@
|
||||
-r requirements.txt
|
||||
|
||||
pytest==9.1.1
|
||||
pytest-asyncio==1.4.0
|
||||
pytest-cov==7.1.0
|
||||
ruff==0.15.22
|
||||
12
requirements.in
Normal file
12
requirements.in
Normal file
@@ -0,0 +1,12 @@
|
||||
# Runtime dependencies. Keep compatible ranges here and tested pins in requirements.txt.
|
||||
Flask>=3.1,<4
|
||||
psycopg2-binary>=2.9,<3
|
||||
bcrypt>=5,<6
|
||||
PyJWT>=2.13,<3
|
||||
discord.py>=2.7,<3
|
||||
openai>=2.46,<3
|
||||
requests>=2.34,<3
|
||||
httpx>=0.28,<1
|
||||
python-dotenv>=1.2,<2
|
||||
gunicorn>=26,<27
|
||||
tzdata>=2026.3
|
||||
13
requirements.txt
Normal file
13
requirements.txt
Normal file
@@ -0,0 +1,13 @@
|
||||
# Tested direct dependency pins for Python 3.11.
|
||||
# Update this file together with requirements.in.
|
||||
Flask==3.1.3
|
||||
psycopg2-binary==2.9.12
|
||||
bcrypt==5.0.0
|
||||
PyJWT==2.13.0
|
||||
discord.py==2.7.1
|
||||
openai==2.46.0
|
||||
requests==2.34.2
|
||||
httpx==0.28.1
|
||||
python-dotenv==1.2.2
|
||||
gunicorn==26.0.0
|
||||
tzdata==2026.3
|
||||
1
scheduler/__init__.py
Normal file
1
scheduler/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Persisted job worker."""
|
||||
72
scheduler/daemon.py
Normal file
72
scheduler/daemon.py
Normal file
@@ -0,0 +1,72 @@
|
||||
"""PostgreSQL-backed worker for registered feature jobs."""
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import socket
|
||||
import time
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from core import jobs
|
||||
from core.registry import discover_modules
|
||||
|
||||
|
||||
load_dotenv(Path(__file__).resolve().parents[1] / ".env", override=False)
|
||||
|
||||
logging.basicConfig(level=os.getenv("LOG_LEVEL", "INFO"))
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
POLL_INTERVAL = float(os.getenv("JOB_POLL_INTERVAL", 5))
|
||||
JOB_BATCH_SIZE = int(os.getenv("JOB_BATCH_SIZE", 20))
|
||||
JOB_LEASE_SECONDS = int(os.getenv("JOB_LEASE_SECONDS", 300))
|
||||
WORKER_ID = f"scheduler:{socket.gethostname()}:{os.getpid()}"
|
||||
|
||||
module_registry = discover_modules()
|
||||
|
||||
|
||||
def runJob(job):
|
||||
handler = module_registry.get_job_handler(job["job_type"])
|
||||
if not handler:
|
||||
jobs.fail_job(job["id"], WORKER_ID, f"unknown job type: {job['job_type']}")
|
||||
return
|
||||
|
||||
try:
|
||||
result = handler(job, WORKER_ID)
|
||||
if inspect.isawaitable(result):
|
||||
asyncio.run(result)
|
||||
current = jobs.get_job(job["id"])
|
||||
if current and current["status"] == "running":
|
||||
jobs.complete_job(job["id"], WORKER_ID)
|
||||
except Exception as error:
|
||||
logger.exception("Job failed: %s", job["id"])
|
||||
jobs.fail_job(job["id"], WORKER_ID, error)
|
||||
|
||||
|
||||
def pollJobs():
|
||||
claimed = jobs.claim_due_jobs(
|
||||
WORKER_ID,
|
||||
limit=JOB_BATCH_SIZE,
|
||||
lease_seconds=JOB_LEASE_SECONDS,
|
||||
)
|
||||
for job in claimed:
|
||||
runJob(job)
|
||||
return len(claimed)
|
||||
|
||||
|
||||
def daemonLoop():
|
||||
logger.info("Scheduler starting as %s", WORKER_ID)
|
||||
while True:
|
||||
try:
|
||||
claimed = pollJobs()
|
||||
except Exception:
|
||||
logger.exception("Scheduler poll failed")
|
||||
claimed = 0
|
||||
if claimed == 0:
|
||||
time.sleep(POLL_INTERVAL)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
daemonLoop()
|
||||
1
tests/integration/__init__.py
Normal file
1
tests/integration/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""PostgreSQL-backed integration tests."""
|
||||
174
tests/integration/conftest.py
Normal file
174
tests/integration/conftest.py
Normal file
@@ -0,0 +1,174 @@
|
||||
"""Isolated PostgreSQL and Flask fixtures for integration tests."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import re
|
||||
import uuid
|
||||
|
||||
import psycopg2
|
||||
from psycopg2 import sql
|
||||
import pytest
|
||||
from dotenv import load_dotenv
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
load_dotenv(PROJECT_ROOT / ".env", override=False)
|
||||
|
||||
TEST_DATABASE_PREFIX = "llm_bot_integration_"
|
||||
SERVICE_KEY = "integration-service-key-0123456789abcdef"
|
||||
ALLOWED_DISCORD_ID = "integration-discord-user"
|
||||
|
||||
|
||||
def _databaseConfig(databaseName=None):
|
||||
return {
|
||||
"host": os.environ.get("DB_HOST", "localhost"),
|
||||
"port": int(os.environ.get("DB_PORT", 5432)),
|
||||
"dbname": databaseName or os.environ.get("DB_NAME", "app"),
|
||||
"user": os.environ.get("DB_USER", "app"),
|
||||
"password": os.environ.get("DB_PASS", ""),
|
||||
"connect_timeout": 3,
|
||||
}
|
||||
|
||||
|
||||
def _databaseUnavailable(message):
|
||||
if os.environ.get("CI", "").lower() == "true":
|
||||
pytest.fail(message)
|
||||
pytest.skip(message)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def postgresServer():
|
||||
"""Skip once locally, but fail CI when its required PostgreSQL is absent."""
|
||||
sourceDatabase = os.environ.get("DB_NAME", "app")
|
||||
try:
|
||||
connection = psycopg2.connect(**_databaseConfig(sourceDatabase))
|
||||
except psycopg2.OperationalError as error:
|
||||
_databaseUnavailable(
|
||||
f"PostgreSQL integration database is unavailable: {error}"
|
||||
)
|
||||
|
||||
with connection:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT rolsuper OR rolcreatedb AS can_create_database
|
||||
FROM pg_roles
|
||||
WHERE rolname = CURRENT_USER
|
||||
"""
|
||||
)
|
||||
canCreateDatabase = cursor.fetchone()[0]
|
||||
connection.close()
|
||||
if not canCreateDatabase:
|
||||
_databaseUnavailable(
|
||||
"PostgreSQL integration user requires CREATEDB for isolated tests"
|
||||
)
|
||||
return sourceDatabase
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def isolatedDatabase(monkeypatch, postgresServer):
|
||||
"""Create a disposable database without altering the configured database."""
|
||||
databaseName = f"{TEST_DATABASE_PREFIX}{uuid.uuid4().hex}"
|
||||
assert re.fullmatch(r"llm_bot_integration_[0-9a-f]{32}", databaseName)
|
||||
|
||||
adminConnection = psycopg2.connect(**_databaseConfig(postgresServer))
|
||||
|
||||
adminConnection.autocommit = True
|
||||
try:
|
||||
with adminConnection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
sql.SQL("CREATE DATABASE {}").format(sql.Identifier(databaseName))
|
||||
)
|
||||
except psycopg2.Error as error:
|
||||
adminConnection.close()
|
||||
_databaseUnavailable(
|
||||
f"PostgreSQL user cannot create an isolated database: {error}"
|
||||
)
|
||||
|
||||
monkeypatch.setenv("DB_NAME", databaseName)
|
||||
try:
|
||||
yield databaseName
|
||||
finally:
|
||||
# Every production helper closes its connection, but terminate any failed-test
|
||||
# leftovers before dropping only the uniquely named test database.
|
||||
with adminConnection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT pg_terminate_backend(pid)
|
||||
FROM pg_stat_activity
|
||||
WHERE datname = %s AND pid != pg_backend_pid()
|
||||
""",
|
||||
(databaseName,),
|
||||
)
|
||||
cursor.execute(
|
||||
sql.SQL("DROP DATABASE {}").format(sql.Identifier(databaseName))
|
||||
)
|
||||
adminConnection.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def migratedDatabase(isolatedDatabase):
|
||||
from core.migrations import upgrade
|
||||
|
||||
applied = upgrade()
|
||||
assert applied
|
||||
return isolatedDatabase
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app(migratedDatabase, monkeypatch):
|
||||
monkeypatch.setenv("JWT_SECRET", "integration-only-jwt-secret")
|
||||
monkeypatch.setenv("BOT_API_KEY", SERVICE_KEY)
|
||||
monkeypatch.setenv(
|
||||
"BOT_API_KEY_SCOPES",
|
||||
"discord:session,outbox:claim,outbox:deliver",
|
||||
)
|
||||
monkeypatch.setenv("DISCORD_ENROLLMENT_MODE", "allowlist")
|
||||
monkeypatch.setenv("DISCORD_ALLOWLIST", ALLOWED_DISCORD_ID)
|
||||
monkeypatch.setenv("DEFAULT_TIMEZONE", "UTC")
|
||||
|
||||
from api.main import createApp
|
||||
|
||||
flaskApp = createApp()
|
||||
flaskApp.config.update(TESTING=True)
|
||||
return flaskApp
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(app):
|
||||
return app.test_client()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def createAuthenticatedUser(client):
|
||||
created = 0
|
||||
|
||||
def create(username=None, password="correct-horse-battery-staple"):
|
||||
nonlocal created
|
||||
created += 1
|
||||
username = username or f"integration-user-{created}-{uuid.uuid4().hex}"
|
||||
registration = client.post(
|
||||
"/api/register",
|
||||
json={"username": username, "password": password, "timezone": "UTC"},
|
||||
)
|
||||
assert registration.status_code == 201, registration.get_json()
|
||||
|
||||
login = client.post(
|
||||
"/api/login", json={"username": username, "password": password}
|
||||
)
|
||||
assert login.status_code == 200, login.get_json()
|
||||
token = login.get_json()["token"]
|
||||
|
||||
uuidResponse = client.get(
|
||||
f"/api/getUserUUID/{username}",
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
assert uuidResponse.status_code == 200, uuidResponse.get_json()
|
||||
return {
|
||||
"username": username,
|
||||
"password": password,
|
||||
"token": token,
|
||||
"user_uuid": uuidResponse.get_json(),
|
||||
}
|
||||
|
||||
return create
|
||||
264
tests/integration/test_api.py
Normal file
264
tests/integration/test_api.py
Normal file
@@ -0,0 +1,264 @@
|
||||
"""Flask API integration tests for authentication and route boundaries."""
|
||||
|
||||
import uuid
|
||||
|
||||
SERVICE_KEY = "integration-service-key-0123456789abcdef"
|
||||
ALLOWED_DISCORD_ID = "integration-discord-user"
|
||||
|
||||
|
||||
def _bearer(token):
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
def test_health_and_error_routes(client):
|
||||
assert client.get("/health/live").get_json() == {"status": "ok"}
|
||||
assert client.get("/health/ready").get_json() == {"status": "ready"}
|
||||
assert client.get("/health").get_json() == {"status": "ready"}
|
||||
|
||||
missing = client.get("/api/route-that-does-not-exist")
|
||||
assert missing.status_code == 404
|
||||
assert missing.get_json() == {"error": "not found"}
|
||||
|
||||
wrongMethod = client.get("/api/register")
|
||||
assert wrongMethod.status_code == 405
|
||||
assert wrongMethod.get_json() == {"error": "method not allowed"}
|
||||
|
||||
invalidJson = client.post(
|
||||
"/api/register", data="[]", content_type="application/json"
|
||||
)
|
||||
assert invalidJson.status_code == 400
|
||||
assert invalidJson.get_json() == {"error": "JSON object required"}
|
||||
|
||||
|
||||
def test_password_registration_keeps_values_parameterized_and_protected(client):
|
||||
from core import postgres
|
||||
|
||||
injectedUsername = "alice'; DROP TABLE users; --"
|
||||
injectedUUID = str(uuid.uuid4())
|
||||
response = client.post(
|
||||
"/api/register",
|
||||
json={
|
||||
"username": injectedUsername,
|
||||
"password": "long-enough-password",
|
||||
"timezone": "America/Chicago",
|
||||
"id": injectedUUID,
|
||||
"password_hashed": "attacker-controlled",
|
||||
"created_at": "1900-01-01T00:00:00Z",
|
||||
"unexpected_admin": True,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 201, response.get_json()
|
||||
assert postgres.table_exists("users")
|
||||
|
||||
record = postgres.select_one("users", {"username": injectedUsername})
|
||||
assert record is not None
|
||||
assert str(record["id"]) != injectedUUID
|
||||
assert bytes(record["password_hashed"]) != b"attacker-controlled"
|
||||
assert record["timezone"] == "America/Chicago"
|
||||
assert record["created_at"].year > 1900
|
||||
|
||||
login = client.post(
|
||||
"/api/login",
|
||||
json={"username": injectedUsername, "password": "long-enough-password"},
|
||||
)
|
||||
assert login.status_code == 200
|
||||
assert login.get_json()["token"]
|
||||
|
||||
|
||||
def test_jwt_user_key_creation_authentication_and_revocation(
|
||||
client, createAuthenticatedUser
|
||||
):
|
||||
from core import postgres
|
||||
|
||||
user = createAuthenticatedUser("api-key-owner")
|
||||
jwtHeaders = _bearer(user["token"])
|
||||
|
||||
profile = client.get(f"/api/user/{user['user_uuid']}", headers=jwtHeaders)
|
||||
assert profile.status_code == 200
|
||||
assert profile.get_json()["username"] == user["username"]
|
||||
assert "password_hashed" not in profile.get_json()
|
||||
|
||||
update = client.put(
|
||||
f"/api/user/{user['user_uuid']}",
|
||||
headers=jwtHeaders,
|
||||
json={"timezone": "America/Chicago", "username": "cannot-change"},
|
||||
)
|
||||
assert update.status_code == 200
|
||||
assert client.get(
|
||||
"/api/user/me/timezone", headers=jwtHeaders
|
||||
).get_json() == {"timezone": "America/Chicago"}
|
||||
|
||||
created = client.post(
|
||||
"/api/keys", headers=jwtHeaders, json={"name": "integration key"}
|
||||
)
|
||||
assert created.status_code == 201, created.get_json()
|
||||
keyRecord = created.get_json()
|
||||
secret = keyRecord["key"]
|
||||
assert secret.startswith("llmbot_user_")
|
||||
assert "key_hash" not in keyRecord
|
||||
storedKey = postgres.select_one("api_keys", {"id": keyRecord["id"]})
|
||||
assert storedKey["key_hash"] != secret
|
||||
assert secret not in {str(value) for value in storedKey.values()}
|
||||
|
||||
listed = client.get("/api/keys", headers=jwtHeaders)
|
||||
assert listed.status_code == 200
|
||||
listedKey = listed.get_json()["keys"][0]
|
||||
assert listedKey["id"] == keyRecord["id"]
|
||||
assert "key" not in listedKey
|
||||
assert "key_hash" not in listedKey
|
||||
|
||||
apiKeyHeaders = _bearer(secret)
|
||||
timezone = client.get("/api/user/me/timezone", headers=apiKeyHeaders)
|
||||
assert timezone.status_code == 200
|
||||
assert timezone.get_json() == {"timezone": "America/Chicago"}
|
||||
assert postgres.select_one("api_keys", {"id": keyRecord["id"]})[
|
||||
"last_used_at"
|
||||
] is not None
|
||||
|
||||
cannotManageKeys = client.post(
|
||||
"/api/keys", headers=apiKeyHeaders, json={"name": "nested key"}
|
||||
)
|
||||
assert cannotManageKeys.status_code == 401
|
||||
|
||||
revoked = client.delete(f"/api/keys/{keyRecord['id']}", headers=jwtHeaders)
|
||||
assert revoked.status_code == 200
|
||||
assert postgres.select_one("api_keys", {"id": keyRecord["id"]})[
|
||||
"revoked_at"
|
||||
] is not None
|
||||
assert client.get("/api/user/me/timezone", headers=apiKeyHeaders).status_code == 401
|
||||
|
||||
|
||||
def test_discord_allowlist_and_service_key_boundaries(client, createAuthenticatedUser):
|
||||
serviceHeaders = _bearer(SERVICE_KEY)
|
||||
|
||||
assert client.post(
|
||||
"/api/auth/discord/session",
|
||||
json={"discord_id": ALLOWED_DISCORD_ID},
|
||||
).status_code == 401
|
||||
assert client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer("not-the-configured-service-key-0123456789"),
|
||||
json={"discord_id": ALLOWED_DISCORD_ID},
|
||||
).status_code == 401
|
||||
|
||||
denied = client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=serviceHeaders,
|
||||
json={"discord_id": "not-allowlisted", "display_name": "No Access"},
|
||||
)
|
||||
assert denied.status_code == 403
|
||||
|
||||
accepted = client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=serviceHeaders,
|
||||
json={"discord_id": ALLOWED_DISCORD_ID, "display_name": "First Name"},
|
||||
)
|
||||
assert accepted.status_code == 200, accepted.get_json()
|
||||
discordSession = accepted.get_json()
|
||||
assert discordSession["timezone"] == "UTC"
|
||||
assert discordSession["token"]
|
||||
|
||||
repeated = client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=serviceHeaders,
|
||||
json={"discord_id": ALLOWED_DISCORD_ID, "display_name": "New Name"},
|
||||
)
|
||||
assert repeated.status_code == 200
|
||||
assert repeated.get_json()["user_uuid"] == discordSession["user_uuid"]
|
||||
|
||||
profile = client.get(
|
||||
f"/api/user/{discordSession['user_uuid']}",
|
||||
headers=_bearer(discordSession["token"]),
|
||||
)
|
||||
assert profile.status_code == 200
|
||||
assert profile.get_json()["username"] is None
|
||||
|
||||
passwordUser = createAuthenticatedUser("service-boundary-user")
|
||||
assert client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer(passwordUser["token"]),
|
||||
json={"discord_id": ALLOWED_DISCORD_ID},
|
||||
).status_code == 401
|
||||
assert client.get(
|
||||
"/api/user/me/timezone", headers=serviceHeaders
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_service_outbox_claim_retry_and_delivery_routes(client):
|
||||
from core import outbox, postgres
|
||||
|
||||
serviceHeaders = _bearer(SERVICE_KEY)
|
||||
session = client.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=serviceHeaders,
|
||||
json={"discord_id": ALLOWED_DISCORD_ID, "display_name": "Recipient"},
|
||||
).get_json()
|
||||
queued = outbox.enqueue_message(
|
||||
session["user_uuid"],
|
||||
"discord_dm",
|
||||
{"content": "integration delivery"},
|
||||
idempotency_key=f"api-outbox-{uuid.uuid4()}",
|
||||
)
|
||||
|
||||
claimed = client.post(
|
||||
"/api/internal/outbox/claim",
|
||||
headers=serviceHeaders,
|
||||
json={"worker_id": "api-worker-one", "limit": 1},
|
||||
)
|
||||
assert claimed.status_code == 200
|
||||
message = claimed.get_json()["messages"][0]
|
||||
assert message["id"] == str(queued["id"])
|
||||
assert message["provider_user_id"] == ALLOWED_DISCORD_ID
|
||||
assert message["content"] == "integration delivery"
|
||||
assert message["worker_id"] == "api-worker-one"
|
||||
|
||||
missingWorker = client.post(
|
||||
f"/api/internal/outbox/{message['id']}/result",
|
||||
headers=serviceHeaders,
|
||||
json={"status": "sent"},
|
||||
)
|
||||
assert missingWorker.status_code == 409
|
||||
|
||||
wrongWorker = client.post(
|
||||
f"/api/internal/outbox/{message['id']}/result",
|
||||
headers=serviceHeaders,
|
||||
json={"status": "sent", "worker_id": "not-the-lease-owner"},
|
||||
)
|
||||
assert wrongWorker.status_code == 409
|
||||
|
||||
stillLeased = outbox.get_message(message["id"])
|
||||
assert stillLeased["status"] == "delivering"
|
||||
assert stillLeased["leased_by"] == "api-worker-one"
|
||||
assert stillLeased["delivered_at"] is None
|
||||
|
||||
retried = client.post(
|
||||
f"/api/internal/outbox/{message['id']}/result",
|
||||
headers=serviceHeaders,
|
||||
json={
|
||||
"status": "retry",
|
||||
"error": "temporary failure",
|
||||
"worker_id": "api-worker-one",
|
||||
},
|
||||
)
|
||||
assert retried.status_code == 200
|
||||
assert retried.get_json()["message"]["status"] == "pending"
|
||||
|
||||
postgres.execute(
|
||||
"UPDATE outbound_messages SET available_at = NOW() WHERE id = %(id)s",
|
||||
{"id": message["id"]},
|
||||
)
|
||||
claimedAgain = client.post(
|
||||
"/api/internal/outbox/claim",
|
||||
headers=serviceHeaders,
|
||||
json={"worker_id": "api-worker-two", "limit": 1},
|
||||
)
|
||||
assert claimedAgain.status_code == 200
|
||||
assert claimedAgain.get_json()["messages"][0]["attempts"] == 2
|
||||
|
||||
delivered = client.post(
|
||||
f"/api/internal/outbox/{message['id']}/result",
|
||||
headers=serviceHeaders,
|
||||
json={"status": "sent", "worker_id": "api-worker-two"},
|
||||
)
|
||||
assert delivered.status_code == 200
|
||||
assert delivered.get_json()["message"]["status"] == "delivered"
|
||||
122
tests/integration/test_leases.py
Normal file
122
tests/integration/test_leases.py
Normal file
@@ -0,0 +1,122 @@
|
||||
"""Concurrent job and outbox lease behavior backed by PostgreSQL."""
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import threading
|
||||
import uuid
|
||||
|
||||
|
||||
def _claimConcurrently(claim, firstWorker, secondWorker):
|
||||
barrier = threading.Barrier(2)
|
||||
|
||||
def run(workerID):
|
||||
barrier.wait(timeout=10)
|
||||
return workerID, claim(workerID)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
futures = [
|
||||
executor.submit(run, firstWorker),
|
||||
executor.submit(run, secondWorker),
|
||||
]
|
||||
return dict(future.result(timeout=20) for future in futures)
|
||||
|
||||
|
||||
def test_job_claims_are_disjoint_and_support_retry_and_cancel(
|
||||
migratedDatabase
|
||||
):
|
||||
from core import jobs, users
|
||||
|
||||
users.registerUser("job-owner", "job-owner-password")
|
||||
userUUID = users.getUserUUID("job-owner")
|
||||
due = datetime.now(timezone.utc) - timedelta(minutes=1)
|
||||
created = [
|
||||
jobs.create_job(
|
||||
"integration.work",
|
||||
{"sequence": index},
|
||||
due,
|
||||
user_uuid=userUUID,
|
||||
idempotency_key=f"job-{uuid.uuid4()}",
|
||||
)
|
||||
for index in range(10)
|
||||
]
|
||||
|
||||
claims = _claimConcurrently(
|
||||
lambda worker: jobs.claim_due_jobs(worker, limit=5, lease_seconds=60),
|
||||
"job-worker-one",
|
||||
"job-worker-two",
|
||||
)
|
||||
firstIDs = {str(item["id"]) for item in claims["job-worker-one"]}
|
||||
secondIDs = {str(item["id"]) for item in claims["job-worker-two"]}
|
||||
assert firstIDs.isdisjoint(secondIDs)
|
||||
assert firstIDs | secondIDs == {str(item["id"]) for item in created}
|
||||
assert jobs.claim_due_jobs("job-worker-three", limit=10) == []
|
||||
|
||||
retriedID = str(claims["job-worker-one"][0]["id"])
|
||||
assert jobs.fail_job(retriedID, "wrong-worker", "must not update") is None
|
||||
retried = jobs.fail_job(
|
||||
retriedID, "job-worker-one", "temporary job failure", retry_seconds=1
|
||||
)
|
||||
assert retried["status"] == "pending"
|
||||
assert retried["attempts"] == 1
|
||||
assert retried["last_error"] == "temporary job failure"
|
||||
assert retried["run_at"] > datetime.now(timezone.utc)
|
||||
|
||||
cancelledID = str(claims["job-worker-two"][0]["id"])
|
||||
assert jobs.cancel_job(cancelledID, user_uuid=uuid.uuid4()) is None
|
||||
cancelled = jobs.cancel_job(cancelledID, user_uuid=userUUID)
|
||||
assert cancelled["status"] == "cancelled"
|
||||
assert cancelled["leased_by"] is None
|
||||
|
||||
|
||||
def test_outbox_claims_are_disjoint_and_support_retry_and_cancel(
|
||||
migratedDatabase
|
||||
):
|
||||
from core import outbox, users
|
||||
|
||||
users.registerUser("outbox-owner", "outbox-owner-password")
|
||||
userUUID = users.getUserUUID("outbox-owner")
|
||||
due = datetime.now(timezone.utc) - timedelta(minutes=1)
|
||||
created = [
|
||||
outbox.enqueue_message(
|
||||
userUUID,
|
||||
"discord_dm",
|
||||
{"content": f"message {index}"},
|
||||
idempotency_key=f"outbox-{uuid.uuid4()}",
|
||||
available_at=due,
|
||||
)
|
||||
for index in range(10)
|
||||
]
|
||||
|
||||
claims = _claimConcurrently(
|
||||
lambda worker: outbox.claim_messages(
|
||||
worker, channel="discord_dm", limit=5, lease_seconds=60
|
||||
),
|
||||
"outbox-worker-one",
|
||||
"outbox-worker-two",
|
||||
)
|
||||
firstIDs = {str(item["id"]) for item in claims["outbox-worker-one"]}
|
||||
secondIDs = {str(item["id"]) for item in claims["outbox-worker-two"]}
|
||||
assert firstIDs.isdisjoint(secondIDs)
|
||||
assert firstIDs | secondIDs == {str(item["id"]) for item in created}
|
||||
assert outbox.claim_messages("outbox-worker-three", limit=10) == []
|
||||
|
||||
retriedID = str(claims["outbox-worker-one"][0]["id"])
|
||||
assert outbox.retry_message(
|
||||
retriedID, "wrong-worker", "must not update"
|
||||
) is None
|
||||
retried = outbox.retry_message(
|
||||
retriedID,
|
||||
"outbox-worker-one",
|
||||
"temporary delivery failure",
|
||||
retry_seconds=1,
|
||||
)
|
||||
assert retried["status"] == "pending"
|
||||
assert retried["attempts"] == 1
|
||||
assert retried["last_error"] == "temporary delivery failure"
|
||||
assert retried["available_at"] > datetime.now(timezone.utc)
|
||||
|
||||
cancelledID = str(claims["outbox-worker-two"][0]["id"])
|
||||
assert outbox.cancel_message(cancelledID, user_uuid=uuid.uuid4()) is None
|
||||
cancelled = outbox.cancel_message(cancelledID, user_uuid=userUUID)
|
||||
assert cancelled["status"] == "cancelled"
|
||||
assert cancelled["leased_by"] is None
|
||||
70
tests/integration/test_migrations.py
Normal file
70
tests/integration/test_migrations.py
Normal file
@@ -0,0 +1,70 @@
|
||||
"""Migration integration coverage against a genuinely fresh database."""
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_fresh_upgrade_is_complete_idempotent_and_reported(isolatedDatabase, capsys):
|
||||
from core import postgres
|
||||
from core.migrations import (
|
||||
MigrationError,
|
||||
discover_migrations,
|
||||
migration_status,
|
||||
upgrade,
|
||||
)
|
||||
from core.migrations.__main__ import main
|
||||
|
||||
discovered = discover_migrations()
|
||||
assert discovered
|
||||
|
||||
applied = upgrade()
|
||||
assert [(item.namespace, item.version) for item in applied] == [
|
||||
(item.namespace, item.version) for item in discovered
|
||||
]
|
||||
|
||||
expectedTables = {
|
||||
"schema_migrations",
|
||||
"users",
|
||||
"notifications",
|
||||
"scheduled_jobs",
|
||||
"outbound_messages",
|
||||
"provider_identities",
|
||||
"api_keys",
|
||||
"reminders",
|
||||
}
|
||||
rows = postgres.execute(
|
||||
"""
|
||||
SELECT table_name
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public'
|
||||
"""
|
||||
)
|
||||
assert expectedTables.issubset({row["table_name"] for row in rows})
|
||||
|
||||
assert upgrade() == []
|
||||
status = migration_status()
|
||||
assert len(status) == len(discovered)
|
||||
assert {record["state"] for record in status} == {"applied"}
|
||||
assert all(record["applied_at"] is not None for record in status)
|
||||
|
||||
assert main(["status"]) == 0
|
||||
output = capsys.readouterr().out
|
||||
assert "applied" in output
|
||||
assert "core:1" in output
|
||||
assert "reminders:1" in output
|
||||
|
||||
postgres.execute(
|
||||
"""
|
||||
UPDATE schema_migrations
|
||||
SET checksum = %(checksum)s
|
||||
WHERE namespace = 'core' AND version = 1
|
||||
""",
|
||||
{"checksum": "0" * 64},
|
||||
)
|
||||
changed = migration_status()
|
||||
coreBaseline = next(
|
||||
row for row in changed if row["namespace"] == "core" and row["version"] == 1
|
||||
)
|
||||
assert coreBaseline["state"] == "changed"
|
||||
assert main(["status"]) == 1
|
||||
with pytest.raises(MigrationError, match="checksum changed"):
|
||||
upgrade()
|
||||
107
tests/integration/test_reminders.py
Normal file
107
tests/integration/test_reminders.py
Normal file
@@ -0,0 +1,107 @@
|
||||
"""Reminder API ownership and transactional cancellation coverage."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import uuid
|
||||
|
||||
|
||||
def _bearer(token):
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
def test_reminders_are_owned_and_cancel_related_work(
|
||||
client, createAuthenticatedUser
|
||||
):
|
||||
from core import outbox, postgres
|
||||
|
||||
owner = createAuthenticatedUser("reminder-owner")
|
||||
stranger = createAuthenticatedUser("reminder-stranger")
|
||||
runAt = datetime.now(timezone.utc) + timedelta(days=1)
|
||||
|
||||
created = client.post(
|
||||
"/api/reminders",
|
||||
headers=_bearer(owner["token"]),
|
||||
json={
|
||||
"message": "renew the certificate",
|
||||
"run_at": runAt.isoformat(),
|
||||
"timezone": "America/Chicago",
|
||||
"recurrence": {"frequency": "weekly", "interval": 1},
|
||||
},
|
||||
)
|
||||
assert created.status_code == 201, created.get_json()
|
||||
reminder = created.get_json()
|
||||
|
||||
ownerList = client.get(
|
||||
"/api/reminders", headers=_bearer(owner["token"])
|
||||
).get_json()["reminders"]
|
||||
assert [item["id"] for item in ownerList] == [reminder["id"]]
|
||||
assert client.get(
|
||||
"/api/reminders", headers=_bearer(stranger["token"])
|
||||
).get_json() == {"reminders": []}
|
||||
|
||||
queued = outbox.enqueue_message(
|
||||
owner["user_uuid"],
|
||||
"discord_dm",
|
||||
{"content": "not delivered yet", "reminder_id": reminder["id"]},
|
||||
idempotency_key=f"pending-reminder-{uuid.uuid4()}",
|
||||
)
|
||||
|
||||
strangerCancel = client.delete(
|
||||
f"/api/reminders/{reminder['id']}",
|
||||
headers=_bearer(stranger["token"]),
|
||||
)
|
||||
assert strangerCancel.status_code == 404
|
||||
storedReminder = postgres.select_one("reminders", {"id": reminder["id"]})
|
||||
assert storedReminder["status"] == "active"
|
||||
assert outbox.get_message(queued["id"])["status"] == "pending"
|
||||
|
||||
ownerCancel = client.delete(
|
||||
f"/api/reminders/{reminder['id']}", headers=_bearer(owner["token"])
|
||||
)
|
||||
assert ownerCancel.status_code == 200
|
||||
assert ownerCancel.get_json()["status"] == "cancelled"
|
||||
|
||||
job = postgres.execute(
|
||||
"""
|
||||
SELECT * FROM scheduled_jobs
|
||||
WHERE payload->>'reminder_id' = %(reminder_id)s
|
||||
""",
|
||||
{"reminder_id": reminder["id"]},
|
||||
)[0]
|
||||
assert job["user_uuid"] == uuid.UUID(owner["user_uuid"])
|
||||
assert job["status"] == "cancelled"
|
||||
assert outbox.get_message(queued["id"])["status"] == "cancelled"
|
||||
|
||||
assert client.delete(
|
||||
f"/api/reminders/{reminder['id']}", headers=_bearer(owner["token"])
|
||||
).status_code == 404
|
||||
|
||||
|
||||
def test_reminder_routes_validate_time_and_authentication(
|
||||
client, createAuthenticatedUser
|
||||
):
|
||||
user = createAuthenticatedUser("reminder-validation")
|
||||
headers = _bearer(user["token"])
|
||||
|
||||
assert client.get("/api/reminders").status_code == 401
|
||||
past = client.post(
|
||||
"/api/reminders",
|
||||
headers=headers,
|
||||
json={
|
||||
"message": "too late",
|
||||
"run_at": (datetime.now(timezone.utc) - timedelta(minutes=1)).isoformat(),
|
||||
"timezone": "UTC",
|
||||
},
|
||||
)
|
||||
assert past.status_code == 400
|
||||
assert "future" in past.get_json()["error"]
|
||||
|
||||
badTimezone = client.post(
|
||||
"/api/reminders",
|
||||
headers=headers,
|
||||
json={
|
||||
"message": "bad timezone",
|
||||
"run_at": (datetime.now(timezone.utc) + timedelta(hours=1)).isoformat(),
|
||||
"timezone": "UTC'; DROP TABLE reminders; --",
|
||||
},
|
||||
)
|
||||
assert badTimezone.status_code == 400
|
||||
77
tests/unit/conftest.py
Normal file
77
tests/unit/conftest.py
Normal file
@@ -0,0 +1,77 @@
|
||||
"""Small import fallbacks for running focused tests without optional services."""
|
||||
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
try:
|
||||
import psycopg2 # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
psycopg2 = types.ModuleType("psycopg2")
|
||||
extras = types.ModuleType("psycopg2.extras")
|
||||
errors = types.ModuleType("psycopg2.errors")
|
||||
|
||||
class Json:
|
||||
def __init__(self, adapted):
|
||||
self.adapted = adapted
|
||||
|
||||
class RealDictCursor:
|
||||
pass
|
||||
|
||||
class UniqueViolation(Exception):
|
||||
pass
|
||||
|
||||
def connect(**_kwargs):
|
||||
raise AssertionError("Tests must mock PostgreSQL connections")
|
||||
|
||||
def execute_values(*_args, **_kwargs):
|
||||
raise AssertionError("Tests must mock bulk PostgreSQL writes")
|
||||
|
||||
extras.Json = Json
|
||||
extras.RealDictCursor = RealDictCursor
|
||||
extras.execute_values = execute_values
|
||||
errors.UniqueViolation = UniqueViolation
|
||||
psycopg2.connect = connect
|
||||
psycopg2.extras = extras
|
||||
psycopg2.errors = errors
|
||||
sys.modules["psycopg2"] = psycopg2
|
||||
sys.modules["psycopg2.extras"] = extras
|
||||
sys.modules["psycopg2.errors"] = errors
|
||||
|
||||
|
||||
try:
|
||||
import bcrypt # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
bcrypt = types.ModuleType("bcrypt")
|
||||
|
||||
def unavailable(*_args, **_kwargs):
|
||||
raise AssertionError("Tests exercising bcrypt require project dependencies")
|
||||
|
||||
bcrypt.gensalt = unavailable
|
||||
bcrypt.hashpw = unavailable
|
||||
bcrypt.checkpw = unavailable
|
||||
sys.modules["bcrypt"] = bcrypt
|
||||
|
||||
|
||||
try:
|
||||
import jwt # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
jwt = types.ModuleType("jwt")
|
||||
exceptions = types.ModuleType("jwt.exceptions")
|
||||
|
||||
class ExpiredSignatureError(Exception):
|
||||
pass
|
||||
|
||||
class InvalidTokenError(Exception):
|
||||
pass
|
||||
|
||||
def unavailableJwt(*_args, **_kwargs):
|
||||
raise AssertionError("Tests exercising JWT encoding require project dependencies")
|
||||
|
||||
jwt.encode = unavailableJwt
|
||||
jwt.decode = unavailableJwt
|
||||
exceptions.ExpiredSignatureError = ExpiredSignatureError
|
||||
exceptions.InvalidTokenError = InvalidTokenError
|
||||
jwt.exceptions = exceptions
|
||||
sys.modules["jwt"] = jwt
|
||||
sys.modules["jwt.exceptions"] = exceptions
|
||||
310
tests/unit/test_adapters.py
Normal file
310
tests/unit/test_adapters.py
Normal file
@@ -0,0 +1,310 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, call
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from bot import api_client as apiClientModule
|
||||
from bot.context import CommandContext
|
||||
from scheduler import daemon
|
||||
|
||||
|
||||
def _response(status, payload=None, jsonError=False):
|
||||
response = MagicMock(status_code=status)
|
||||
if jsonError:
|
||||
response.json.side_effect = ValueError("not JSON")
|
||||
else:
|
||||
response.json.return_value = payload
|
||||
return response
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def apiClient(monkeypatch):
|
||||
transport = SimpleNamespace(
|
||||
post=AsyncMock(),
|
||||
request=AsyncMock(),
|
||||
aclose=AsyncMock(),
|
||||
)
|
||||
clientFactory = MagicMock(return_value=transport)
|
||||
monkeypatch.setattr(apiClientModule.httpx, "AsyncClient", clientFactory)
|
||||
monkeypatch.setenv("API_URL", "http://api.test/")
|
||||
monkeypatch.setenv("BOT_API_KEY", "service-secret")
|
||||
monkeypatch.setenv("DEFAULT_TIMEZONE", "UTC")
|
||||
|
||||
client = apiClientModule.ApiClient("123", "Test User")
|
||||
|
||||
clientFactory.assert_called_once_with(timeout=10.0)
|
||||
return client, transport
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_authenticate_stores_session(apiClient):
|
||||
client, transport = apiClient
|
||||
transport.post.return_value = _response(
|
||||
200,
|
||||
{
|
||||
"token": "user-token",
|
||||
"user_uuid": "user-1",
|
||||
"timezone": "America/Chicago",
|
||||
},
|
||||
)
|
||||
|
||||
result, status = await client.authenticate()
|
||||
|
||||
assert status == 200
|
||||
assert result["user_uuid"] == "user-1"
|
||||
assert client.token == "user-token"
|
||||
assert client.user_uuid == "user-1"
|
||||
assert client.timezone == "America/Chicago"
|
||||
transport.post.assert_awaited_once_with(
|
||||
"http://api.test/api/auth/discord/session",
|
||||
headers={"Authorization": "Bearer service-secret"},
|
||||
json={"discord_id": "123", "display_name": "Test User"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_authentication_error_stops_user_request(apiClient):
|
||||
client, transport = apiClient
|
||||
transport.post.side_effect = httpx.ConnectError("offline")
|
||||
|
||||
result, status = await client.authenticate()
|
||||
|
||||
assert (result, status) == ({"error": "API unavailable"}, 503)
|
||||
assert client.token is None
|
||||
|
||||
result, status = await client.request("get", "/api/reminders")
|
||||
|
||||
assert (result, status) == ({"error": "authentication failed"}, 503)
|
||||
transport.request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_request_sends_user_token_and_payload(apiClient):
|
||||
client, transport = apiClient
|
||||
client.token = "user-token"
|
||||
transport.request.return_value = _response(201, {"id": "item-1"})
|
||||
|
||||
result, status = await client.request(
|
||||
"post",
|
||||
"/api/items",
|
||||
{"name": "example"},
|
||||
params={"source": "test"},
|
||||
)
|
||||
|
||||
assert (result, status) == ({"id": "item-1"}, 201)
|
||||
transport.request.assert_awaited_once_with(
|
||||
"POST",
|
||||
"http://api.test/api/items",
|
||||
headers={"Authorization": "Bearer user-token"},
|
||||
params={"source": "test"},
|
||||
json={"name": "example"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_request_returns_safe_transport_error(apiClient):
|
||||
client, transport = apiClient
|
||||
client.token = "user-token"
|
||||
transport.request.side_effect = httpx.ReadTimeout("timed out")
|
||||
|
||||
assert await client.request("get", "/api/items") == (
|
||||
{"error": "API unavailable"},
|
||||
503,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_reauthenticates_once_after_401(apiClient):
|
||||
client, transport = apiClient
|
||||
client.token = "expired-token"
|
||||
transport.request.side_effect = [
|
||||
_response(401, {"error": "unauthorized"}),
|
||||
_response(200, {"items": [1]}),
|
||||
]
|
||||
transport.post.return_value = _response(
|
||||
200,
|
||||
{"token": "fresh-token", "user_uuid": "user-1", "timezone": "UTC"},
|
||||
)
|
||||
|
||||
result, status = await client.request("get", "/api/items", params={"page": 2})
|
||||
|
||||
assert (result, status) == ({"items": [1]}, 200)
|
||||
assert client.token == "fresh-token"
|
||||
assert transport.request.await_args_list == [
|
||||
call(
|
||||
"GET",
|
||||
"http://api.test/api/items",
|
||||
headers={"Authorization": "Bearer expired-token"},
|
||||
params={"page": 2},
|
||||
),
|
||||
call(
|
||||
"GET",
|
||||
"http://api.test/api/items",
|
||||
headers={"Authorization": "Bearer fresh-token"},
|
||||
params={"page": 2},
|
||||
),
|
||||
]
|
||||
transport.post.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_returns_failed_refresh_result_after_401(apiClient):
|
||||
client, transport = apiClient
|
||||
client.token = "expired-token"
|
||||
transport.request.return_value = _response(401, {"error": "unauthorized"})
|
||||
transport.post.return_value = _response(403, {"error": "not enrolled"})
|
||||
|
||||
assert await client.request("get", "/api/items") == (
|
||||
{"error": "not enrolled"},
|
||||
403,
|
||||
)
|
||||
transport.request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_service_request_uses_service_key(apiClient):
|
||||
client, transport = apiClient
|
||||
transport.request.return_value = _response(200, {"messages": []})
|
||||
|
||||
result, status = await client.service_request(
|
||||
"post",
|
||||
"/api/internal/outbox/claim",
|
||||
{"worker_id": "worker-1"},
|
||||
)
|
||||
|
||||
assert (result, status) == ({"messages": []}, 200)
|
||||
transport.request.assert_awaited_once_with(
|
||||
"POST",
|
||||
"http://api.test/api/internal/outbox/claim",
|
||||
headers={"Authorization": "Bearer service-secret"},
|
||||
params=None,
|
||||
json={"worker_id": "worker-1"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_client_handles_non_json_response_and_closes(apiClient):
|
||||
client, transport = apiClient
|
||||
transport.request.return_value = _response(502, jsonError=True)
|
||||
|
||||
assert await client.service_request("get", "/bad-response") == ({}, 502)
|
||||
|
||||
await client.close()
|
||||
transport.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_context_exposes_platform_neutral_values_and_replies():
|
||||
channel = SimpleNamespace(send=AsyncMock(return_value="sent-message"))
|
||||
message = SimpleNamespace(author=SimpleNamespace(id=987), channel=channel)
|
||||
api = SimpleNamespace(user_uuid="user-1", timezone="America/Chicago")
|
||||
|
||||
context = CommandContext(message, api)
|
||||
|
||||
assert context.api is api
|
||||
assert context.user_uuid == "user-1"
|
||||
assert context.discord_user_id == "987"
|
||||
assert context.timezone == "America/Chicago"
|
||||
assert await context.reply("hello") == "sent-message"
|
||||
channel.send.assert_awaited_once_with("hello")
|
||||
|
||||
|
||||
def test_scheduler_fails_job_with_unknown_type(monkeypatch):
|
||||
getHandler = MagicMock(return_value=None)
|
||||
failJob = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
daemon,
|
||||
"module_registry",
|
||||
SimpleNamespace(get_job_handler=getHandler),
|
||||
)
|
||||
monkeypatch.setattr(daemon.jobs, "fail_job", failJob)
|
||||
|
||||
daemon.runJob({"id": "job-1", "job_type": "missing"})
|
||||
|
||||
getHandler.assert_called_once_with("missing")
|
||||
failJob.assert_called_once_with(
|
||||
"job-1",
|
||||
daemon.WORKER_ID,
|
||||
"unknown job type: missing",
|
||||
)
|
||||
|
||||
|
||||
def test_scheduler_runs_handler_and_completes_running_job(monkeypatch):
|
||||
job = {"id": "job-1", "job_type": "example"}
|
||||
handler = MagicMock(return_value=None)
|
||||
completeJob = MagicMock()
|
||||
failJob = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
daemon,
|
||||
"module_registry",
|
||||
SimpleNamespace(get_job_handler=MagicMock(return_value=handler)),
|
||||
)
|
||||
monkeypatch.setattr(daemon.jobs, "get_job", MagicMock(return_value={"status": "running"}))
|
||||
monkeypatch.setattr(daemon.jobs, "complete_job", completeJob)
|
||||
monkeypatch.setattr(daemon.jobs, "fail_job", failJob)
|
||||
|
||||
daemon.runJob(job)
|
||||
|
||||
handler.assert_called_once_with(job, daemon.WORKER_ID)
|
||||
completeJob.assert_called_once_with("job-1", daemon.WORKER_ID)
|
||||
failJob.assert_not_called()
|
||||
|
||||
|
||||
def test_scheduler_awaits_handler_without_double_completion(monkeypatch):
|
||||
job = {"id": "job-1", "job_type": "async-example"}
|
||||
handler = AsyncMock(return_value=None)
|
||||
completeJob = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
daemon,
|
||||
"module_registry",
|
||||
SimpleNamespace(get_job_handler=MagicMock(return_value=handler)),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
daemon.jobs,
|
||||
"get_job",
|
||||
MagicMock(return_value={"status": "completed"}),
|
||||
)
|
||||
monkeypatch.setattr(daemon.jobs, "complete_job", completeJob)
|
||||
|
||||
daemon.runJob(job)
|
||||
|
||||
handler.assert_awaited_once_with(job, daemon.WORKER_ID)
|
||||
completeJob.assert_not_called()
|
||||
|
||||
|
||||
def test_scheduler_records_handler_failure(monkeypatch):
|
||||
job = {"id": "job-1", "job_type": "broken"}
|
||||
failure = RuntimeError("handler failed")
|
||||
handler = MagicMock(side_effect=failure)
|
||||
failJob = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
daemon,
|
||||
"module_registry",
|
||||
SimpleNamespace(get_job_handler=MagicMock(return_value=handler)),
|
||||
)
|
||||
monkeypatch.setattr(daemon.jobs, "fail_job", failJob)
|
||||
monkeypatch.setattr(daemon.logger, "exception", MagicMock())
|
||||
|
||||
daemon.runJob(job)
|
||||
|
||||
failJob.assert_called_once_with("job-1", daemon.WORKER_ID, failure)
|
||||
|
||||
|
||||
def test_scheduler_poll_claims_configured_batch_and_runs_each_job(monkeypatch):
|
||||
claimed = [
|
||||
{"id": "job-1", "job_type": "one"},
|
||||
{"id": "job-2", "job_type": "two"},
|
||||
]
|
||||
claimJobs = MagicMock(return_value=claimed)
|
||||
runJob = MagicMock()
|
||||
monkeypatch.setattr(daemon.jobs, "claim_due_jobs", claimJobs)
|
||||
monkeypatch.setattr(daemon, "runJob", runJob)
|
||||
|
||||
assert daemon.pollJobs() == 2
|
||||
claimJobs.assert_called_once_with(
|
||||
daemon.WORKER_ID,
|
||||
limit=daemon.JOB_BATCH_SIZE,
|
||||
lease_seconds=daemon.JOB_LEASE_SECONDS,
|
||||
)
|
||||
assert runJob.call_args_list == [call(claimed[0]), call(claimed[1])]
|
||||
592
tests/unit/test_api_routes.py
Normal file
592
tests/unit/test_api_routes.py
Normal file
@@ -0,0 +1,592 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, call
|
||||
|
||||
import flask
|
||||
import pytest
|
||||
|
||||
from api import main as apiMain
|
||||
from api import security
|
||||
|
||||
|
||||
def _bearer(token="token"):
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
def _userPrincipal(userUUID="user-1", authentication="jwt"):
|
||||
return {
|
||||
"type": "user",
|
||||
"authentication": authentication,
|
||||
"user_uuid": userUUID,
|
||||
"scopes": [],
|
||||
}
|
||||
|
||||
|
||||
def _servicePrincipal(*scopes):
|
||||
return {
|
||||
"type": "service",
|
||||
"authentication": "api_key",
|
||||
"service_name": "test-service",
|
||||
"scopes": list(scopes),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def securityClient():
|
||||
app = flask.Flask("security-test")
|
||||
app.config["TESTING"] = True
|
||||
|
||||
@app.route("/json", methods=["POST"])
|
||||
def parseJson():
|
||||
return flask.jsonify({"parsed": security.jsonObject()})
|
||||
|
||||
@app.route("/user")
|
||||
@security.requireUser()
|
||||
def userRoute():
|
||||
return flask.jsonify(
|
||||
{
|
||||
"user_uuid": flask.g.user_uuid,
|
||||
"principal_type": flask.g.principal["type"],
|
||||
}
|
||||
)
|
||||
|
||||
@app.route("/login-user")
|
||||
@security.requireUser(requireLogin=True)
|
||||
def loginUserRoute():
|
||||
return flask.jsonify({"ok": True})
|
||||
|
||||
@app.route("/service")
|
||||
@security.requireService("jobs:claim")
|
||||
def serviceRoute():
|
||||
return flask.jsonify({"service": flask.g.principal["service_name"]})
|
||||
|
||||
return app.test_client()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def apiClient(monkeypatch):
|
||||
registry = SimpleNamespace(route_registrars=[])
|
||||
monkeypatch.setattr(apiMain, "discover_modules", MagicMock(return_value=registry))
|
||||
app = apiMain.createApp()
|
||||
app.config.update(TESTING=True, SERVICE_KEY_READY=True)
|
||||
return app.test_client()
|
||||
|
||||
|
||||
def test_json_object_accepts_only_json_objects(securityClient):
|
||||
assert securityClient.post("/json", json={"value": 1}).get_json() == {
|
||||
"parsed": {"value": 1}
|
||||
}
|
||||
assert securityClient.post("/json", json=[1, 2]).get_json() == {"parsed": None}
|
||||
assert securityClient.post(
|
||||
"/json",
|
||||
data="not-json",
|
||||
content_type="application/json",
|
||||
).get_json() == {"parsed": None}
|
||||
|
||||
|
||||
def test_require_user_sets_context_and_forwards_login_requirement(
|
||||
monkeypatch, securityClient
|
||||
):
|
||||
principal = _userPrincipal()
|
||||
authenticate = MagicMock(return_value=principal)
|
||||
isUser = MagicMock(return_value=True)
|
||||
monkeypatch.setattr(security.auth, "authenticateBearerToken", authenticate)
|
||||
monkeypatch.setattr(security.auth, "isUserPrincipal", isUser)
|
||||
|
||||
response = securityClient.get("/user", headers=_bearer("user-token"))
|
||||
loginResponse = securityClient.get(
|
||||
"/login-user",
|
||||
headers=_bearer("user-token"),
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.get_json() == {
|
||||
"user_uuid": "user-1",
|
||||
"principal_type": "user",
|
||||
}
|
||||
assert loginResponse.status_code == 200
|
||||
assert authenticate.call_args_list == [
|
||||
call("Bearer user-token", allowService=False),
|
||||
call("Bearer user-token", allowService=False),
|
||||
]
|
||||
assert isUser.call_args_list == [
|
||||
call(principal, requireLogin=False),
|
||||
call(principal, requireLogin=True),
|
||||
]
|
||||
|
||||
|
||||
def test_require_user_rejects_invalid_principal(monkeypatch, securityClient):
|
||||
monkeypatch.setattr(
|
||||
security.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
security.auth,
|
||||
"isUserPrincipal",
|
||||
MagicMock(return_value=False),
|
||||
)
|
||||
|
||||
response = securityClient.get("/user")
|
||||
|
||||
assert response.status_code == 401
|
||||
assert response.get_json() == {"error": "unauthorized"}
|
||||
|
||||
|
||||
def test_require_service_enforces_scope_and_sets_principal(
|
||||
monkeypatch, securityClient
|
||||
):
|
||||
principal = _servicePrincipal("jobs:claim")
|
||||
authenticate = MagicMock(side_effect=[principal, None])
|
||||
hasScope = MagicMock(side_effect=[True, False])
|
||||
monkeypatch.setattr(security.auth, "authenticateBearerToken", authenticate)
|
||||
monkeypatch.setattr(security.auth, "hasServiceScope", hasScope)
|
||||
|
||||
accepted = securityClient.get("/service", headers=_bearer("service-key"))
|
||||
rejected = securityClient.get("/service", headers=_bearer("wrong-key"))
|
||||
|
||||
assert accepted.status_code == 200
|
||||
assert accepted.get_json() == {"service": "test-service"}
|
||||
assert rejected.status_code == 401
|
||||
assert rejected.get_json() == {"error": "unauthorized"}
|
||||
assert authenticate.call_args_list == [
|
||||
call(
|
||||
"Bearer service-key",
|
||||
requiredScopes=["jobs:claim"],
|
||||
allowUser=False,
|
||||
),
|
||||
call(
|
||||
"Bearer wrong-key",
|
||||
requiredScopes=["jobs:claim"],
|
||||
allowUser=False,
|
||||
),
|
||||
]
|
||||
assert hasScope.call_args_list == [
|
||||
call(principal, "jobs:claim"),
|
||||
call(None, "jobs:claim"),
|
||||
]
|
||||
|
||||
|
||||
def test_registration_and_login_success_and_errors(monkeypatch, apiClient):
|
||||
registerUser = MagicMock(
|
||||
side_effect=[True, False, ValueError("invalid registration")]
|
||||
)
|
||||
getToken = MagicMock(side_effect=["login-token", False])
|
||||
monkeypatch.setattr(apiMain.users, "registerUser", registerUser)
|
||||
monkeypatch.setattr(apiMain.auth, "getLoginToken", getToken)
|
||||
|
||||
registered = apiClient.post(
|
||||
"/api/register",
|
||||
json={"username": "alice", "password": "password123", "timezone": "UTC"},
|
||||
)
|
||||
duplicate = apiClient.post(
|
||||
"/api/register",
|
||||
json={"username": "alice", "password": "password123"},
|
||||
)
|
||||
invalid = apiClient.post(
|
||||
"/api/register",
|
||||
json={"username": "", "password": "password123"},
|
||||
)
|
||||
invalidJson = apiClient.post("/api/register", json=["not", "an", "object"])
|
||||
loggedIn = apiClient.post(
|
||||
"/api/login",
|
||||
json={"username": "alice", "password": "password123"},
|
||||
)
|
||||
denied = apiClient.post(
|
||||
"/api/login",
|
||||
json={"username": "alice", "password": "wrong-password"},
|
||||
)
|
||||
|
||||
assert registered.status_code == 201
|
||||
assert registered.get_json() == {"success": True}
|
||||
assert duplicate.status_code == 409
|
||||
assert duplicate.get_json() == {"error": "username taken"}
|
||||
assert invalid.status_code == 400
|
||||
assert invalid.get_json() == {"error": "invalid registration"}
|
||||
assert invalidJson.status_code == 400
|
||||
assert loggedIn.get_json() == {"token": "login-token"}
|
||||
assert denied.status_code == 401
|
||||
assert denied.get_json() == {"error": "invalid credentials"}
|
||||
|
||||
|
||||
def test_discord_session_success_and_enrollment_errors(monkeypatch, apiClient):
|
||||
principal = _servicePrincipal("discord:session")
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=principal),
|
||||
)
|
||||
getOrCreate = MagicMock(
|
||||
side_effect=[ValueError("invalid Discord ID"), None, {"id": "user-1"}]
|
||||
)
|
||||
monkeypatch.setattr(apiMain.identity, "getOrCreateDiscordUser", getOrCreate)
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"createLoginToken",
|
||||
MagicMock(return_value="discord-token"),
|
||||
)
|
||||
|
||||
missing = apiClient.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer("service-key"),
|
||||
json={},
|
||||
)
|
||||
invalid = apiClient.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer("service-key"),
|
||||
json={"discord_id": "bad"},
|
||||
)
|
||||
denied = apiClient.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer("service-key"),
|
||||
json={"discord_id": "456"},
|
||||
)
|
||||
accepted = apiClient.post(
|
||||
"/api/auth/discord/session",
|
||||
headers=_bearer("service-key"),
|
||||
json={"discord_id": "123", "display_name": "Alice"},
|
||||
)
|
||||
|
||||
assert missing.status_code == 400
|
||||
assert invalid.get_json() == {"error": "invalid Discord ID"}
|
||||
assert denied.status_code == 403
|
||||
assert accepted.status_code == 200
|
||||
assert accepted.get_json() == {
|
||||
"token": "discord-token",
|
||||
"user_uuid": "user-1",
|
||||
"timezone": "UTC",
|
||||
}
|
||||
apiMain.auth.createLoginToken.assert_called_once_with(
|
||||
"user-1",
|
||||
name="Alice",
|
||||
extraClaims={"provider": "discord"},
|
||||
)
|
||||
|
||||
|
||||
def test_api_key_routes_use_authenticated_owner(monkeypatch, apiClient):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_userPrincipal()),
|
||||
)
|
||||
listKeys = MagicMock(return_value=[{"id": "key-1", "name": "CLI"}])
|
||||
createKey = MagicMock(
|
||||
side_effect=[
|
||||
{"id": "key-2", "name": "new", "key": "secret"},
|
||||
ValueError("invalid expiry"),
|
||||
]
|
||||
)
|
||||
revokeKey = MagicMock(side_effect=[False, True])
|
||||
monkeypatch.setattr(apiMain.apiKeys, "listUserApiKeys", listKeys)
|
||||
monkeypatch.setattr(apiMain.apiKeys, "createUserApiKey", createKey)
|
||||
monkeypatch.setattr(apiMain.apiKeys, "revokeUserApiKey", revokeKey)
|
||||
|
||||
listed = apiClient.get("/api/keys", headers=_bearer())
|
||||
created = apiClient.post(
|
||||
"/api/keys",
|
||||
headers=_bearer(),
|
||||
json={"name": "new", "expires_at": "2099-01-01T00:00:00Z"},
|
||||
)
|
||||
invalid = apiClient.post(
|
||||
"/api/keys",
|
||||
headers=_bearer(),
|
||||
json={"name": "bad"},
|
||||
)
|
||||
missing = apiClient.delete("/api/keys/missing", headers=_bearer())
|
||||
revoked = apiClient.delete("/api/keys/key-1", headers=_bearer())
|
||||
|
||||
assert listed.get_json() == {"keys": [{"id": "key-1", "name": "CLI"}]}
|
||||
assert created.status_code == 201
|
||||
assert created.get_json()["key"] == "secret"
|
||||
assert invalid.status_code == 400
|
||||
assert invalid.get_json() == {"error": "invalid expiry"}
|
||||
assert missing.status_code == 404
|
||||
assert revoked.get_json() == {"success": True}
|
||||
listKeys.assert_called_once_with("user-1")
|
||||
createKey.assert_has_calls(
|
||||
[
|
||||
call("user-1", "new", expiresAt="2099-01-01T00:00:00Z"),
|
||||
call("user-1", "bad", expiresAt=None),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_user_profile_and_username_routes_enforce_ownership(monkeypatch, apiClient):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_userPrincipal()),
|
||||
)
|
||||
getUser = MagicMock(
|
||||
side_effect=[
|
||||
{
|
||||
"id": "user-1",
|
||||
"username": "alice",
|
||||
"password_hashed": b"secret-hash",
|
||||
},
|
||||
None,
|
||||
]
|
||||
)
|
||||
getUserUUID = MagicMock(side_effect=["user-1", "user-2", False])
|
||||
monkeypatch.setattr(apiMain.users, "getUser", getUser)
|
||||
monkeypatch.setattr(apiMain.users, "getUserUUID", getUserUUID)
|
||||
|
||||
foreign = apiClient.get("/api/user/user-2", headers=_bearer())
|
||||
own = apiClient.get("/api/user/user-1", headers=_bearer())
|
||||
missing = apiClient.get("/api/user/user-1", headers=_bearer())
|
||||
username = apiClient.get("/api/getUserUUID/alice", headers=_bearer())
|
||||
foreignUsername = apiClient.get("/api/getUserUUID/bob", headers=_bearer())
|
||||
missingUsername = apiClient.get("/api/getUserUUID/missing", headers=_bearer())
|
||||
|
||||
assert foreign.status_code == 403
|
||||
assert getUser.call_count == 2
|
||||
assert own.status_code == 200
|
||||
assert own.get_json() == {"id": "user-1", "username": "alice"}
|
||||
assert missing.status_code == 404
|
||||
assert username.get_json() == "user-1"
|
||||
assert foreignUsername.status_code == 403
|
||||
assert missingUsername.status_code == 404
|
||||
|
||||
|
||||
def test_user_update_and_delete_success_and_errors(monkeypatch, apiClient):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_userPrincipal()),
|
||||
)
|
||||
updateUser = MagicMock(side_effect=[ValueError("bad timezone"), False, True])
|
||||
unregisterUser = MagicMock(side_effect=[False, True])
|
||||
monkeypatch.setattr(apiMain.users, "updateUser", updateUser)
|
||||
monkeypatch.setattr(apiMain.auth, "unregisterUser", unregisterUser)
|
||||
|
||||
assert apiClient.put(
|
||||
"/api/user/user-2", headers=_bearer(), json={"timezone": "UTC"}
|
||||
).status_code == 403
|
||||
assert apiClient.put(
|
||||
"/api/user/user-1", headers=_bearer(), json=[]
|
||||
).status_code == 400
|
||||
invalid = apiClient.put(
|
||||
"/api/user/user-1", headers=_bearer(), json={"timezone": "Invalid"}
|
||||
)
|
||||
empty = apiClient.put(
|
||||
"/api/user/user-1", headers=_bearer(), json={"username": "ignored"}
|
||||
)
|
||||
updated = apiClient.put(
|
||||
"/api/user/user-1", headers=_bearer(), json={"timezone": "UTC"}
|
||||
)
|
||||
missingPassword = apiClient.delete(
|
||||
"/api/user/user-1", headers=_bearer(), json={}
|
||||
)
|
||||
wrongPassword = apiClient.delete(
|
||||
"/api/user/user-1", headers=_bearer(), json={"password": "wrong"}
|
||||
)
|
||||
deleted = apiClient.delete(
|
||||
"/api/user/user-1", headers=_bearer(), json={"password": "correct"}
|
||||
)
|
||||
|
||||
assert invalid.get_json() == {"error": "bad timezone"}
|
||||
assert empty.get_json() == {"error": "no valid fields to update"}
|
||||
assert updated.get_json() == {"success": True}
|
||||
assert missingPassword.status_code == 400
|
||||
assert wrongPassword.status_code == 401
|
||||
assert deleted.get_json() == {"success": True}
|
||||
|
||||
|
||||
def test_timezone_routes_read_validate_and_update(monkeypatch, apiClient):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_userPrincipal()),
|
||||
)
|
||||
normalize = MagicMock(
|
||||
side_effect=["America/Chicago", ValueError("invalid timezone")]
|
||||
)
|
||||
updateUser = MagicMock(return_value=True)
|
||||
monkeypatch.setattr(
|
||||
apiMain.users,
|
||||
"getUserTimezone",
|
||||
MagicMock(return_value="UTC"),
|
||||
)
|
||||
monkeypatch.setattr(apiMain.users, "normalizeTimezone", normalize)
|
||||
monkeypatch.setattr(apiMain.users, "updateUser", updateUser)
|
||||
|
||||
current = apiClient.get("/api/user/me/timezone", headers=_bearer())
|
||||
updated = apiClient.put(
|
||||
"/api/user/me/timezone",
|
||||
headers=_bearer(),
|
||||
json={"timezone": "America/Chicago"},
|
||||
)
|
||||
invalid = apiClient.put(
|
||||
"/api/user/me/timezone",
|
||||
headers=_bearer(),
|
||||
json={"timezone": "Invalid"},
|
||||
)
|
||||
|
||||
assert current.get_json() == {"timezone": "UTC"}
|
||||
assert updated.get_json() == {"timezone": "America/Chicago"}
|
||||
assert invalid.status_code == 400
|
||||
assert invalid.get_json() == {"error": "invalid timezone"}
|
||||
updateUser.assert_called_once_with(
|
||||
"user-1",
|
||||
{"timezone": "America/Chicago"},
|
||||
)
|
||||
|
||||
|
||||
def test_outbox_claim_filters_missing_discord_identities(monkeypatch, apiClient):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_servicePrincipal("outbox:claim")),
|
||||
)
|
||||
claimMessages = MagicMock(
|
||||
side_effect=[
|
||||
ValueError("invalid limit"),
|
||||
[
|
||||
{
|
||||
"id": "message-1",
|
||||
"user_uuid": "user-1",
|
||||
"payload": {"content": "hello"},
|
||||
"attempts": 1,
|
||||
},
|
||||
{
|
||||
"id": "message-2",
|
||||
"user_uuid": "user-2",
|
||||
"payload": {},
|
||||
"attempts": 2,
|
||||
},
|
||||
],
|
||||
]
|
||||
)
|
||||
identities = MagicMock(
|
||||
side_effect=[
|
||||
[{"provider": "discord", "provider_user_id": "discord-1"}],
|
||||
[],
|
||||
]
|
||||
)
|
||||
retryMessage = MagicMock()
|
||||
monkeypatch.setattr(apiMain.outbox, "claim_messages", claimMessages)
|
||||
monkeypatch.setattr(apiMain.identity, "listUserIdentities", identities)
|
||||
monkeypatch.setattr(apiMain.outbox, "retry_message", retryMessage)
|
||||
|
||||
missingWorker = apiClient.post(
|
||||
"/api/internal/outbox/claim",
|
||||
headers=_bearer(),
|
||||
json={},
|
||||
)
|
||||
invalid = apiClient.post(
|
||||
"/api/internal/outbox/claim",
|
||||
headers=_bearer(),
|
||||
json={"worker_id": "worker-1", "limit": 0},
|
||||
)
|
||||
claimed = apiClient.post(
|
||||
"/api/internal/outbox/claim",
|
||||
headers=_bearer(),
|
||||
json={"worker_id": "worker-1", "channel": "discord_dm", "limit": 2},
|
||||
)
|
||||
|
||||
assert missingWorker.status_code == 400
|
||||
assert invalid.get_json() == {"error": "invalid limit"}
|
||||
assert claimed.get_json() == {
|
||||
"messages": [
|
||||
{
|
||||
"id": "message-1",
|
||||
"provider_user_id": "discord-1",
|
||||
"content": "hello",
|
||||
"attempts": 1,
|
||||
"worker_id": "worker-1",
|
||||
}
|
||||
]
|
||||
}
|
||||
retryMessage.assert_called_once_with(
|
||||
"message-2",
|
||||
"worker-1",
|
||||
"user has no Discord identity",
|
||||
)
|
||||
|
||||
|
||||
def test_outbox_result_enforces_lease_owner_and_records_results(
|
||||
monkeypatch, apiClient
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
apiMain.auth,
|
||||
"authenticateBearerToken",
|
||||
MagicMock(return_value=_servicePrincipal("outbox:deliver")),
|
||||
)
|
||||
leased = {"id": "message-1", "status": "delivering", "leased_by": "worker-1"}
|
||||
getMessage = MagicMock(side_effect=[None, leased, leased, leased])
|
||||
deliveredRecord = {"id": "message-1", "status": "delivered"}
|
||||
retryRecord = {"id": "message-1", "status": "pending"}
|
||||
markDelivered = MagicMock(return_value=deliveredRecord)
|
||||
retryMessage = MagicMock(return_value=retryRecord)
|
||||
monkeypatch.setattr(apiMain.outbox, "get_message", getMessage)
|
||||
monkeypatch.setattr(apiMain.outbox, "mark_delivered", markDelivered)
|
||||
monkeypatch.setattr(apiMain.outbox, "retry_message", retryMessage)
|
||||
|
||||
invalid = apiClient.post(
|
||||
"/api/internal/outbox/message-1/result",
|
||||
headers=_bearer(),
|
||||
json={"status": "unknown"},
|
||||
)
|
||||
missing = apiClient.post(
|
||||
"/api/internal/outbox/message-1/result",
|
||||
headers=_bearer(),
|
||||
json={"status": "sent", "worker_id": "worker-1"},
|
||||
)
|
||||
wrongWorker = apiClient.post(
|
||||
"/api/internal/outbox/message-1/result",
|
||||
headers=_bearer(),
|
||||
json={"status": "sent", "worker_id": "worker-2"},
|
||||
)
|
||||
delivered = apiClient.post(
|
||||
"/api/internal/outbox/message-1/result",
|
||||
headers=_bearer(),
|
||||
json={
|
||||
"status": "sent",
|
||||
"worker_id": "worker-1",
|
||||
"external_message_id": "x" * 300,
|
||||
},
|
||||
)
|
||||
retried = apiClient.post(
|
||||
"/api/internal/outbox/message-1/result",
|
||||
headers=_bearer(),
|
||||
json={
|
||||
"status": "retry",
|
||||
"worker_id": "worker-1",
|
||||
"error": "temporary failure",
|
||||
},
|
||||
)
|
||||
|
||||
assert invalid.status_code == 400
|
||||
assert missing.status_code == 404
|
||||
assert wrongWorker.status_code == 409
|
||||
assert delivered.get_json()["message"] == deliveredRecord
|
||||
assert retried.get_json()["message"] == retryRecord
|
||||
markDelivered.assert_called_once_with(
|
||||
"message-1",
|
||||
"worker-1",
|
||||
external_message_id="x" * 255,
|
||||
)
|
||||
retryMessage.assert_called_once_with(
|
||||
"message-1",
|
||||
"worker-1",
|
||||
"temporary failure",
|
||||
)
|
||||
|
||||
|
||||
def test_health_and_standard_error_responses(monkeypatch, apiClient):
|
||||
execute = MagicMock(side_effect=[[{"ready": 1}], RuntimeError("offline")])
|
||||
monkeypatch.setattr(apiMain.postgres, "execute", execute)
|
||||
|
||||
live = apiClient.get("/health/live")
|
||||
ready = apiClient.get("/health/ready")
|
||||
unavailable = apiClient.get("/health")
|
||||
missing = apiClient.get("/missing")
|
||||
wrongMethod = apiClient.get("/api/register")
|
||||
|
||||
assert live.get_json() == {"status": "ok"}
|
||||
assert ready.get_json() == {"status": "ready"}
|
||||
assert unavailable.status_code == 503
|
||||
assert unavailable.get_json() == {"status": "not ready"}
|
||||
assert missing.status_code == 404
|
||||
assert missing.get_json() == {"error": "not found"}
|
||||
assert wrongMethod.status_code == 405
|
||||
assert wrongMethod.get_json() == {"error": "method not allowed"}
|
||||
262
tests/unit/test_bot_runtime.py
Normal file
262
tests/unit/test_bot_runtime.py
Normal file
@@ -0,0 +1,262 @@
|
||||
"""Network-free tests for Discord adapter routing and delivery behavior."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from bot import bot as botModule
|
||||
|
||||
|
||||
class TypingContext:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return False
|
||||
|
||||
|
||||
class FakeChannel:
|
||||
def __init__(self):
|
||||
self.sent = []
|
||||
|
||||
async def send(self, content):
|
||||
self.sent.append(content)
|
||||
return content
|
||||
|
||||
def typing(self):
|
||||
return TypingContext()
|
||||
|
||||
|
||||
def _message(content="hello", authorID=123, channel=None):
|
||||
author = SimpleNamespace(
|
||||
id=authorID,
|
||||
display_name="Alice",
|
||||
__str__=lambda self: "Alice",
|
||||
)
|
||||
return SimpleNamespace(
|
||||
content=content,
|
||||
author=author,
|
||||
channel=channel or FakeChannel(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanBotCaches():
|
||||
botModule.api_clients.clear()
|
||||
botModule.message_history.clear()
|
||||
botModule.user_locks.clear()
|
||||
yield
|
||||
botModule.api_clients.clear()
|
||||
botModule.message_history.clear()
|
||||
botModule.user_locks.clear()
|
||||
|
||||
|
||||
def test_api_clients_are_cached_per_stable_discord_id(monkeypatch):
|
||||
created = []
|
||||
|
||||
def clientFactory(discordID, displayName):
|
||||
client = SimpleNamespace(discord_id=str(discordID), display_name=displayName)
|
||||
created.append(client)
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(botModule, "ApiClient", clientFactory)
|
||||
first = botModule.getApiClient(_message(authorID=123))
|
||||
second = botModule.getApiClient(_message(authorID=123))
|
||||
third = botModule.getApiClient(_message(authorID=456))
|
||||
assert first is second and third is not first
|
||||
assert len(created) == 2 and first.display_name == "Alice"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_help_message_uses_registered_module_help(monkeypatch):
|
||||
message = _message("help")
|
||||
monkeypatch.setattr(
|
||||
botModule.module_registry,
|
||||
"help_lines",
|
||||
MagicMock(return_value=["- remind me", "- list reminders"]),
|
||||
)
|
||||
await botModule.sendHelpMessage(message)
|
||||
assert "remind me" in message.channel.sent[0]
|
||||
assert "Just talk naturally" in message.channel.sent[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authentication_reuses_token_and_maps_api_failures(monkeypatch):
|
||||
message = _message()
|
||||
existing = SimpleNamespace(token="already-authenticated")
|
||||
monkeypatch.setattr(botModule, "getApiClient", lambda _message: existing)
|
||||
assert await botModule.authenticateMessage(message) is existing
|
||||
|
||||
for status, expected in [
|
||||
(403, "not enabled"),
|
||||
(503, "still starting"),
|
||||
(401, "couldn't start"),
|
||||
]:
|
||||
message = _message()
|
||||
client = SimpleNamespace(
|
||||
token=None,
|
||||
authenticate=AsyncMock(return_value=({"error": "denied"}, status)),
|
||||
)
|
||||
monkeypatch.setattr(botModule, "getApiClient", lambda _message, value=client: value)
|
||||
assert await botModule.authenticateMessage(message) is None
|
||||
assert expected in message.channel.sent[0]
|
||||
|
||||
accepted = SimpleNamespace(
|
||||
token=None,
|
||||
authenticate=AsyncMock(return_value=({"token": "new"}, 200)),
|
||||
)
|
||||
monkeypatch.setattr(botModule, "getApiClient", lambda _message: accepted)
|
||||
assert await botModule.authenticateMessage(_message()) is accepted
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_routing_handles_help_parser_errors_and_unknowns(monkeypatch):
|
||||
apiClient = SimpleNamespace(timezone="UTC", user_uuid="user-one")
|
||||
helpMessage = _message("help")
|
||||
helpCall = AsyncMock()
|
||||
monkeypatch.setattr(botModule, "sendHelpMessage", helpCall)
|
||||
await botModule.routeCommand(helpMessage, apiClient)
|
||||
helpCall.assert_awaited_once_with(helpMessage)
|
||||
|
||||
parse = AsyncMock(
|
||||
side_effect=[
|
||||
{"needs_clarification": "When should I do that?"},
|
||||
{"error": "provider unavailable"},
|
||||
{"interaction_type": "missing"},
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(botModule.ai_parser, "parse_command_async", parse)
|
||||
monkeypatch.setattr(botModule.module_registry, "get_command", lambda _name: None)
|
||||
|
||||
clarification = _message("do the thing")
|
||||
await botModule.routeCommand(clarification, apiClient)
|
||||
assert clarification.channel.sent == ["When should I do that?"]
|
||||
failed = _message("do the other thing")
|
||||
await botModule.routeCommand(failed, apiClient)
|
||||
assert "provider unavailable" in failed.channel.sent[0]
|
||||
unknown = _message("unknown feature")
|
||||
await botModule.routeCommand(unknown, apiClient)
|
||||
assert unknown.channel.sent == ["Unknown command type: missing"]
|
||||
assert len(botModule.message_history[123]) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_handler_receives_context_and_failure_is_isolated(monkeypatch):
|
||||
apiClient = SimpleNamespace(timezone="UTC", user_uuid="user-one")
|
||||
parse = AsyncMock(return_value={"interaction_type": "sample", "value": 1})
|
||||
handler = AsyncMock()
|
||||
monkeypatch.setattr(botModule.ai_parser, "parse_command_async", parse)
|
||||
monkeypatch.setattr(
|
||||
botModule.module_registry,
|
||||
"get_command",
|
||||
lambda _name: {"handler": handler},
|
||||
)
|
||||
message = _message("run sample")
|
||||
await botModule.routeCommand(message, apiClient)
|
||||
context, parsed = handler.await_args.args
|
||||
assert context.user_uuid == "user-one" and parsed["value"] == 1
|
||||
|
||||
handler.side_effect = RuntimeError("module failed")
|
||||
message = _message("run sample again")
|
||||
await botModule.routeCommand(message, apiClient)
|
||||
assert "failed unexpectedly" in message.channel.sent[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_message_filters_and_serializes_dm_work(monkeypatch):
|
||||
class FakeDMChannel(FakeChannel):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(botModule.discord, "DMChannel", FakeDMChannel)
|
||||
botUser = SimpleNamespace(id=999)
|
||||
monkeypatch.setattr(botModule, "client", SimpleNamespace(user=botUser))
|
||||
authenticate = AsyncMock(return_value=SimpleNamespace())
|
||||
route = AsyncMock()
|
||||
monkeypatch.setattr(botModule, "authenticateMessage", authenticate)
|
||||
monkeypatch.setattr(botModule, "routeCommand", route)
|
||||
|
||||
ownMessage = _message(channel=FakeDMChannel())
|
||||
ownMessage.author = botUser
|
||||
await botModule.on_message(ownMessage)
|
||||
await botModule.on_message(_message(channel=FakeChannel()))
|
||||
authenticate.assert_not_awaited()
|
||||
|
||||
dm = _message(channel=FakeDMChannel())
|
||||
await botModule.on_message(dm)
|
||||
authenticate.assert_awaited_once_with(dm)
|
||||
route.assert_awaited_once()
|
||||
assert dm.author.id in botModule.user_locks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_outbound_delivery_reports_success_and_retry(monkeypatch):
|
||||
sent = SimpleNamespace(id=987)
|
||||
discordUser = SimpleNamespace(send=AsyncMock(return_value=sent))
|
||||
client = SimpleNamespace(fetch_user=AsyncMock(return_value=discordUser))
|
||||
service = SimpleNamespace(service_request=AsyncMock(return_value=({}, 200)))
|
||||
monkeypatch.setattr(botModule, "client", client)
|
||||
monkeypatch.setattr(botModule, "service_client", service)
|
||||
outbound = {
|
||||
"id": "message-one",
|
||||
"provider_user_id": "123",
|
||||
"content": "hello",
|
||||
"worker_id": "worker-one",
|
||||
}
|
||||
await botModule.deliverOutboundMessage(outbound)
|
||||
payload = service.service_request.await_args.args[2]
|
||||
assert payload == {
|
||||
"status": "sent",
|
||||
"external_message_id": "987",
|
||||
"worker_id": "worker-one",
|
||||
}
|
||||
|
||||
client.fetch_user.side_effect = RuntimeError("Discord offline")
|
||||
service.service_request.reset_mock()
|
||||
await botModule.deliverOutboundMessage(outbound)
|
||||
payload = service.service_request.await_args.args[2]
|
||||
assert payload["status"] == "retry"
|
||||
assert payload["worker_id"] == "worker-one"
|
||||
assert "Discord offline" in payload["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ready_starts_outbox_loop_once(monkeypatch):
|
||||
loop = SimpleNamespace(is_running=MagicMock(side_effect=[False, True]), start=MagicMock())
|
||||
monkeypatch.setattr(botModule, "outboxLoop", loop)
|
||||
monkeypatch.setattr(botModule, "client", SimpleNamespace(user="bot-user"))
|
||||
await botModule.on_ready()
|
||||
await botModule.on_ready()
|
||||
loop.start.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_outbox_poll_handles_api_failure_and_delivers_batch(monkeypatch):
|
||||
outbound = {"id": "message-one"}
|
||||
service = SimpleNamespace(
|
||||
service_request=AsyncMock(
|
||||
side_effect=[
|
||||
({"error": "offline"}, 503),
|
||||
({"messages": [outbound]}, 200),
|
||||
]
|
||||
)
|
||||
)
|
||||
deliver = AsyncMock()
|
||||
monkeypatch.setattr(botModule, "service_client", service)
|
||||
monkeypatch.setattr(botModule, "deliverOutboundMessage", deliver)
|
||||
|
||||
await botModule.outboxLoop.coro()
|
||||
deliver.assert_not_awaited()
|
||||
await botModule.outboxLoop.coro()
|
||||
deliver.assert_awaited_once_with(outbound)
|
||||
request = service.service_request.await_args.args
|
||||
assert request[0:2] == ("post", "/api/internal/outbox/claim")
|
||||
assert request[2]["worker_id"] == botModule.OUTBOX_WORKER_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_outbox_before_loop_waits_for_discord_ready(monkeypatch):
|
||||
wait = AsyncMock()
|
||||
monkeypatch.setattr(botModule, "client", SimpleNamespace(wait_until_ready=wait))
|
||||
await botModule.beforeOutboxLoop()
|
||||
wait.assert_awaited_once()
|
||||
290
tests/unit/test_durable_queues.py
Normal file
290
tests/unit/test_durable_queues.py
Normal file
@@ -0,0 +1,290 @@
|
||||
"""Focused unit coverage for durable job and outbound-message state changes."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from core import jobs, outbox
|
||||
|
||||
|
||||
NOW = datetime.now(timezone.utc)
|
||||
MESSAGE_ID = "00000000-0000-0000-0000-000000000001"
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
def __init__(self, one=None, many=None):
|
||||
self.one = list(one or [])
|
||||
self.many = list(many or [])
|
||||
self.executed = []
|
||||
|
||||
def execute(self, query, params=None):
|
||||
self.executed.append((" ".join(query.split()), params))
|
||||
|
||||
def fetchone(self):
|
||||
return self.one.pop(0) if self.one else None
|
||||
|
||||
def fetchall(self):
|
||||
return self.many
|
||||
|
||||
|
||||
@pytest.mark.parametrize("queue", [jobs, outbox])
|
||||
def test_shared_timestamp_and_positive_validation(queue):
|
||||
assert queue._timestamp("2026-01-02T03:04:05Z") == datetime(
|
||||
2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc
|
||||
)
|
||||
assert queue._timestamp(datetime(2026, 1, 2, 3, 4, 5)).tzinfo == timezone.utc
|
||||
|
||||
with pytest.raises(ValueError, match="datetime"):
|
||||
queue._timestamp(123)
|
||||
with pytest.raises(ValueError, match="whole number"):
|
||||
queue._positive(None, "limit")
|
||||
with pytest.raises(ValueError, match="at least 1"):
|
||||
queue._positive(0, "limit")
|
||||
with pytest.raises(ValueError, match="at most 2"):
|
||||
queue._positive(3, "limit", 2)
|
||||
|
||||
|
||||
def test_create_get_and_list_jobs_with_parameterized_filters():
|
||||
created = {"id": "job-one", "status": "pending"}
|
||||
cursor = RecordingCursor(one=[created])
|
||||
result = jobs.create_job(
|
||||
" sample.work ",
|
||||
{"value": 1},
|
||||
NOW,
|
||||
user_uuid="user-one",
|
||||
max_attempts="4",
|
||||
idempotency_key="unique-work",
|
||||
job_id="job-one",
|
||||
cursor=cursor,
|
||||
)
|
||||
assert result == created
|
||||
params = cursor.executed[0][1]
|
||||
assert params["job_type"] == "sample.work"
|
||||
assert params["max_attempts"] == 4
|
||||
assert params["payload"].adapted == {"value": 1}
|
||||
|
||||
with pytest.raises(ValueError, match="job_type"):
|
||||
jobs.create_job(" ", {}, NOW, cursor=cursor)
|
||||
|
||||
cursor = RecordingCursor(one=[created])
|
||||
assert jobs.get_job("job-one", cursor=cursor) == created
|
||||
|
||||
cursor = RecordingCursor(many=[created])
|
||||
assert jobs.list_jobs(
|
||||
user_uuid="user-one",
|
||||
status="pending",
|
||||
job_type="sample.work",
|
||||
limit=3,
|
||||
cursor=cursor,
|
||||
) == [created]
|
||||
query, params = cursor.executed[0]
|
||||
assert "user_uuid = %s" in query and "job_type = %s" in query
|
||||
assert params == ["user-one", "pending", "sample.work", 3]
|
||||
|
||||
|
||||
def test_job_claim_renew_complete_retry_and_cancel_paths():
|
||||
claimed = [{"id": "job-one", "status": "running"}]
|
||||
cursor = RecordingCursor(many=claimed)
|
||||
assert jobs.claim_due_jobs(
|
||||
"worker-one",
|
||||
limit=2,
|
||||
lease_seconds=60,
|
||||
job_types="sample.work",
|
||||
cursor=cursor,
|
||||
) == claimed
|
||||
assert len(cursor.executed) == 2
|
||||
assert cursor.executed[1][1]["job_types"] == ["sample.work"]
|
||||
assert jobs.claim_due_jobs("worker", job_types=[], cursor=cursor) == []
|
||||
with pytest.raises(ValueError, match="worker_id"):
|
||||
jobs.claim_due_jobs("", cursor=cursor)
|
||||
|
||||
updated = {"id": "job-one", "status": "running"}
|
||||
cursor = RecordingCursor(one=[updated, {**updated, "status": "completed"}])
|
||||
assert jobs.renew_job_lease("job-one", "worker-one", 30, cursor=cursor) == updated
|
||||
assert jobs.complete_job("job-one", "worker-one", cursor=cursor)[
|
||||
"status"
|
||||
] == "completed"
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 2, "max_attempts": 3},
|
||||
{"id": "job-one", "status": "pending"},
|
||||
]
|
||||
)
|
||||
retried = jobs.fail_job(
|
||||
"job-one", "worker-one", "temporary", retry_seconds=10, cursor=cursor
|
||||
)
|
||||
assert retried["status"] == "pending"
|
||||
assert cursor.executed[1][1]["delay"] == 20
|
||||
assert cursor.executed[1][1]["exhausted"] is False
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 3, "max_attempts": 3},
|
||||
{"id": "job-one", "status": "failed"},
|
||||
]
|
||||
)
|
||||
assert jobs.fail_job("job-one", "worker-one", "fatal", cursor=cursor)[
|
||||
"status"
|
||||
] == "failed"
|
||||
assert cursor.executed[1][1]["exhausted"] is True
|
||||
assert jobs.fail_job(
|
||||
"missing", "worker-one", "ignored", cursor=RecordingCursor()
|
||||
) is None
|
||||
|
||||
cursor = RecordingCursor(one=[{"id": "job-one", "status": "cancelled"}])
|
||||
assert jobs.cancel_job("job-one", user_uuid="user-one", cursor=cursor)[
|
||||
"status"
|
||||
] == "cancelled"
|
||||
cursor = RecordingCursor(many=[{"id": "job-two"}])
|
||||
assert jobs.cancel_jobs(job_type="sample.work", cursor=cursor) == [
|
||||
{"id": "job-two"}
|
||||
]
|
||||
with pytest.raises(ValueError, match="filter"):
|
||||
jobs.cancel_jobs(cursor=cursor)
|
||||
|
||||
|
||||
def test_enqueue_get_and_list_messages_with_parameterized_filters():
|
||||
created = {"id": "message-one", "status": "pending"}
|
||||
cursor = RecordingCursor(one=[created])
|
||||
result = outbox.enqueue_message(
|
||||
"user-one",
|
||||
" discord_dm ",
|
||||
{"content": "hello"},
|
||||
"unique-message",
|
||||
available_at=NOW,
|
||||
max_attempts=4,
|
||||
message_id="message-one",
|
||||
cursor=cursor,
|
||||
)
|
||||
assert result == created
|
||||
params = cursor.executed[0][1]
|
||||
assert params["channel"] == "discord_dm"
|
||||
assert params["payload"].adapted == {"content": "hello"}
|
||||
|
||||
invalidValues = [
|
||||
(None, "discord_dm", {}, "key", "user_uuid"),
|
||||
("user", " ", {}, "key", "channel"),
|
||||
("user", "discord_dm", {}, None, "idempotency_key"),
|
||||
("user", "discord_dm", None, "key", "payload"),
|
||||
]
|
||||
for userUUID, channel, payload, key, error in invalidValues:
|
||||
with pytest.raises(ValueError, match=error):
|
||||
outbox.enqueue_message(
|
||||
userUUID, channel, payload, key, available_at=NOW, cursor=cursor
|
||||
)
|
||||
|
||||
cursor = RecordingCursor(one=[created])
|
||||
assert outbox.get_message(MESSAGE_ID, cursor=cursor) == created
|
||||
assert outbox.get_message("not-a-uuid", cursor=cursor) is None
|
||||
cursor = RecordingCursor(many=[created])
|
||||
assert outbox.list_messages(
|
||||
user_uuid="user-one",
|
||||
status="pending",
|
||||
channel="discord_dm",
|
||||
limit=2,
|
||||
cursor=cursor,
|
||||
) == [created]
|
||||
|
||||
|
||||
def test_outbox_claim_renew_delivery_retry_and_cancel_paths():
|
||||
claimed = [{"id": "message-one", "status": "delivering"}]
|
||||
cursor = RecordingCursor(many=claimed)
|
||||
assert outbox.claim_messages(
|
||||
"worker-one",
|
||||
channel="discord_dm",
|
||||
limit=2,
|
||||
lease_seconds=60,
|
||||
cursor=cursor,
|
||||
) == claimed
|
||||
assert len(cursor.executed) == 2
|
||||
with pytest.raises(ValueError, match="worker_id"):
|
||||
outbox.claim_messages(None, cursor=cursor)
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"id": "message-one", "status": "delivering"},
|
||||
{"id": "message-one", "status": "delivered"},
|
||||
]
|
||||
)
|
||||
assert outbox.renew_message_lease(
|
||||
"message-one", "worker-one", 60, cursor=cursor
|
||||
)["status"] == "delivering"
|
||||
delivered = outbox.mark_delivered(
|
||||
"message-one",
|
||||
"worker-one",
|
||||
external_message_id="discord-123",
|
||||
cursor=cursor,
|
||||
)
|
||||
assert delivered["status"] == "delivered"
|
||||
assert cursor.executed[1][1][0] == "discord-123"
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 1, "max_attempts": 3},
|
||||
{"id": "message-one", "status": "pending"},
|
||||
]
|
||||
)
|
||||
retried = outbox.retry_message(
|
||||
"message-one", "worker-one", "temporary", retry_seconds=15, cursor=cursor
|
||||
)
|
||||
assert retried["status"] == "pending"
|
||||
assert cursor.executed[1][1]["delay"] == 15
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 3, "max_attempts": 3},
|
||||
{"id": "message-one", "status": "failed"},
|
||||
]
|
||||
)
|
||||
assert outbox.retry_message(
|
||||
"message-one", "worker-one", "fatal", cursor=cursor
|
||||
)["status"] == "failed"
|
||||
assert outbox.retry_message(
|
||||
"missing", "worker-one", "ignored", cursor=RecordingCursor()
|
||||
) is None
|
||||
|
||||
cursor = RecordingCursor(one=[{"id": "message-one", "status": "cancelled"}])
|
||||
assert outbox.cancel_message(
|
||||
"message-one", user_uuid="user-one", cursor=cursor
|
||||
)["status"] == "cancelled"
|
||||
cursor = RecordingCursor(many=[{"id": "message-two"}])
|
||||
assert outbox.cancel_messages(channel="discord_dm", cursor=cursor) == [
|
||||
{"id": "message-two"}
|
||||
]
|
||||
with pytest.raises(ValueError, match="filter"):
|
||||
outbox.cancel_messages(cursor=cursor)
|
||||
|
||||
|
||||
def test_retry_backoff_is_capped():
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 20, "max_attempts": 30},
|
||||
{"id": "job", "status": "pending"},
|
||||
]
|
||||
)
|
||||
jobs.fail_job(
|
||||
"job",
|
||||
"worker",
|
||||
"retry",
|
||||
retry_seconds=30,
|
||||
max_retry_seconds=90,
|
||||
cursor=cursor,
|
||||
)
|
||||
assert cursor.executed[1][1]["delay"] == 90
|
||||
|
||||
cursor = RecordingCursor(
|
||||
one=[
|
||||
{"attempts": 20, "max_attempts": 30},
|
||||
{"id": "message", "status": "pending"},
|
||||
]
|
||||
)
|
||||
outbox.retry_message(
|
||||
"message",
|
||||
"worker",
|
||||
"retry",
|
||||
retry_seconds=30,
|
||||
max_retry_seconds=90,
|
||||
cursor=cursor,
|
||||
)
|
||||
assert cursor.executed[1][1]["delay"] == 90
|
||||
170
tests/unit/test_migration_engine.py
Normal file
170
tests/unit/test_migration_engine.py
Normal file
@@ -0,0 +1,170 @@
|
||||
"""Unit tests for migration discovery and history decisions."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
|
||||
import pytest
|
||||
|
||||
from core import migrations
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def migrationRoot():
|
||||
with TemporaryDirectory(prefix=".migration-test-", dir=Path.cwd()) as directory:
|
||||
yield Path(directory)
|
||||
|
||||
|
||||
def _writeMigration(root, relativePath, content="SELECT 1;"):
|
||||
path = root / relativePath
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(content, encoding="utf-8")
|
||||
return path
|
||||
|
||||
|
||||
class FakeCursor:
|
||||
def __init__(self, applied=None, historyTable=True):
|
||||
self.applied = list(applied or [])
|
||||
self.historyTable = historyTable
|
||||
self.executed = []
|
||||
self._one = None
|
||||
self._all = []
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def execute(self, query, params=None):
|
||||
normalized = " ".join(query.split())
|
||||
self.executed.append((normalized, params))
|
||||
if "SELECT to_regclass" in normalized:
|
||||
self._one = {
|
||||
"table_name": "schema_migrations" if self.historyTable else None
|
||||
}
|
||||
elif normalized.startswith("SELECT namespace, version, checksum"):
|
||||
self._all = self.applied
|
||||
|
||||
def fetchone(self):
|
||||
return self._one
|
||||
|
||||
def fetchall(self):
|
||||
return self._all
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self, cursor):
|
||||
self.activeCursor = cursor
|
||||
|
||||
def cursor(self, **_kwargs):
|
||||
return self.activeCursor
|
||||
|
||||
|
||||
def _connectionFor(cursor):
|
||||
@contextmanager
|
||||
def fakeConnection():
|
||||
yield FakeConnection(cursor)
|
||||
|
||||
return fakeConnection
|
||||
|
||||
|
||||
def test_discovery_orders_core_before_feature_namespaces(migrationRoot):
|
||||
_writeMigration(migrationRoot, "config/migrations/0002_second.sql", "SELECT 2;")
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_first.sql", "SELECT 1;")
|
||||
_writeMigration(migrationRoot, "modules/zeta/migrations/0001_zeta.sql")
|
||||
_writeMigration(migrationRoot, "modules/alpha/migrations/0002_alpha.sql")
|
||||
|
||||
found = migrations.discover_migrations(migrationRoot)
|
||||
|
||||
assert [(item.namespace, item.version) for item in found] == [
|
||||
("core", 1),
|
||||
("core", 2),
|
||||
("alpha", 2),
|
||||
("zeta", 1),
|
||||
]
|
||||
assert len(found[0].checksum) == 64
|
||||
assert found[0].path == Path(
|
||||
migrationRoot, "config/migrations/0001_first.sql"
|
||||
)
|
||||
|
||||
|
||||
def test_discovery_rejects_bad_names_and_duplicate_versions(migrationRoot):
|
||||
_writeMigration(migrationRoot, "config/migrations/not-numbered.sql")
|
||||
with pytest.raises(migrations.MigrationError, match="Invalid migration filename"):
|
||||
migrations.discover_migrations(migrationRoot)
|
||||
|
||||
Path(migrationRoot, "config/migrations/not-numbered.sql").unlink()
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_first.sql")
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_duplicate.sql")
|
||||
with pytest.raises(migrations.MigrationError, match="Duplicate migration"):
|
||||
migrations.discover_migrations(migrationRoot)
|
||||
|
||||
|
||||
def test_upgrade_applies_pending_and_skips_matching_history(
|
||||
migrationRoot, monkeypatch
|
||||
):
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_first.sql", "SELECT 42;")
|
||||
migration = migrations.discover_migrations(migrationRoot)[0]
|
||||
cursor = FakeCursor()
|
||||
monkeypatch.setattr(migrations.postgres, "get_connection", _connectionFor(cursor))
|
||||
|
||||
assert migrations.upgrade(migrationRoot) == [migration]
|
||||
queries = [query for query, _params in cursor.executed]
|
||||
assert any("pg_advisory_xact_lock" in query for query in queries)
|
||||
assert "SELECT 42;" in queries
|
||||
assert any(query.startswith("INSERT INTO schema_migrations") for query in queries)
|
||||
|
||||
cursor = FakeCursor(
|
||||
applied=[
|
||||
{
|
||||
"namespace": "core",
|
||||
"version": 1,
|
||||
"checksum": migration.checksum,
|
||||
"applied_at": None,
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(migrations.postgres, "get_connection", _connectionFor(cursor))
|
||||
assert migrations.upgrade(migrationRoot) == []
|
||||
assert "SELECT 42;" not in [query for query, _params in cursor.executed]
|
||||
|
||||
|
||||
def test_upgrade_rejects_changed_applied_migration(migrationRoot, monkeypatch):
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_first.sql", "SELECT 1;")
|
||||
cursor = FakeCursor(
|
||||
applied=[
|
||||
{
|
||||
"namespace": "core",
|
||||
"version": 1,
|
||||
"checksum": "0" * 64,
|
||||
"applied_at": None,
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(migrations.postgres, "get_connection", _connectionFor(cursor))
|
||||
|
||||
with pytest.raises(migrations.MigrationError, match="checksum changed"):
|
||||
migrations.upgrade(migrationRoot)
|
||||
|
||||
|
||||
def test_status_handles_new_database_and_missing_source(migrationRoot, monkeypatch):
|
||||
_writeMigration(migrationRoot, "config/migrations/0001_first.sql", "SELECT 1;")
|
||||
cursor = FakeCursor(historyTable=False)
|
||||
monkeypatch.setattr(migrations.postgres, "get_connection", _connectionFor(cursor))
|
||||
assert migrations.migration_status(migrationRoot)[0]["state"] == "pending"
|
||||
|
||||
cursor = FakeCursor(
|
||||
applied=[
|
||||
{
|
||||
"namespace": "removed_feature",
|
||||
"version": 3,
|
||||
"checksum": "a" * 64,
|
||||
"applied_at": "earlier",
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(migrations.postgres, "get_connection", _connectionFor(cursor))
|
||||
status = migrations.migration_status(migrationRoot)
|
||||
assert [item["state"] for item in status] == ["pending", "missing"]
|
||||
assert status[1]["namespace"] == "removed_feature"
|
||||
223
tests/unit/test_notifications.py
Normal file
223
tests/unit/test_notifications.py
Normal file
@@ -0,0 +1,223 @@
|
||||
from unittest.mock import MagicMock, call
|
||||
|
||||
import pytest
|
||||
|
||||
from core import notifications
|
||||
|
||||
|
||||
def test_get_notification_settings_returns_record_or_false(monkeypatch):
|
||||
selectOne = MagicMock(
|
||||
side_effect=[{"user_uuid": "user-1", "ntfy_enabled": True}, None]
|
||||
)
|
||||
monkeypatch.setattr(notifications.postgres, "select_one", selectOne)
|
||||
|
||||
assert notifications.getNotificationSettings("user-1") == {
|
||||
"user_uuid": "user-1",
|
||||
"ntfy_enabled": True,
|
||||
}
|
||||
assert notifications.getNotificationSettings("user-2") is False
|
||||
assert selectOne.call_args_list == [
|
||||
call("notifications", {"user_uuid": "user-1"}),
|
||||
call("notifications", {"user_uuid": "user-2"}),
|
||||
]
|
||||
|
||||
|
||||
def test_notification_settings_filter_fields_and_update_existing(monkeypatch):
|
||||
update = MagicMock()
|
||||
insert = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
notifications.postgres,
|
||||
"select_one",
|
||||
MagicMock(return_value={"id": "notification-1"}),
|
||||
)
|
||||
monkeypatch.setattr(notifications.postgres, "update", update)
|
||||
monkeypatch.setattr(notifications.postgres, "insert", insert)
|
||||
|
||||
result = notifications.setNotificationSettings(
|
||||
"user-1",
|
||||
{
|
||||
"ntfy_topic": "team-alerts",
|
||||
"ntfy_enabled": True,
|
||||
"user_uuid": "another-user",
|
||||
"created_at": "not-allowed",
|
||||
},
|
||||
)
|
||||
|
||||
assert result is True
|
||||
update.assert_called_once_with(
|
||||
"notifications",
|
||||
{"ntfy_topic": "team-alerts", "ntfy_enabled": True},
|
||||
{"user_uuid": "user-1"},
|
||||
)
|
||||
insert.assert_not_called()
|
||||
|
||||
|
||||
def test_notification_settings_insert_new_record(monkeypatch):
|
||||
insert = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
notifications.postgres,
|
||||
"select_one",
|
||||
MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(notifications.postgres, "insert", insert)
|
||||
monkeypatch.setattr(notifications.uuid, "uuid4", lambda: "notification-1")
|
||||
|
||||
result = notifications.setNotificationSettings(
|
||||
"user-1",
|
||||
{"discord_enabled": False, "ntfy_topic": "personal"},
|
||||
)
|
||||
|
||||
assert result is True
|
||||
insert.assert_called_once_with(
|
||||
"notifications",
|
||||
{
|
||||
"discord_enabled": False,
|
||||
"ntfy_topic": "personal",
|
||||
"id": "notification-1",
|
||||
"user_uuid": "user-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"settings",
|
||||
[None, [], "invalid", {}, {"created_at": "not-allowed"}],
|
||||
)
|
||||
def test_notification_settings_reject_invalid_or_empty_updates(monkeypatch, settings):
|
||||
selectOne = MagicMock()
|
||||
monkeypatch.setattr(notifications.postgres, "select_one", selectOne)
|
||||
|
||||
assert notifications.setNotificationSettings("user-1", settings) is False
|
||||
selectOne.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"webhook",
|
||||
[
|
||||
"https://discord.com/api/webhooks/123/token",
|
||||
"https://canary.discord.com/api/webhooks/123/token",
|
||||
"https://ptb.discord.com/api/webhooks/123/token",
|
||||
],
|
||||
)
|
||||
def test_discord_webhook_validation_accepts_official_https_urls(webhook):
|
||||
assert notifications._validateDiscordWebhook(webhook) == webhook
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("webhook", "error"),
|
||||
[
|
||||
(
|
||||
"http://discord.com/api/webhooks/123/token",
|
||||
"official HTTPS Discord host",
|
||||
),
|
||||
(
|
||||
"https://discord.com.evil.example/api/webhooks/123/token",
|
||||
"official HTTPS Discord host",
|
||||
),
|
||||
("https://discord.com/channels/123", "Invalid Discord webhook path"),
|
||||
],
|
||||
)
|
||||
def test_discord_webhook_validation_rejects_unsafe_urls(webhook, error):
|
||||
with pytest.raises(ValueError, match=error):
|
||||
notifications._validateDiscordWebhook(webhook)
|
||||
|
||||
|
||||
def test_discord_webhook_delivery_posts_content(monkeypatch):
|
||||
post = MagicMock(return_value=MagicMock(status_code=204))
|
||||
monkeypatch.setattr(notifications.requests, "post", post)
|
||||
webhook = "https://discord.com/api/webhooks/123/token"
|
||||
|
||||
assert notifications.discord.send(webhook, 42) is True
|
||||
post.assert_called_once_with(
|
||||
webhook,
|
||||
json={"content": "42"},
|
||||
timeout=notifications.REQUEST_TIMEOUT,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", [429, notifications.requests.ConnectionError("offline")])
|
||||
def test_discord_webhook_delivery_reports_failures(monkeypatch, failure):
|
||||
post = MagicMock()
|
||||
if isinstance(failure, int):
|
||||
post.return_value = MagicMock(status_code=failure)
|
||||
else:
|
||||
post.side_effect = failure
|
||||
monkeypatch.setattr(notifications.requests, "post", post)
|
||||
|
||||
assert notifications.discord.send(
|
||||
"https://discord.com/api/webhooks/123/token",
|
||||
"hello",
|
||||
) is False
|
||||
|
||||
|
||||
def test_ntfy_encodes_topic_and_sends_bearer_token(monkeypatch):
|
||||
post = MagicMock(return_value=MagicMock(status_code=201))
|
||||
monkeypatch.setattr(notifications.requests, "post", post)
|
||||
monkeypatch.setenv("NTFY_BASE_URL", "https://notify.example/base/")
|
||||
monkeypatch.setenv("NTFY_TOKEN", "ntfy-secret")
|
||||
|
||||
assert notifications.ntfy.send(" alerts/team #1 ", 42) is True
|
||||
post.assert_called_once_with(
|
||||
"https://notify.example/base/alerts%2Fteam%20%231",
|
||||
data=b"42",
|
||||
headers={"Authorization": "Bearer ntfy-secret"},
|
||||
timeout=notifications.REQUEST_TIMEOUT,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", [500, notifications.requests.Timeout("slow")])
|
||||
def test_ntfy_reports_http_and_transport_failures(monkeypatch, failure):
|
||||
post = MagicMock()
|
||||
if isinstance(failure, int):
|
||||
post.return_value = MagicMock(status_code=failure)
|
||||
else:
|
||||
post.side_effect = failure
|
||||
monkeypatch.setattr(notifications.requests, "post", post)
|
||||
monkeypatch.delenv("NTFY_TOKEN", raising=False)
|
||||
|
||||
assert notifications.ntfy.send("alerts", "hello") is False
|
||||
|
||||
|
||||
def test_ntfy_rejects_empty_topic_without_request(monkeypatch):
|
||||
post = MagicMock()
|
||||
monkeypatch.setattr(notifications.requests, "post", post)
|
||||
|
||||
assert notifications.ntfy.send(" ", "hello") is False
|
||||
assert notifications.ntfy.send(None, "hello") is False
|
||||
post.assert_not_called()
|
||||
|
||||
|
||||
def test_channel_aggregation_tries_each_enabled_channel(monkeypatch):
|
||||
discordSend = MagicMock(return_value=False)
|
||||
ntfySend = MagicMock(return_value=True)
|
||||
monkeypatch.setattr(notifications.discord, "send", discordSend)
|
||||
monkeypatch.setattr(notifications.ntfy, "send", ntfySend)
|
||||
|
||||
result = notifications._sendToEnabledChannels(
|
||||
{
|
||||
"discord_enabled": True,
|
||||
"discord_webhook": "https://discord.com/api/webhooks/123/token",
|
||||
"ntfy_enabled": True,
|
||||
"ntfy_topic": "alerts",
|
||||
},
|
||||
"hello",
|
||||
)
|
||||
|
||||
assert result is True
|
||||
discordSend.assert_called_once_with(
|
||||
"https://discord.com/api/webhooks/123/token",
|
||||
"hello",
|
||||
)
|
||||
ntfySend.assert_called_once_with("alerts", "hello")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("settings", [None, {}, {"discord_enabled": False}])
|
||||
def test_channel_aggregation_skips_unconfigured_channels(monkeypatch, settings):
|
||||
discordSend = MagicMock()
|
||||
ntfySend = MagicMock()
|
||||
monkeypatch.setattr(notifications.discord, "send", discordSend)
|
||||
monkeypatch.setattr(notifications.ntfy, "send", ntfySend)
|
||||
|
||||
assert notifications._sendToEnabledChannels(settings, "hello") is False
|
||||
discordSend.assert_not_called()
|
||||
ntfySend.assert_not_called()
|
||||
275
tests/unit/test_parser.py
Normal file
275
tests/unit/test_parser.py
Normal file
@@ -0,0 +1,275 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from ai import parser
|
||||
from core.registry import FrameworkRegistry
|
||||
|
||||
|
||||
def _response(content=None, reasoning=None, choices=True):
|
||||
message = SimpleNamespace(content=content, reasoning=reasoning)
|
||||
return SimpleNamespace(choices=[SimpleNamespace(message=message)] if choices else [])
|
||||
|
||||
|
||||
def test_template_rendering_preserves_literal_json_braces():
|
||||
template = (
|
||||
'Schema: {"action":"create","payload":{"items":[1, 2]}}\n'
|
||||
"Input: {user_input}\nUnknown: {leave_this_alone}"
|
||||
)
|
||||
|
||||
rendered = parser._render_template(template, {"user_input": "remind me"})
|
||||
|
||||
assert '{"action":"create","payload":{"items":[1, 2]}}' in rendered
|
||||
assert "Input: remind me" in rendered
|
||||
assert "Unknown: {leave_this_alone}" in rendered
|
||||
|
||||
|
||||
def test_json_extraction_handles_fenced_text_and_nested_values():
|
||||
response = """```json
|
||||
Here is the structured result:
|
||||
{"action":"create","payload":{"items":[{"id":1},{"id":2}]}}
|
||||
```"""
|
||||
|
||||
assert parser._extract_json_from_text(response) == {
|
||||
"action": "create",
|
||||
"payload": {"items": [{"id": 1}, {"id": 2}]},
|
||||
}
|
||||
|
||||
|
||||
def test_parse_retries_invalid_json_and_validator_failures(monkeypatch):
|
||||
responses = iter(
|
||||
[
|
||||
"not JSON",
|
||||
'{"valid":false}',
|
||||
'```json\n{"valid":true,"nested":{"value":7}}\n```',
|
||||
]
|
||||
)
|
||||
prompts = []
|
||||
|
||||
def fake_llm(_system_prompt, user_prompt):
|
||||
prompts.append(user_prompt)
|
||||
return next(responses)
|
||||
|
||||
prompt = {
|
||||
"system": "Return an object",
|
||||
"user_template": (
|
||||
'Literal schema: {"valid":true}\nUser message: {user_input}'
|
||||
),
|
||||
}
|
||||
validator = lambda value: [] if value.get("valid") else ["valid must be true"]
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 3)
|
||||
monkeypatch.setattr(parser, "_call_llm", fake_llm)
|
||||
|
||||
result = parser.parse(
|
||||
"test input",
|
||||
"focused",
|
||||
prompt_override=prompt,
|
||||
validator=validator,
|
||||
)
|
||||
|
||||
assert result == {"valid": True, "nested": {"value": 7}}
|
||||
assert len(prompts) == 3
|
||||
assert 'Literal schema: {"valid":true}' in prompts[0]
|
||||
assert "Response was not valid JSON" in prompts[1]
|
||||
assert "valid must be true" in prompts[2]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_parser_routes_then_uses_focused_module_prompt(monkeypatch):
|
||||
target = FrameworkRegistry()
|
||||
target.begin_module("reminders", "modules.reminders")
|
||||
target.register_command(
|
||||
"reminder",
|
||||
lambda _context, _parsed: None,
|
||||
{
|
||||
"system": "Focused reminder parser",
|
||||
"user_template": (
|
||||
'Timezone: {timezone}\nMessage: {user_input}\n'
|
||||
'Schema: {"action":"create"}'
|
||||
),
|
||||
},
|
||||
validator=lambda value: (
|
||||
[] if value.get("action") == "create" else ["invalid action"]
|
||||
),
|
||||
description="Create and manage reminders",
|
||||
)
|
||||
target.finish_module()
|
||||
|
||||
llm = AsyncMock(
|
||||
side_effect=[
|
||||
'{"interaction_type":"reminder","confidence":0.99}',
|
||||
(
|
||||
'```json\n{"action":"create","message":"call home",'
|
||||
'"payload":{"source":"dm"}}\n```'
|
||||
),
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(parser, "_call_llm_async", llm)
|
||||
|
||||
result = await parser.parse_command_async(
|
||||
"remind me to call home",
|
||||
target,
|
||||
history=[("hello", {"interaction_type": "reminder"})],
|
||||
timezone_name="America/Chicago",
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"action": "create",
|
||||
"message": "call home",
|
||||
"payload": {"source": "dm"},
|
||||
"interaction_type": "reminder",
|
||||
}
|
||||
assert llm.await_count == 2
|
||||
route_system, route_user = llm.await_args_list[0].args
|
||||
focused_system, focused_user = llm.await_args_list[1].args
|
||||
assert route_system == parser.AI_CONFIG["prompts"]["command_parser"]["system"]
|
||||
assert "- reminder: Create and manage reminders" in route_user
|
||||
assert focused_system == "Focused reminder parser"
|
||||
assert "Timezone: America/Chicago" in focused_user
|
||||
assert 'Schema: {"action":"create"}' in focused_user
|
||||
|
||||
|
||||
def test_clients_are_created_lazily_and_cached(monkeypatch):
|
||||
syncClient = SimpleNamespace()
|
||||
asyncClient = SimpleNamespace()
|
||||
syncFactory = MagicMock(return_value=syncClient)
|
||||
asyncFactory = MagicMock(return_value=asyncClient)
|
||||
monkeypatch.setattr(parser, "OpenAI", syncFactory)
|
||||
monkeypatch.setattr(parser, "AsyncOpenAI", asyncFactory)
|
||||
monkeypatch.setattr(parser, "_sync_client", None)
|
||||
monkeypatch.setattr(parser, "_async_client", None)
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "provider-key")
|
||||
|
||||
assert parser._get_client() is syncClient
|
||||
assert parser._get_client() is syncClient
|
||||
assert parser._get_client(async_client=True) is asyncClient
|
||||
assert parser._get_client(async_client=True) is asyncClient
|
||||
assert syncFactory.call_count == 1
|
||||
assert asyncFactory.call_count == 1
|
||||
|
||||
|
||||
def test_response_text_and_request_arguments(monkeypatch):
|
||||
assert parser._response_text(_response(" result ")) == "result"
|
||||
assert parser._response_text(_response(reasoning=" reason ")) == "reason"
|
||||
assert parser._response_text(_response(choices=False)) is None
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG, "json_mode", True)
|
||||
arguments = parser._request_args("system", "user")
|
||||
assert arguments["messages"] == [
|
||||
{"role": "system", "content": "system"},
|
||||
{"role": "user", "content": "user"},
|
||||
]
|
||||
assert arguments["response_format"] == {"type": "json_object"}
|
||||
|
||||
|
||||
def test_sync_llm_call_returns_content_and_absorbs_provider_failure(monkeypatch):
|
||||
completions = MagicMock()
|
||||
completions.create.side_effect = [
|
||||
_response('{"ok":true}'),
|
||||
RuntimeError("provider unavailable"),
|
||||
]
|
||||
client = SimpleNamespace(chat=SimpleNamespace(completions=completions))
|
||||
monkeypatch.setattr(parser, "_get_client", lambda **_kwargs: client)
|
||||
|
||||
assert parser._call_llm("system", "user") == '{"ok":true}'
|
||||
assert parser._call_llm("system", "user") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_llm_call_returns_content_and_absorbs_failure(monkeypatch):
|
||||
create = AsyncMock(
|
||||
side_effect=[_response('{"ok":true}'), RuntimeError("offline")]
|
||||
)
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(completions=SimpleNamespace(create=create))
|
||||
)
|
||||
monkeypatch.setattr(parser, "_get_client", lambda **_kwargs: client)
|
||||
|
||||
assert await parser._call_llm_async("system", "user") == '{"ok":true}'
|
||||
assert await parser._call_llm_async("system", "user") is None
|
||||
|
||||
|
||||
def test_parse_reports_unknown_unavailable_and_exhausted_results(monkeypatch):
|
||||
assert parser.parse("hello", "missing") == {
|
||||
"error": "Unknown interaction type: missing"
|
||||
}
|
||||
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: None)
|
||||
assert parser.parse("hello", "test", prompt_override=prompt) == {
|
||||
"error": "AI service unavailable",
|
||||
"user_input": "hello",
|
||||
}
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 2)
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: "not json")
|
||||
exhausted = parser.parse("hello", "test", prompt_override=prompt)
|
||||
assert exhausted["error"] == "Failed to parse after 2 attempts"
|
||||
assert exhausted["validation_errors"] == ["Response was not valid JSON"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_parse_unknown_unavailable_and_exhausted(monkeypatch):
|
||||
assert await parser.parse_async("hello", "missing") == {
|
||||
"error": "Unknown interaction type: missing"
|
||||
}
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
call = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(parser, "_call_llm_async", call)
|
||||
assert await parser.parse_async("hello", "test", prompt_override=prompt) == {
|
||||
"error": "AI service unavailable",
|
||||
"user_input": "hello",
|
||||
}
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 1)
|
||||
call.return_value = "[]"
|
||||
exhausted = await parser.parse_async(
|
||||
"hello", "test", prompt_override=prompt
|
||||
)
|
||||
assert exhausted["validation_errors"] == [
|
||||
"Response must be a JSON object"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_router_handles_early_and_ambiguous_results(monkeypatch):
|
||||
target = FrameworkRegistry()
|
||||
target.begin_module("one", "modules.one")
|
||||
target.register_command(
|
||||
"one",
|
||||
lambda *_args: None,
|
||||
{"system": "system", "user_template": "{user_input}"},
|
||||
)
|
||||
target.finish_module()
|
||||
|
||||
call = AsyncMock(return_value={"error": "offline"})
|
||||
monkeypatch.setattr(parser, "parse_async", call)
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"error": "offline"
|
||||
}
|
||||
|
||||
call.return_value = {
|
||||
"needs_clarification": "which one?",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"needs_clarification": "which one?",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
|
||||
call.return_value = {"interaction_type": "one", "confidence": 0.2}
|
||||
ambiguous = await parser.parse_command_async("hello", target)
|
||||
assert "needs_clarification" in ambiguous
|
||||
|
||||
call.return_value = {"interaction_type": "missing", "confidence": 1.0}
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"error": "Unknown command type: missing"
|
||||
}
|
||||
|
||||
|
||||
def test_registered_validator_is_used(monkeypatch):
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
parser.register_validator("registered", lambda value: [] if value["ok"] else ["bad"])
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: '{"ok":true}')
|
||||
assert parser.parse("hello", "registered", prompt_override=prompt) == {"ok": True}
|
||||
78
tests/unit/test_postgres.py
Normal file
78
tests/unit/test_postgres.py
Normal file
@@ -0,0 +1,78 @@
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from core import postgres
|
||||
|
||||
|
||||
def test_safe_identifier_quotes_names_and_rejects_sql_expressions():
|
||||
assert postgres._safe_id("scheduled_jobs") == '"scheduled_jobs"'
|
||||
|
||||
for unsafe in ["jobs.id", "jobs; DROP TABLE jobs", "two words", "", 7, None]:
|
||||
with pytest.raises(ValueError, match="Invalid SQL identifier"):
|
||||
postgres._safe_id(unsafe)
|
||||
|
||||
|
||||
def test_order_clause_allows_only_identifiers_and_directions():
|
||||
assert postgres._order_clause(
|
||||
["created_at desc", ("id", "ASC")]
|
||||
) == '"created_at" DESC, "id" ASC'
|
||||
|
||||
unsafe_values = [
|
||||
"created_at DESC NULLS LAST",
|
||||
"created_at;drop DESC",
|
||||
[("created_at", "SIDEWAYS")],
|
||||
[("created_at", "ASC", "extra")],
|
||||
]
|
||||
for value in unsafe_values:
|
||||
with pytest.raises(ValueError):
|
||||
postgres._order_clause(value)
|
||||
|
||||
|
||||
def test_select_builds_parameterized_where_and_safe_order(monkeypatch):
|
||||
cursor = MagicMock()
|
||||
cursor.fetchall.return_value = [{"id": "job-1"}]
|
||||
|
||||
@contextmanager
|
||||
def fake_cursor():
|
||||
yield cursor
|
||||
|
||||
monkeypatch.setattr(postgres, "get_cursor", fake_cursor)
|
||||
|
||||
rows = postgres.select(
|
||||
"scheduled_jobs",
|
||||
where={"user_uuid": "user-1", "status": ("IN", ["pending", "running"])},
|
||||
order_by=[("run_at", "ASC"), ("id", "DESC")],
|
||||
limit=10,
|
||||
)
|
||||
|
||||
assert rows == [{"id": "job-1"}]
|
||||
query, params = cursor.execute.call_args.args
|
||||
assert 'FROM "scheduled_jobs"' in query
|
||||
assert '"user_uuid" = %(user_uuid_0)s' in query
|
||||
assert '"status" IN (%(status_1_0)s, %(status_1_1)s)' in query
|
||||
assert 'ORDER BY "run_at" ASC, "id" DESC' in query
|
||||
assert "LIMIT %(query_limit)s" in query
|
||||
assert params == {
|
||||
"user_uuid_0": "user-1",
|
||||
"status_1_0": "pending",
|
||||
"status_1_1": "running",
|
||||
"query_limit": 10,
|
||||
}
|
||||
|
||||
|
||||
def test_empty_update_and_delete_conditions_fail_before_opening_cursor(monkeypatch):
|
||||
cursor_factory = MagicMock(
|
||||
side_effect=AssertionError("a database cursor must not be opened")
|
||||
)
|
||||
monkeypatch.setattr(postgres, "get_cursor", cursor_factory)
|
||||
|
||||
with pytest.raises(ValueError, match="update data cannot be empty"):
|
||||
postgres.update("users", {}, {"id": "user-1"})
|
||||
with pytest.raises(ValueError, match="non-empty where"):
|
||||
postgres.update("users", {"timezone": "UTC"}, {})
|
||||
with pytest.raises(ValueError, match="non-empty where"):
|
||||
postgres.delete("users", {})
|
||||
|
||||
cursor_factory.assert_not_called()
|
||||
156
tests/unit/test_registry.py
Normal file
156
tests/unit/test_registry.py
Normal file
@@ -0,0 +1,156 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
import core.registry as registry_module
|
||||
from core.registry import FrameworkRegistry
|
||||
|
||||
|
||||
PROMPT = {"system": "Return JSON", "user_template": "Message: {user_input}"}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clean_global_registry():
|
||||
registry_module.reset_registry()
|
||||
yield
|
||||
registry_module.reset_registry()
|
||||
|
||||
|
||||
def test_discovery_is_alphabetical_ignores_private_packages_and_runs_once(monkeypatch):
|
||||
loaded = []
|
||||
|
||||
def feature(name):
|
||||
def register(target):
|
||||
loaded.append(name)
|
||||
target.describe(f"{name} feature")
|
||||
target.register_command(
|
||||
name,
|
||||
lambda _context, _parsed: None,
|
||||
PROMPT,
|
||||
description=f"Handle {name}",
|
||||
help_text=[f"use {name}"],
|
||||
)
|
||||
|
||||
return SimpleNamespace(register=register)
|
||||
|
||||
fake_modules = {
|
||||
"modules.alpha": feature("alpha"),
|
||||
"modules.zeta": feature("zeta"),
|
||||
}
|
||||
|
||||
def fake_import(name):
|
||||
if name == "modules":
|
||||
return SimpleNamespace(__path__=["unused"])
|
||||
return fake_modules[name]
|
||||
|
||||
discovered = [
|
||||
SimpleNamespace(name="zeta", ispkg=True),
|
||||
SimpleNamespace(name="_private", ispkg=True),
|
||||
SimpleNamespace(name="single_file", ispkg=False),
|
||||
SimpleNamespace(name="alpha", ispkg=True),
|
||||
]
|
||||
monkeypatch.setattr(registry_module.importlib, "import_module", fake_import)
|
||||
monkeypatch.setattr(
|
||||
registry_module.pkgutil, "iter_modules", lambda _path: discovered
|
||||
)
|
||||
|
||||
result = registry_module.discover_modules()
|
||||
|
||||
assert loaded == ["alpha", "zeta"]
|
||||
assert list(result.modules) == ["alpha", "zeta"]
|
||||
assert result.list_commands() == ["alpha", "zeta"]
|
||||
assert registry_module.discover_modules() is result
|
||||
assert loaded == ["alpha", "zeta"]
|
||||
|
||||
|
||||
def test_duplicate_names_and_malformed_registrations_are_rejected():
|
||||
target = FrameworkRegistry()
|
||||
handler = lambda _context, _parsed: None
|
||||
|
||||
target.begin_module("first", "modules.first")
|
||||
target.register_command("shared", handler, PROMPT)
|
||||
target.register_job("shared.job", handler)
|
||||
target.finish_module()
|
||||
|
||||
with pytest.raises(ValueError, match="Duplicate module name: first"):
|
||||
target.begin_module("first", "modules.again")
|
||||
|
||||
target.begin_module("second", "modules.second")
|
||||
with pytest.raises(ValueError, match="Duplicate command type: shared"):
|
||||
target.register_command("shared", handler, PROMPT)
|
||||
with pytest.raises(ValueError, match="Duplicate job type: shared.job"):
|
||||
target.register_job("shared.job", handler)
|
||||
with pytest.raises(TypeError, match="must be callable"):
|
||||
target.register_command("not_callable", None, PROMPT)
|
||||
with pytest.raises(TypeError, match="Validator .* must be callable"):
|
||||
target.register_command("bad_validator", handler, PROMPT, validator="bad")
|
||||
with pytest.raises(ValueError, match="system and user_template"):
|
||||
target.register_command("bad_prompt", handler, {"system": "only one"})
|
||||
target.finish_module()
|
||||
|
||||
with pytest.raises(RuntimeError, match="inside a module register"):
|
||||
target.describe("orphan metadata")
|
||||
|
||||
|
||||
def test_help_and_router_context_are_generated_from_sorted_metadata():
|
||||
target = FrameworkRegistry()
|
||||
handler = lambda _context, _parsed: None
|
||||
target.begin_module("examples", "modules.examples")
|
||||
target.register_command(
|
||||
"zeta",
|
||||
handler,
|
||||
PROMPT,
|
||||
description="Last command",
|
||||
)
|
||||
target.register_command(
|
||||
"alpha",
|
||||
handler,
|
||||
PROMPT,
|
||||
description="First command",
|
||||
help_text=["say alpha", "ask alpha for help"],
|
||||
)
|
||||
target.finish_module()
|
||||
|
||||
assert target.router_context() == (
|
||||
"- alpha: First command\n- zeta: Last command"
|
||||
)
|
||||
assert target.help_lines() == [
|
||||
"- say alpha",
|
||||
"- ask alpha for help",
|
||||
"- zeta: Last command",
|
||||
]
|
||||
|
||||
|
||||
def test_failed_forced_discovery_clears_partial_state_and_can_retry(monkeypatch):
|
||||
broken = SimpleNamespace(register=lambda target: target.describe("partial"))
|
||||
monkeypatch.setattr(
|
||||
registry_module.importlib,
|
||||
"import_module",
|
||||
lambda name: (
|
||||
SimpleNamespace(__path__=["unused"]) if name == "modules" else broken
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
registry_module.pkgutil,
|
||||
"iter_modules",
|
||||
lambda _path: [SimpleNamespace(name="broken", ispkg=True)],
|
||||
)
|
||||
|
||||
registry_module.discover_modules()
|
||||
assert registry_module.registry.modules == {
|
||||
"broken": {
|
||||
"name": "broken",
|
||||
"package": "modules.broken",
|
||||
"description": "partial",
|
||||
}
|
||||
}
|
||||
|
||||
del broken.register
|
||||
with pytest.raises(RuntimeError, match="must expose register"):
|
||||
registry_module.discover_modules(force=True)
|
||||
assert registry_module.registry.modules == {}
|
||||
|
||||
broken.register = lambda target: target.describe("recovered")
|
||||
assert registry_module.discover_modules().modules["broken"][
|
||||
"description"
|
||||
] == "recovered"
|
||||
220
tests/unit/test_reminders.py
Normal file
220
tests/unit/test_reminders.py
Normal file
@@ -0,0 +1,220 @@
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from modules.reminders.commands import _errorMessage, handleReminder, validateReminder
|
||||
from modules.reminders.service import _nextRun, normalizeRecurrence
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"command",
|
||||
[
|
||||
{"action": "list"},
|
||||
{"action": "cancel", "reminder_id": "reminder-1"},
|
||||
{"action": "set_timezone", "timezone": "America/Chicago"},
|
||||
{
|
||||
"action": "create",
|
||||
"message": "call home",
|
||||
"run_at": "2999-03-08T09:00:00-05:00",
|
||||
"recurrence": {"frequency": "weekly", "interval": 2},
|
||||
},
|
||||
{"needs_clarification": "What time should I use?"},
|
||||
],
|
||||
)
|
||||
def test_reminder_validator_accepts_supported_commands(command):
|
||||
assert validateReminder(command) == []
|
||||
|
||||
|
||||
def test_reminder_validator_reports_all_invalid_create_fields():
|
||||
errors = validateReminder(
|
||||
{
|
||||
"action": "create",
|
||||
"message": " ",
|
||||
"run_at": "2026-03-08T09:00:00",
|
||||
"recurrence": {"frequency": "hourly"},
|
||||
}
|
||||
)
|
||||
|
||||
assert "create requires a reminder message" in errors
|
||||
assert "run_at must include a timezone offset" in errors
|
||||
assert "recurrence frequency must be daily or weekly" in errors
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("value", "expected"),
|
||||
[
|
||||
(None, None),
|
||||
("once", None),
|
||||
("daily", {"frequency": "daily", "interval": 1}),
|
||||
(
|
||||
{"frequency": "WEEKLY", "interval": "3"},
|
||||
{"frequency": "weekly", "interval": 3},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_recurrence_normalization(value, expected):
|
||||
assert normalizeRecurrence(value) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value",
|
||||
[
|
||||
"hourly",
|
||||
{"frequency": "daily", "interval": 0},
|
||||
{"frequency": "weekly", "interval": 366},
|
||||
{"frequency": "daily", "interval": "many"},
|
||||
],
|
||||
)
|
||||
def test_recurrence_normalization_rejects_invalid_values(value):
|
||||
with pytest.raises(ValueError):
|
||||
normalizeRecurrence(value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("scheduled_for", "expected"),
|
||||
[
|
||||
# America/Chicago enters daylight time on March 8, 2026.
|
||||
(
|
||||
datetime(2026, 3, 7, 15, 0, tzinfo=timezone.utc),
|
||||
datetime(2026, 3, 8, 14, 0, tzinfo=timezone.utc),
|
||||
),
|
||||
# It returns to standard time on November 1, 2026.
|
||||
(
|
||||
datetime(2026, 10, 31, 14, 0, tzinfo=timezone.utc),
|
||||
datetime(2026, 11, 1, 15, 0, tzinfo=timezone.utc),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_next_run_preserves_local_wall_clock_across_dst(scheduled_for, expected):
|
||||
result = _nextRun(
|
||||
scheduled_for,
|
||||
{"frequency": "daily", "interval": 1},
|
||||
"America/Chicago",
|
||||
now=scheduled_for,
|
||||
)
|
||||
|
||||
assert result == expected
|
||||
assert result.astimezone(__import__("zoneinfo").ZoneInfo("America/Chicago")).hour == 9
|
||||
|
||||
|
||||
def test_next_run_skips_missed_intervals_after_downtime():
|
||||
scheduled_for = datetime(2026, 3, 7, 15, 0, tzinfo=timezone.utc)
|
||||
|
||||
result = _nextRun(
|
||||
scheduled_for,
|
||||
{"frequency": "daily", "interval": 1},
|
||||
"America/Chicago",
|
||||
now=datetime(2026, 3, 9, 14, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
assert result == datetime(2026, 3, 10, 14, 0, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def _context(response, status):
|
||||
api = SimpleNamespace(
|
||||
request=AsyncMock(return_value=(response, status)),
|
||||
timezone="UTC",
|
||||
)
|
||||
return SimpleNamespace(api=api, timezone="UTC", reply=AsyncMock())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_handler_reports_success_and_api_error():
|
||||
context = _context(
|
||||
{
|
||||
"message": "call home",
|
||||
"next_run_at": "2099-01-01T12:00:00+00:00",
|
||||
"recurrence": {"frequency": "daily"},
|
||||
},
|
||||
201,
|
||||
)
|
||||
parsed = {
|
||||
"action": "create",
|
||||
"message": "call home",
|
||||
"run_at": "2099-01-01T12:00:00Z",
|
||||
"recurrence": {"frequency": "daily"},
|
||||
}
|
||||
await handleReminder(context, parsed)
|
||||
assert "recurring" in context.reply.await_args.args[0]
|
||||
context.api.request.assert_awaited_once_with(
|
||||
"post",
|
||||
"/api/reminders",
|
||||
{
|
||||
"message": "call home",
|
||||
"run_at": "2099-01-01T12:00:00Z",
|
||||
"recurrence": {"frequency": "daily"},
|
||||
},
|
||||
)
|
||||
|
||||
context = _context({"error": "database unavailable"}, 503)
|
||||
await handleReminder(context, parsed)
|
||||
assert "database unavailable" in context.reply.await_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_handler_formats_results_and_empty_state():
|
||||
context = _context(
|
||||
{
|
||||
"reminders": [
|
||||
{
|
||||
"id": "reminder-one",
|
||||
"next_run_at": "2099-01-01T12:00:00+00:00",
|
||||
"message": "call home",
|
||||
}
|
||||
]
|
||||
},
|
||||
200,
|
||||
)
|
||||
await handleReminder(context, {"action": "list"})
|
||||
reply = context.reply.await_args.args[0]
|
||||
assert "Active reminders" in reply and "reminder-one" in reply
|
||||
|
||||
context = _context({"reminders": []}, 200)
|
||||
await handleReminder(context, {"action": "list"})
|
||||
assert context.reply.await_args.args[0] == "You have no active reminders."
|
||||
|
||||
context = _context({"error": "offline"}, 503)
|
||||
await handleReminder(context, {"action": "list"})
|
||||
assert "offline" in context.reply.await_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_and_timezone_handlers_update_context():
|
||||
context = _context({"message": "call home"}, 200)
|
||||
await handleReminder(
|
||||
context,
|
||||
{"action": "cancel", "reminder_id": "reminder-one"},
|
||||
)
|
||||
assert context.reply.await_args.args[0] == "Cancelled reminder: call home"
|
||||
context.api.request.assert_awaited_once_with(
|
||||
"delete", "/api/reminders/reminder-one"
|
||||
)
|
||||
|
||||
context = _context({"error": "not found"}, 404)
|
||||
await handleReminder(
|
||||
context,
|
||||
{"action": "cancel", "reminder_id": "missing"},
|
||||
)
|
||||
assert "not found" in context.reply.await_args.args[0]
|
||||
|
||||
context = _context({"timezone": "America/Chicago"}, 200)
|
||||
await handleReminder(
|
||||
context,
|
||||
{"action": "set_timezone", "timezone": "America/Chicago"},
|
||||
)
|
||||
assert context.timezone == "America/Chicago"
|
||||
assert context.api.timezone == "America/Chicago"
|
||||
|
||||
context = _context({}, 400)
|
||||
await handleReminder(
|
||||
context,
|
||||
{"action": "set_timezone", "timezone": "bad"},
|
||||
)
|
||||
assert "couldn't update" in context.reply.await_args.args[0]
|
||||
|
||||
|
||||
def test_error_message_handles_structured_and_unstructured_results():
|
||||
assert _errorMessage({"error": "detail"}, "fallback") == "fallback detail"
|
||||
assert _errorMessage("not an object", "fallback") == "fallback"
|
||||
382
tests/unit/test_security_core.py
Normal file
382
tests/unit/test_security_core.py
Normal file
@@ -0,0 +1,382 @@
|
||||
"""Unit tests for user, token, API-key, and provider identity boundaries."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
import datetime
|
||||
import json
|
||||
from unittest.mock import MagicMock
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from psycopg2.errors import UniqueViolation
|
||||
|
||||
from core import api_keys as apiKeys
|
||||
from core import auth, identity, users
|
||||
|
||||
|
||||
USER_ID = "00000000-0000-0000-0000-000000000101"
|
||||
KEY_ID = "00000000-0000-0000-0000-000000000201"
|
||||
|
||||
|
||||
class FakeCursor:
|
||||
def __init__(self, rows=None):
|
||||
self.rows = list(rows or [])
|
||||
self.executed = []
|
||||
|
||||
def execute(self, query, params=None):
|
||||
self.executed.append((" ".join(query.split()), params))
|
||||
|
||||
def fetchone(self):
|
||||
return self.rows.pop(0) if self.rows else None
|
||||
|
||||
|
||||
def _cursorContext(cursor):
|
||||
@contextmanager
|
||||
def context():
|
||||
yield cursor
|
||||
|
||||
return context
|
||||
|
||||
|
||||
def test_user_timezone_and_password_validation(monkeypatch):
|
||||
assert users.isValidTimezone("America/Chicago")
|
||||
assert not users.isValidTimezone("not/a-zone")
|
||||
assert users.normalizeTimezone(" UTC ") == "UTC"
|
||||
with pytest.raises(ValueError, match="IANA"):
|
||||
users.normalizeTimezone(None)
|
||||
|
||||
monkeypatch.setenv("DEFAULT_TIMEZONE", "America/Chicago")
|
||||
assert users.getDefaultTimezone() == "America/Chicago"
|
||||
assert users.validatePassword("long-enough")[0]
|
||||
assert not users.validatePassword("short")[0]
|
||||
assert not users.validatePassword("é" * 40)[0]
|
||||
assert users.validatePassword(None) == (False, ["password"])
|
||||
|
||||
|
||||
def test_user_lookup_helpers_and_registration_boundaries(monkeypatch):
|
||||
select = MagicMock(
|
||||
side_effect=[
|
||||
{"id": USER_ID, "username": "alice", "timezone": "UTC"},
|
||||
{"id": USER_ID},
|
||||
None,
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(users.postgres, "select_one", select)
|
||||
monkeypatch.setattr(users.postgres, "exists", MagicMock(side_effect=[False, True]))
|
||||
assert users.getUser(USER_ID)["username"] == "alice"
|
||||
assert users.getUserUUID("alice") == USER_ID
|
||||
assert users.getUserUUID("missing") is False
|
||||
assert users.isUsernameAvailable("alice")
|
||||
assert users.doesUserUUIDExist(USER_ID)
|
||||
|
||||
monkeypatch.setattr(users, "isUsernameAvailable", lambda _name: True)
|
||||
monkeypatch.setattr(users.bcrypt, "gensalt", lambda: b"salt")
|
||||
monkeypatch.setattr(users.bcrypt, "hashpw", lambda value, _salt: b"hash:" + value)
|
||||
create = MagicMock()
|
||||
monkeypatch.setattr(users, "createUser", create)
|
||||
assert users.registerUser(
|
||||
" alice ",
|
||||
"password123",
|
||||
{"timezone": "UTC", "id": "attacker", "admin": True},
|
||||
)
|
||||
created = create.call_args.args[0]
|
||||
assert created["username"] == "alice"
|
||||
assert created["password_hashed"] == b"hash:password123"
|
||||
assert created["timezone"] == "UTC"
|
||||
assert created["id"] != "attacker" and "admin" not in created
|
||||
|
||||
monkeypatch.setattr(users, "isUsernameAvailable", lambda _name: False)
|
||||
assert users.registerUser("alice", "password123") is False
|
||||
monkeypatch.setattr(users, "isUsernameAvailable", lambda _name: True)
|
||||
create.side_effect = UniqueViolation()
|
||||
assert users.registerUser("alice", "password123") is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("username", [None, " ", "x" * 256])
|
||||
def test_registration_rejects_invalid_usernames(username):
|
||||
with pytest.raises(ValueError, match="username"):
|
||||
users.registerUser(username, "password123")
|
||||
|
||||
|
||||
def test_user_creation_updates_password_and_delete(monkeypatch):
|
||||
monkeypatch.setattr(users, "getUser", lambda _user: {"id": USER_ID})
|
||||
update = MagicMock(return_value=[{"id": USER_ID}])
|
||||
delete = MagicMock(return_value=[{"id": USER_ID}])
|
||||
insert = MagicMock(return_value={"id": USER_ID})
|
||||
monkeypatch.setattr(users.postgres, "update", update)
|
||||
monkeypatch.setattr(users.postgres, "delete", delete)
|
||||
monkeypatch.setattr(users.postgres, "insert", insert)
|
||||
|
||||
assert users.updateUser(USER_ID, {"timezone": "UTC", "username": "ignored"})
|
||||
update.assert_called_with("users", {"timezone": "UTC"}, {"id": USER_ID})
|
||||
assert not users.updateUser(USER_ID, {"username": "ignored"})
|
||||
|
||||
monkeypatch.setattr(users.bcrypt, "gensalt", lambda: b"salt")
|
||||
monkeypatch.setattr(users.bcrypt, "hashpw", lambda _value, _salt: b"new-hash")
|
||||
assert users.changePassword(USER_ID, "new-password")
|
||||
assert users.deleteUser(USER_ID)
|
||||
|
||||
valid = {
|
||||
"id": USER_ID,
|
||||
"username": "alice",
|
||||
"password_hashed": b"hash",
|
||||
"timezone": "UTC",
|
||||
"unexpected": True,
|
||||
}
|
||||
assert users.createUser(valid) == {"id": USER_ID}
|
||||
assert "unexpected" not in insert.call_args.args[1]
|
||||
assert users.validateUser(valid)[0]
|
||||
assert not users.validateUser({"id": "bad", "timezone": "bad"})[0]
|
||||
|
||||
|
||||
def test_login_token_creation_decoding_and_protected_claims(monkeypatch):
|
||||
monkeypatch.setenv("JWT_SECRET", "test-secret")
|
||||
monkeypatch.setattr(users, "doesUserUUIDExist", lambda value: value == USER_ID)
|
||||
monkeypatch.setattr(users, "getUserFirstName", lambda _value: "Alice")
|
||||
encode = MagicMock(return_value="encoded-token")
|
||||
monkeypatch.setattr(auth.jwt, "encode", encode)
|
||||
|
||||
assert auth.createLoginToken(
|
||||
USER_ID,
|
||||
expiresIn=60,
|
||||
extraClaims={"provider": "discord", "sub": "attacker"},
|
||||
) == "encoded-token"
|
||||
payload = encode.call_args.args[0]
|
||||
assert payload["sub"] == USER_ID and payload["provider"] == "discord"
|
||||
assert payload["exp"] > payload["iat"]
|
||||
with pytest.raises(ValueError, match="positive"):
|
||||
auth.createLoginToken(USER_ID, expiresIn=0)
|
||||
with pytest.raises(ValueError, match="extra claims"):
|
||||
auth.createLoginToken(USER_ID, extraClaims="bad")
|
||||
|
||||
monkeypatch.setattr(auth.jwt, "decode", lambda *_args, **_kwargs: {"sub": USER_ID})
|
||||
assert auth.decodeLoginToken("token")["sub"] == USER_ID
|
||||
monkeypatch.setattr(auth.jwt, "decode", lambda *_args, **_kwargs: {"sub": 123})
|
||||
assert auth.decodeLoginToken("token") is None
|
||||
monkeypatch.delenv("JWT_SECRET")
|
||||
assert auth.decodeLoginToken("token") is None
|
||||
with pytest.raises(RuntimeError, match="JWT_SECRET"):
|
||||
auth.createLoginToken(USER_ID)
|
||||
|
||||
|
||||
def test_password_login_and_account_removal(monkeypatch):
|
||||
monkeypatch.setattr(users, "getUserUUID", lambda name: USER_ID if name == "alice" else False)
|
||||
monkeypatch.setattr(auth, "getUserpasswordHash", lambda _value: b"stored")
|
||||
monkeypatch.setattr(auth.bcrypt, "checkpw", lambda password, hashed: password == b"password123")
|
||||
monkeypatch.setattr(auth, "createLoginToken", lambda _value: "login-token")
|
||||
assert auth.getLoginToken("alice", "password123") == "login-token"
|
||||
assert auth.getLoginToken("alice", "wrong-password") is False
|
||||
assert auth.getLoginToken(None, "password123") is False
|
||||
|
||||
monkeypatch.setattr(users, "deleteUser", MagicMock(return_value=True))
|
||||
assert auth.unregisterUser(USER_ID, "password123")
|
||||
assert not auth.unregisterUser(USER_ID, "wrong-password")
|
||||
|
||||
|
||||
def test_bearer_authentication_and_principal_helpers(monkeypatch):
|
||||
assert auth.getBearerToken("Bearer token") == "token"
|
||||
assert auth.getBearerToken("basic token") is None
|
||||
assert auth.getBearerToken(None) is None
|
||||
|
||||
monkeypatch.setattr(auth, "decodeLoginToken", lambda _token: {"sub": USER_ID})
|
||||
principal = auth.authenticateBearerToken("Bearer a.b.c")
|
||||
assert principal["authentication"] == "jwt"
|
||||
assert auth.isUserPrincipal(principal, USER_ID, requireLogin=True)
|
||||
assert auth.authenticateBearerToken(
|
||||
"Bearer a.b.c", requiredScopes=["service:scope"]
|
||||
) is None
|
||||
|
||||
service = {
|
||||
"type": "service",
|
||||
"scopes": ["jobs:claim"],
|
||||
"authentication": "api_key",
|
||||
}
|
||||
monkeypatch.setattr(apiKeys, "authenticateApiKey", lambda *_args, **_kwargs: service)
|
||||
assert auth.authenticateBearerToken("Bearer service-key") == service
|
||||
assert auth.authenticateBearerToken("Bearer service-key", allowService=False) is None
|
||||
assert auth.hasServiceScope(service, "jobs:claim")
|
||||
assert not auth.hasServiceScope(principal, "jobs:claim")
|
||||
|
||||
|
||||
def test_api_key_normalization_public_shape_and_ids(monkeypatch):
|
||||
assert apiKeys._normalizeScopes("one, two one") == ["one", "two"]
|
||||
assert apiKeys._normalizeScopes(None) == []
|
||||
with pytest.raises(ValueError, match="scopes"):
|
||||
apiKeys._normalizeScopes(123)
|
||||
with pytest.raises(ValueError, match="future"):
|
||||
apiKeys._normalizeExpiry("2000-01-01T00:00:00Z")
|
||||
|
||||
record = {"id": KEY_ID, "key_hash": "secret", "scopes": '["one"]'}
|
||||
assert apiKeys._publicKey(record, includeSecret="raw") == {
|
||||
"id": KEY_ID,
|
||||
"scopes": ["one"],
|
||||
"key": "raw",
|
||||
}
|
||||
assert apiKeys._keyID(KEY_ID) == KEY_ID
|
||||
assert apiKeys._keyID("bad") is None
|
||||
monkeypatch.setattr(apiKeys.postgres, "select_one", MagicMock(return_value=record))
|
||||
assert apiKeys.getApiKey("bad") is None
|
||||
assert apiKeys.getApiKey(KEY_ID)["id"] == KEY_ID
|
||||
|
||||
|
||||
def test_api_key_creation_and_owner_rules(monkeypatch):
|
||||
monkeypatch.setattr(users, "doesUserUUIDExist", lambda value: value == USER_ID)
|
||||
secret = "x" * 32
|
||||
stored = {
|
||||
"id": KEY_ID,
|
||||
"name": "cli",
|
||||
"key_type": "user",
|
||||
"user_uuid": USER_ID,
|
||||
"key_hash": apiKeys._hashToken(secret),
|
||||
"scopes": [],
|
||||
}
|
||||
cursor = FakeCursor([stored])
|
||||
monkeypatch.setattr(apiKeys.postgres, "get_cursor", _cursorContext(cursor))
|
||||
created = apiKeys.createApiKey(
|
||||
"cli", "user", userUUID=USER_ID, secret=secret
|
||||
)
|
||||
assert created["key"] == secret and "key_hash" not in created
|
||||
assert cursor.executed[0][1]["key_prefix"] == secret[:20]
|
||||
|
||||
with pytest.raises(ValueError, match="user does not exist"):
|
||||
apiKeys.createUserApiKey(str(uuid.uuid4()), "missing")
|
||||
with pytest.raises(ValueError, match="cannot have a service"):
|
||||
apiKeys.createApiKey(
|
||||
"bad", "user", userUUID=USER_ID, serviceName="service"
|
||||
)
|
||||
with pytest.raises(ValueError, match="service name"):
|
||||
apiKeys.createApiKey("bad", "service", scopes=[])
|
||||
with pytest.raises(ValueError, match="cannot have a user"):
|
||||
apiKeys.createApiKey(
|
||||
"bad", "service", userUUID=USER_ID, serviceName="service"
|
||||
)
|
||||
|
||||
|
||||
def test_api_key_authentication_scope_expiry_and_revoke(monkeypatch):
|
||||
secret = "service-secret-that-is-at-least-32-characters"
|
||||
record = {
|
||||
"id": KEY_ID,
|
||||
"key_type": "service",
|
||||
"user_uuid": None,
|
||||
"service_name": "worker",
|
||||
"key_hash": apiKeys._hashToken(secret),
|
||||
"scopes": ["jobs:claim"],
|
||||
"expires_at": None,
|
||||
}
|
||||
select = MagicMock(return_value=[record])
|
||||
update = MagicMock()
|
||||
monkeypatch.setattr(apiKeys.postgres, "select", select)
|
||||
monkeypatch.setattr(apiKeys.postgres, "update", update)
|
||||
principal = apiKeys.authenticateApiKey(secret, ["jobs:claim"])
|
||||
assert principal["service_name"] == "worker"
|
||||
assert update.called
|
||||
assert apiKeys.authenticateApiKey("short") is None
|
||||
assert apiKeys.authenticateApiKey(secret, ["jobs:deliver"]) is None
|
||||
|
||||
expired = dict(record, expires_at=apiKeys._utcNow() - datetime.timedelta(seconds=1))
|
||||
select.return_value = [expired]
|
||||
assert apiKeys.authenticateApiKey(secret) is None
|
||||
select.return_value = []
|
||||
assert apiKeys.authenticateApiKey(secret) is None
|
||||
|
||||
execute = MagicMock(return_value=[{"id": KEY_ID}])
|
||||
monkeypatch.setattr(apiKeys.postgres, "execute", execute)
|
||||
assert apiKeys.revokeApiKey(KEY_ID, serviceName="worker")
|
||||
assert not apiKeys.revokeApiKey("invalid")
|
||||
|
||||
|
||||
def test_service_key_bootstrap_is_idempotent_configuration(monkeypatch):
|
||||
monkeypatch.delenv("BOT_API_KEY", raising=False)
|
||||
assert apiKeys.bootstrapServiceApiKey() is None
|
||||
secret = "configured-service-key-at-least-32-characters"
|
||||
monkeypatch.setenv("BOT_API_KEY", secret)
|
||||
monkeypatch.setenv("BOT_API_KEY_SCOPES", "discord:session, outbox:claim")
|
||||
record = {
|
||||
"id": KEY_ID,
|
||||
"key_hash": apiKeys._hashToken(secret),
|
||||
"scopes": json.dumps(["discord:session", "outbox:claim"]),
|
||||
}
|
||||
cursor = FakeCursor([record])
|
||||
monkeypatch.setattr(apiKeys.postgres, "get_cursor", _cursorContext(cursor))
|
||||
result = apiKeys.bootstrapServiceApiKey()
|
||||
assert result["scopes"] == ["discord:session", "outbox:claim"]
|
||||
assert len(cursor.executed) == 2
|
||||
|
||||
|
||||
def test_identity_normalization_enrollment_and_lookup(monkeypatch):
|
||||
assert identity._normalizeProvider(" Discord ") == "discord"
|
||||
assert identity._normalizeProviderUserID(123) == "123"
|
||||
assert identity._normalizeDisplayName(" Alice ") == "Alice"
|
||||
with pytest.raises(ValueError, match="provider"):
|
||||
identity._normalizeProvider("bad provider")
|
||||
with pytest.raises(ValueError, match="provider user ID"):
|
||||
identity._normalizeProviderUserID("")
|
||||
|
||||
monkeypatch.setenv("DISCORD_ENROLLMENT_MODE", "allowlist")
|
||||
monkeypatch.setenv("DISCORD_ALLOWLIST", "123, 456\n789")
|
||||
assert identity.getDiscordAllowlist() == {"123", "456", "789"}
|
||||
assert identity.isDiscordEnrollmentAllowed("456")
|
||||
assert not identity.isDiscordEnrollmentAllowed("000")
|
||||
monkeypatch.setenv("DISCORD_ENROLLMENT_MODE", "open")
|
||||
assert identity.isDiscordEnrollmentAllowed("000")
|
||||
|
||||
monkeypatch.setattr(
|
||||
identity.postgres,
|
||||
"select_one",
|
||||
MagicMock(return_value={"user_uuid": USER_ID}),
|
||||
)
|
||||
monkeypatch.setattr(users, "getUser", MagicMock(return_value={"id": USER_ID}))
|
||||
assert identity.getDiscordUser("123") == {"id": USER_ID}
|
||||
|
||||
|
||||
def test_link_provider_identity_existing_and_new_paths(monkeypatch):
|
||||
monkeypatch.setattr(users, "doesUserUUIDExist", lambda _value: True)
|
||||
existing = {
|
||||
"id": "identity-one",
|
||||
"user_uuid": USER_ID,
|
||||
"display_name": "Old",
|
||||
}
|
||||
monkeypatch.setattr(identity, "getProviderIdentity", MagicMock(return_value=existing))
|
||||
update = MagicMock(return_value=[dict(existing, display_name="New")])
|
||||
monkeypatch.setattr(identity.postgres, "update", update)
|
||||
assert identity.linkProviderIdentity(
|
||||
USER_ID, "discord", "123", displayName="New"
|
||||
)["display_name"] == "New"
|
||||
|
||||
identity.getProviderIdentity.return_value = dict(existing, user_uuid=str(uuid.uuid4()))
|
||||
with pytest.raises(ValueError, match="another user"):
|
||||
identity.linkProviderIdentity(USER_ID, "discord", "123")
|
||||
|
||||
identity.getProviderIdentity.return_value = None
|
||||
inserted = {"id": "identity-two", "user_uuid": USER_ID}
|
||||
monkeypatch.setattr(identity.postgres, "insert", MagicMock(return_value=inserted))
|
||||
assert identity.linkDiscordIdentity(USER_ID, "456") == inserted
|
||||
monkeypatch.setattr(identity.postgres, "delete", MagicMock(return_value=[inserted]))
|
||||
assert identity.unlinkProviderIdentity(USER_ID, "discord", "456")
|
||||
|
||||
|
||||
def test_get_or_create_discord_user_existing_denied_and_created(monkeypatch):
|
||||
monkeypatch.setattr(users, "getDefaultTimezone", lambda: "UTC")
|
||||
existingIdentity = {
|
||||
"id": "identity-one",
|
||||
"user_uuid": USER_ID,
|
||||
"display_name": "Old",
|
||||
}
|
||||
monkeypatch.setattr(
|
||||
identity, "getProviderIdentity", MagicMock(return_value=existingIdentity)
|
||||
)
|
||||
monkeypatch.setattr(users, "getUser", MagicMock(return_value={"id": USER_ID}))
|
||||
update = MagicMock()
|
||||
monkeypatch.setattr(identity.postgres, "update", update)
|
||||
assert identity.getOrCreateDiscordUser("123", "New") == {"id": USER_ID}
|
||||
assert update.called
|
||||
|
||||
identity.getProviderIdentity.return_value = None
|
||||
monkeypatch.setattr(identity, "isDiscordEnrollmentAllowed", lambda _value: False)
|
||||
assert identity.getOrCreateDiscordUser("999") is None
|
||||
|
||||
monkeypatch.setattr(identity, "isDiscordEnrollmentAllowed", lambda _value: True)
|
||||
cursor = FakeCursor([None, {"id": USER_ID, "timezone": "UTC"}])
|
||||
monkeypatch.setattr(identity.postgres, "get_cursor", _cursorContext(cursor))
|
||||
created = identity.getOrCreateDiscordUser("999", "New")
|
||||
assert created["id"] == USER_ID
|
||||
assert len(cursor.executed) == 4
|
||||
Reference in New Issue
Block a user