diff --git a/.env.example b/.env.example index 74e6c73141289e94a6eee61d506b7cd3ebca9687..4ad43b43977247c8eb87aa05157c92387916da24 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,13 @@ APP_NAME=MediaRouter +APP_VERSION=1.0.0 +# Docker sets this to production. Local development keeps development so the +# established SQLite workflow remains available. +APP_ENVIRONMENT=development +HOST=0.0.0.0 +PORT=7860 +# Comma-separated exact origins. Production requires at least one HTTPS origin +# such as the separately hosted Vercel frontend; wildcard origins are rejected. +CORS_ALLOWED_ORIGINS=http://localhost:3000 TEMP_DIR=./temp OUTPUT_DIR=./outputs TEMPLATE_DIR=./app/templates/categories @@ -17,6 +26,13 @@ FFMPEG_BINARY=ffmpeg FFPROBE_BINARY=ffprobe AUTH_ENABLED=true DATABASE_URL=sqlite+aiosqlite:///./data/mediarouter.db +# SQLite initializes this schema locally. PostgreSQL must receive the SQL files +# under app/security/migrations/ through your deployment migration process. +SECURITY_AUTO_MIGRATE=false +# PostgreSQL only: backend-only BYPASSRLS role used for authoritative users, +# memberships, API-key principals, and canonical asset records. +SECURITY_DATABASE_ROLE=mediarouter_security_service +SECURITY_ENFORCE_RLS=true AUTH_ROLE_SCOPES={} AUTH_BOOTSTRAP_KEY_HASH= AUTH_BOOTSTRAP_KEY_PREFIX= @@ -37,6 +53,14 @@ MCP_STDIO_API_KEY= SOCIAL_ENABLED=true # Example: postgresql+asyncpg://postgres:password@db.example:5432/postgres SOCIAL_DATABASE_URL= +# Required for PostgreSQL production. This URL must authenticate as a normal +# non-owner/non-BYPASSRLS API role. Do not use a Supabase service-role URL here. +SOCIAL_TENANT_DATABASE_ROLE=mediarouter_tenant +# Trusted backend-only worker connection. It must use the explicitly named +# BYPASSRLS role and must never be exposed through REST, MCP, n8n, SDKs, or UI. +SOCIAL_WORKER_DATABASE_URL= +SOCIAL_WORKER_DATABASE_ROLE=mediarouter_social_worker +SOCIAL_ENFORCE_RLS=true SOCIAL_AUTO_MIGRATE=false SOCIAL_WORKER_ENABLED=true SOCIAL_SCHEDULER_INTERVAL_SECONDS=30 @@ -53,6 +77,47 @@ SUPABASE_URL= SUPABASE_SERVICE_ROLE_KEY= SUPABASE_VAULT_ENABLED=false SOCIAL_OAUTH_REDIRECT_BASE_URL= + +# Provider-neutral generation runtime. WAN and FLUX are optional; neither is +# enabled merely by the shared timeout/retry settings below. +GENERATION_ENABLED=true +GENERATION_JOB_RETRY_LIMIT=3 +# Shared remote AI worker client defaults. No worker is configured by these +# values alone; do not add browser-visible worker URLs or credentials here. +AI_WORKER_CONNECT_TIMEOUT_SECONDS=10 +AI_WORKER_REQUEST_TIMEOUT_SECONDS=60 +AI_WORKER_READ_TIMEOUT_SECONDS=300 +AI_WORKER_MAX_RETRIES=3 +AI_WORKER_RETRY_BACKOFF_SECONDS=0.5 +# Optional authenticated WAN 2.2 image-to-video worker. Set both values only +# in backend/Hugging Face deployment secrets; never expose them to a browser, +# MCP client, n8n node, SDK output, or logs. WAN stays unavailable if either +# value is omitted or malformed, without blocking application startup. +WAN_SPACE_URL= +WAN_SPACE_TOKEN= +# Optional authenticated FLUX.2 Klein image-generation worker. Set both +# backend-only values in deployment secrets; never expose either value in a +# browser, MCP/n8n request, SDK result, audit record, or log. +FLUX_SPACE_URL= +FLUX_SPACE_TOKEN= +GENERATION_WORKER_ENABLED=true +GENERATION_WORKER_INTERVAL_SECONDS=5 +GENERATION_WORKER_POLL_BACKOFF_SECONDS=2 +GENERATION_WORKER_BATCH_SIZE=8 +GENERATION_JOB_STALE_AFTER_SECONDS=900 +# Content Studio persistence/render limits. Editor state and render jobs are +# durable PostgreSQL records; render workers use bounded temporary files only. +EDITOR_STATE_MAX_BYTES=1048576 +RENDER_WORKER_ENABLED=true +RENDER_WORKER_INTERVAL_SECONDS=2 +RENDER_JOB_STALE_AFTER_SECONDS=900 +RENDER_JOB_TIMEOUT_SECONDS=7200 +RENDER_JOB_RETRY_LIMIT=2 +RENDER_MAX_ACTIVE_JOBS_PER_PROJECT=1 +RENDER_MAX_TRACKS=32 +RENDER_MAX_CLIPS=500 +RENDER_MAX_DURATION_SECONDS=3600 +RENDER_MAX_INPUT_BYTES=4294967296 GOOGLE_CLIENT_ID= GOOGLE_CLIENT_SECRET= # YouTube Data API v3 worker controls. Keep the client secret backend-only. @@ -63,6 +128,11 @@ YOUTUBE_REQUEST_TIMEOUT_SECONDS=60 YOUTUBE_PROCESSING_POLL_SECONDS=30 META_CLIENT_ID= META_CLIENT_SECRET= +# Preferred Meta application names. META_CLIENT_* remains a compatibility +# alias for deployments created before this configuration contract. +META_APP_ID= +META_APP_SECRET= +META_GRAPH_API_VERSION=v25.0 TIKTOK_CLIENT_KEY= TIKTOK_CLIENT_SECRET= # Exact backend callback registered under TikTok Login Kit. It must be: diff --git a/Dockerfile b/Dockerfile index e17a4063bffb6c39315bcfdb1996ec8d5aede542..9770045b04b02e51218d886c7cba436198cd40fd 100644 --- a/Dockerfile +++ b/Dockerfile @@ -6,9 +6,9 @@ ENV PYTHONDONTWRITEBYTECODE=1 \ PIP_DISABLE_PIP_VERSION_CHECK=1 \ HF_HOME=/home/user/.cache/huggingface \ APP_NAME=MediaRouter \ + APP_ENVIRONMENT=production \ TEMP_DIR=/app/temp \ OUTPUT_DIR=/app/outputs \ - DATABASE_URL=sqlite+aiosqlite:////app/data/mediarouter.db \ PORT=7860 RUN apt-get update \ @@ -39,6 +39,8 @@ USER user EXPOSE 7860 +STOPSIGNAL SIGTERM + HEALTHCHECK --interval=30s --timeout=5s --start-period=20s --retries=3 \ CMD python -c "import urllib.request; urllib.request.urlopen('http://127.0.0.1:7860/health', timeout=4)" || exit 1 diff --git a/README.md b/README.md index 0cb08b8e927bdef8949dfcd9632af26d4f48fafa..f675ab3282af70dd7b0c8414530478bfae07a31f 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,7 @@ media-api/ ├── app/ │ ├── api/ # versioned, thin FastAPI routes │ ├── core/ # settings, errors, logging, response models +│ ├── generation/ # provider-neutral durable generation foundation │ ├── mcp/ # MCP server, registry, tools, resources, prompts │ ├── models/ # InputMedia and request/result models │ ├── operations/ # reusable FFmpeg operation functions @@ -83,6 +84,19 @@ media-api/ The top-level `api/`, `services/`, `operations/`, `workers/`, `core/`, and `models/` packages mirror the canonical `app/` modules as import-compatible entry points for integrations that use the requested layout. Runtime composition uses the single implementation under `app/`, so business logic is not duplicated. +## Generation foundation + +MediaRouter includes a tenant-scoped generation request/job foundation at +`/v1/generation`. Optional `wan` and `flux` providers implement audited WAN +2.2 image-to-video and FLUX.2 Klein image-generation worker contracts. Each +model stays unavailable until its backend-only URL/token are configured and +live readiness verifies its exact worker identity. Requests use an idempotency +key and may reference only canonical MediaRouter assets by opaque ID. See +[`docs/generation-foundation.md`](docs/generation-foundation.md), +[`docs/generation-wan.md`](docs/generation-wan.md), and +[`docs/generation-flux.md`](docs/generation-flux.md) for the state machine, +PostgreSQL migration, scopes, worker contracts, and security boundary. + Each request creates `TEMP_DIR//{uploads,outputs,logs}`. Completed files are atomically moved to `OUTPUT_DIR/` so they can be streamed. Both locations expire after `CLEANUP_MINUTES`; active requests are protected from the cleanup worker. ## Run locally @@ -110,19 +124,29 @@ docker run --rm -p 7860:7860 --env-file .env media-api The image uses `python:3.10-slim`, installs FFmpeg, FFprobe, ImageMagick and the system libraries needed by CTranslate2/faster-whisper, clears apt and pip caches, runs as UID 1000, exposes port 7860, and includes a container health check. +The Docker image sets `APP_ENVIRONMENT=production`. It intentionally has no +SQLite database default: configure external PostgreSQL/Supabase, apply the +explicit migrations, and set an exact Vercel CORS origin before startup. See +[the production deployment gate](docs/production-deployment-gate.md). + ## Deploy to Hugging Face Spaces 1. Create a new Space and select **Docker** as the SDK. 2. Push the contents of this directory to the Space repository. Keep the YAML block at the top of this README; `app_port` is already `7860`. -3. Generate the initial administrator key locally with `python -m app.security.cli generate-bootstrap --environment live`. Store `API_KEY` in a password manager. Add only the printed `AUTH_BOOTSTRAP_KEY_HASH`, `AUTH_BOOTSTRAP_KEY_PREFIX`, and `AUTH_BOOTSTRAP_ENVIRONMENT` under **Settings → Secrets**. Do not commit them. -4. Wait for the Docker build. The first Whisper call downloads the selected model to the Hugging Face cache. Persistent storage is optional, but avoids downloading models again after a cold rebuild. -5. Check the public `https://-.hf.space/health`, then use the saved key for every `/v1/*` or `/mcp/` request. +3. Apply the PostgreSQL migrations in the documented dependency order, using a controlled administrative connection. +4. Configure the required PostgreSQL/RLS/CORS variables and roles in Space Settings. Keep `SECURITY_AUTO_MIGRATE=false` and `SOCIAL_AUTO_MIGRATE=false`. +5. Generate the initial administrator key locally with `python -m app.security.cli generate-bootstrap --environment live`. Store `API_KEY` in a password manager. Add only the printed `AUTH_BOOTSTRAP_KEY_HASH`, `AUTH_BOOTSTRAP_KEY_PREFIX`, and `AUTH_BOOTSTRAP_ENVIRONMENT` under **Settings → Secrets**. Do not commit them. +6. Wait for the Docker build. The first Whisper call downloads the selected model to the Hugging Face cache. Persistent storage is optional for the model cache, but avoids downloading models again after a cold rebuild. +7. Check the public `https://-.hf.space/health`, then run `scripts/deployment_smoke.py` against the Space before accepting traffic. If using the supplied `media-api-huggingface.zip`, extract it first and push the extracted files so `Dockerfile` and this `README.md` are at the Space repository root. Do not commit the ZIP as the only repository file. Use one Uvicorn process in a CPU Space. FFmpeg and Whisper concurrency is managed in-process by `MAX_WORKERS`; multiple Uvicorn workers duplicate Whisper models and memory. -For durable keys, audit records, and rate aggregates, attach Hugging Face persistent storage and set `DATABASE_URL=sqlite+aiosqlite:////data/mediarouter.db`. The default `./data/mediarouter.db` is appropriate locally but follows the Space filesystem lifecycle. The security layer uses SQLAlchemy so a future external database migration does not change authentication contracts; this release includes and supports the `aiosqlite` driver. +Production keys, workspaces, projects, audit records, and rate aggregates live +in external PostgreSQL/Supabase. SQLite remains a local-development and test +backend only. Media files in `TEMP_DIR`/`OUTPUT_DIR` are expiring working data; +see the deployment gate for their filesystem lifecycle. ## Connect the Vercel frontend to Hugging Face @@ -155,13 +179,21 @@ The bootstrap fields are optional after the first successful start. If a Space f Recommended backend deployment settings are: ```env +APP_ENVIRONMENT=production BASE_URL=https://basyx-mediarouter.hf.space -DATABASE_URL=sqlite+aiosqlite:////data/mediarouter.db +DATABASE_URL=postgresql+asyncpg://:@/ +SECURITY_DATABASE_ROLE=mediarouter_security_service +SECURITY_ENFORCE_RLS=true +SECURITY_AUTO_MIGRATE=false +CORS_ALLOWED_ORIGINS=https://.vercel.app +SOCIAL_AUTO_MIGRATE=false AUTH_ROLE_SCOPES={} MCP_STDIO_API_KEY= ``` -`DATABASE_URL` uses `/data` so keys, audit logs, and rate-limit state survive only when Hugging Face persistent storage is attached. `BASE_URL` is the public Hugging Face origin, not the Vercel frontend URL. +When social automation is enabled, configure its existing separate tenant and +worker PostgreSQL URLs/roles as described in the deployment gate. `BASE_URL` +is the public Hugging Face origin, not the Vercel frontend URL. ### Vercel environment variables @@ -284,6 +316,7 @@ Scopes are enforced before route execution. Explicit scopes are combined with th | Jobs | `jobs:read`, `jobs:create`, `jobs:cancel` | | Assets | `assets:read`, `assets:write`, `assets:delete` | | MCP | `mcp:read`, `mcp:execute` | +| Generation foundation | `generation:providers:read`, `generation:requests:read`, `generation:requests:create`, `generation:jobs:cancel` | | System | `system:read` | | Administration | `admin` | @@ -1284,6 +1317,12 @@ curl -X POST https://your-space.hf.space/v1/social/posts \ | Variable | Default | Purpose | |---|---:|---| +| `APP_NAME` | `MediaRouter` | OpenAPI/application name | +| `APP_VERSION` | `1.0.0` | Runtime and health version | +| `APP_ENVIRONMENT` | `development` (`production` in Docker) | Enables fail-closed production configuration validation | +| `HOST` | `0.0.0.0` | Documented bind host; Docker command binds explicitly to this address | +| `PORT` | `7860` | Hugging Face application port | +| `CORS_ALLOWED_ORIGINS` | empty | Comma-separated exact frontend origins; production requires HTTPS and rejects wildcard CORS | | `TEMP_DIR` | `./temp` | Request workspaces and in-progress files | | `OUTPUT_DIR` | `./outputs` | Published files served by download URLs | | `TEMPLATE_DIR` | built-in `app/templates/categories` | Recursively scanned YAML workflow catalog | @@ -1301,7 +1340,10 @@ curl -X POST https://your-space.hf.space/v1/social/posts \ | `FFMPEG_BINARY` | `ffmpeg` | FFmpeg executable name/path | | `FFPROBE_BINARY` | `ffprobe` | FFprobe executable name/path | | `AUTH_ENABLED` | `true` | Fail-closed API-key enforcement; disable only for isolated development/tests | -| `DATABASE_URL` | `sqlite+aiosqlite:///./data/mediarouter.db` | Hash-only keys, audit records, and rate aggregates | +| `DATABASE_URL` | `sqlite+aiosqlite:///./data/mediarouter.db` | Backend-only security/tenant/asset store; bare `postgresql://` is normalized to `postgresql+asyncpg://` | +| `SECURITY_AUTO_MIGRATE` | `false` | Permit metadata creation for controlled local/testing use; PostgreSQL production must apply `app/security/migrations/` | +| `SECURITY_DATABASE_ROLE` | empty | Required with PostgreSQL RLS; backend-only role for authoritative tenancy and canonical asset administration | +| `SECURITY_ENFORCE_RLS` | `true` | Fail startup if the security-store role cannot administer its forced-RLS tables | | `AUTH_ROLE_SCOPES` | `{}` | JSON custom role-to-scope mappings | | `AUTH_BOOTSTRAP_KEY_HASH` | empty | SHA-256 hash for first-start administrator | | `AUTH_BOOTSTRAP_KEY_PREFIX` | empty | Safe display prefix matching the bootstrap key | @@ -1314,8 +1356,28 @@ curl -X POST https://your-space.hf.space/v1/social/posts \ | `AUTH_DEFAULT_PROCESSING_BYTES_PER_DAY` | `107374182400` | Default daily uploaded processing bytes | | `AUTH_TRUST_PROXY_HEADERS` | `true` | Use first `X-Forwarded-For` address for audits behind HF/Vercel | | `MCP_STDIO_API_KEY` | empty | Existing API key required by authenticated standalone stdio MCP | +| `GENERATION_ENABLED` | `true` | Enable the additive durable generation domain; it does not configure a model by itself | +| `GENERATION_JOB_RETRY_LIMIT` | `3` | Bounded durable submission retry count; the initial submission is separate | +| `AI_WORKER_CONNECT_TIMEOUT_SECONDS` | `10` | Remote generation-worker connection timeout | +| `AI_WORKER_REQUEST_TIMEOUT_SECONDS` | `60` | Remote generation-worker request/write timeout | +| `AI_WORKER_READ_TIMEOUT_SECONDS` | `300` | Remote generation-worker response/read timeout | +| `AI_WORKER_MAX_RETRIES` | `3` | Bounded transport retries for safe idempotent worker operations | +| `AI_WORKER_RETRY_BACKOFF_SECONDS` | `0.5` | Base exponential backoff for worker transport retries | +| `WAN_SPACE_URL` | empty | Trusted backend-only WAN worker origin; requires `WAN_SPACE_TOKEN` and is never client-selectable | +| `WAN_SPACE_TOKEN` | empty | Backend-only Bearer token for the WAN worker; never return, log, or store in browser/MCP/n8n/SDK output | +| `FLUX_SPACE_URL` | empty | Trusted backend-only FLUX worker origin; requires `FLUX_SPACE_TOKEN` and is never client-selectable | +| `FLUX_SPACE_TOKEN` | empty | Backend-only Bearer token for FLUX; never return, log, or store in browser/MCP/n8n/SDK output | +| `GENERATION_WORKER_ENABLED` | `true` | Run durable generation dispatch, polling, and output-ingestion worker | +| `GENERATION_WORKER_INTERVAL_SECONDS` | `5` | Generation worker dispatch/reconciliation interval | +| `GENERATION_WORKER_POLL_BACKOFF_SECONDS` | `2` | Base bounded reconciliation poll backoff after a transient worker error | +| `GENERATION_WORKER_BATCH_SIZE` | `8` | Maximum generation jobs claimed per dispatch/reconciliation cycle | +| `GENERATION_JOB_STALE_AFTER_SECONDS` | `900` | Submission ambiguity threshold and durable reconciliation lease duration | | `SOCIAL_ENABLED` | `true` | Enable the additive social domain; existing media APIs remain independent | | `SOCIAL_DATABASE_URL` | `DATABASE_URL` | Async SQLAlchemy URL; use Supabase/Postgres in production | +| `SOCIAL_TENANT_DATABASE_ROLE` | empty | Required with PostgreSQL RLS; non-owner/non-`BYPASSRLS` role for API tenant sessions | +| `SOCIAL_WORKER_DATABASE_URL` | empty | Backend-only trusted worker connection; required for PostgreSQL scheduler/Vault usage | +| `SOCIAL_WORKER_DATABASE_ROLE` | empty | Expected `BYPASSRLS` role for the worker URL; checked at social startup | +| `SOCIAL_ENFORCE_RLS` | `true` | Fail startup if PostgreSQL API/worker role separation cannot be verified | | `SOCIAL_AUTO_MIGRATE` | `false` | Local/test metadata creation only; never use for production migration management | | `SOCIAL_WORKER_ENABLED` | `true` | Run durable scheduler/publisher claim loop when schema is ready | | `SOCIAL_SCHEDULER_INTERVAL_SECONDS` | `30` | Scheduler polling interval | diff --git a/app/ai/__init__.py b/app/ai/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..09a37466e95600db1f1c9076e3dfcbcedcbb30e2 --- /dev/null +++ b/app/ai/__init__.py @@ -0,0 +1 @@ +"""Provider-neutral AI Studio application boundary.""" diff --git a/app/ai/api.py b/app/ai/api.py new file mode 100644 index 0000000000000000000000000000000000000000..14d1e2dd657e819b78ef0c2e026df5b3700abbb0 --- /dev/null +++ b/app/ai/api.py @@ -0,0 +1,87 @@ +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Header, Path, Query, Request, status + +from app.ai.schemas import AiCapabilities, AiGenerationRequest, AiHistory, AiJob +from app.security.errors import ForbiddenError + +router = APIRouter(prefix="/v1/ai", tags=["ai"]) + + +def _identity(request: Request) -> tuple[str, str, str | None, str]: + context = request.state.auth + if not context.workspace_id or not context.user_id: + raise ForbiddenError + return ( + context.workspace_id, + context.user_id, + context.api_key_id, + request.state.request_id, + ) + + +@router.get("/capabilities", response_model=AiCapabilities) +async def capabilities(request: Request) -> AiCapabilities: + return request.app.state.container.ai.capabilities() + + +@router.get("/jobs", response_model=AiHistory) +async def list_jobs( + request: Request, + offset: Annotated[int, Query(ge=0)] = 0, + limit: Annotated[int, Query(ge=1, le=100)] = 25, +) -> AiHistory: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.ai.history( + workspace_id=workspace_id, + user_id=user_id, + offset=offset, + limit=limit, + ) + + +@router.post("/jobs", response_model=AiJob, status_code=status.HTTP_202_ACCEPTED) +async def create_job( + request: Request, + payload: AiGenerationRequest, + idempotency_key: Annotated[str, Header(alias="Idempotency-Key", min_length=8, max_length=255)], +) -> AiJob: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.ai.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + payload=payload, + idempotency_key=idempotency_key, + ) + + +@router.get("/jobs/{generation_id}", response_model=AiJob) +async def get_job( + request: Request, + generation_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> AiJob: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.ai.get( + workspace_id=workspace_id, + user_id=user_id, + generation_id=generation_id, + ) + + +@router.post("/jobs/{generation_id}/cancel", response_model=AiJob) +async def cancel_job( + request: Request, + generation_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> AiJob: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.ai.cancel( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + generation_id=generation_id, + ) diff --git a/app/ai/schemas.py b/app/ai/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..8e79551d2a5e4444d0f9b137ea0dfc98051b4322 --- /dev/null +++ b/app/ai/schemas.py @@ -0,0 +1,171 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Annotated, Literal +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field + + +AiCategory = Literal["generate", "transform", "understand", "create"] +AiOperation = Literal["generate_image", "generate_video"] +AiJobStatus = Literal[ + "queued", "processing", "retrying", "completed", "failed", "cancelling", "cancelled" +] + + +class AiModel(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + display_name: str + operation: AiOperation + input_types: list[str] + output_types: list[str] + input_schema: dict[str, object] = Field(default_factory=dict) + available: bool + provider_display_name: str + pricing: None = None + + +class AiTool(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + category: AiCategory + operation: AiOperation + name: str + description: str + input_types: list[str] + output_types: list[str] + required_permission: str + job_type: str = "generation" + supports_project_context: bool = True + supports_asset_input: bool + supports_prompt_input: bool = True + supports_multiple_inputs: bool = False + supports_batch: bool = False + available: bool + models: list[AiModel] = Field(default_factory=list) + + +class AiProvider(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + display_name: str + available: bool + operations: list[AiOperation] = Field(default_factory=list) + supports_cancellation: bool = False + + +class AiCapabilities(BaseModel): + model_config = ConfigDict(extra="forbid") + + available: bool + categories: list[AiCategory] + tools: list[AiTool] + providers: list[AiProvider] + permissions: list[str] + feature_flag: str = "ai" + + +class AiOutputPreferences(BaseModel): + model_config = ConfigDict(extra="forbid") + + attach_to_project: bool = True + + +class AiImageParameters(BaseModel): + model_config = ConfigDict(extra="forbid") + + mode: Literal["fast", "quality"] = "fast" + seed: int = Field(default=42, ge=0, le=2_147_483_647) + randomize_seed: bool = False + width: int = Field(default=1024, ge=256, le=1024, multiple_of=8) + height: int = Field(default=1024, ge=256, le=1024, multiple_of=8) + steps: int = Field(default=4, ge=1, le=100) + guidance: float = Field(default=1.0, ge=0.0, le=10.0) + enhance_prompt: bool = False + + +class AiVideoParameters(BaseModel): + model_config = ConfigDict(extra="forbid") + + negative_prompt: str | None = Field(default=None, max_length=4_000) + duration_seconds: float = Field(default=5.0, ge=0.5, le=5.0) + steps: int = Field(default=4, ge=1, le=30) + guidance: float = Field(default=1.0, ge=0.0, le=10.0) + secondary_guidance: float = Field(default=1.0, ge=0.0, le=10.0) + seed: int = Field(default=42, ge=0, le=2_147_483_647) + randomize_seed: bool = False + + +class AiGenerateImageRequest(BaseModel): + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + operation: Literal["generate_image"] + prompt: str = Field(min_length=1, max_length=4_000) + model: str | None = Field(default=None, min_length=1, max_length=255) + project_id: UUID | None = None + source_asset_ids: list[UUID] = Field(default_factory=list, max_length=1) + parameters: AiImageParameters = Field(default_factory=AiImageParameters) + output_preferences: AiOutputPreferences = Field(default_factory=AiOutputPreferences) + + +class AiGenerateVideoRequest(BaseModel): + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + operation: Literal["generate_video"] + prompt: str = Field(min_length=1, max_length=4_000) + model: str | None = Field(default=None, min_length=1, max_length=255) + project_id: UUID | None = None + source_asset_ids: list[UUID] = Field(min_length=1, max_length=1) + parameters: AiVideoParameters = Field(default_factory=AiVideoParameters) + output_preferences: AiOutputPreferences = Field(default_factory=AiOutputPreferences) + + +AiGenerationRequest = Annotated[ + AiGenerateImageRequest | AiGenerateVideoRequest, + Field(discriminator="operation"), +] + + +class AiOutput(BaseModel): + model_config = ConfigDict(extra="forbid") + + asset_id: str + media_type: Literal["image", "video"] + + +class AiJob(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + generation_id: str + operation: AiOperation + status: AiJobStatus + prompt: str + project_id: str | None = None + source_asset_id: str | None = None + output: AiOutput | None = None + model: str + provider_display_name: str + progress: None = None + error_code: str | None = None + error_message: str | None = None + retryable: bool = False + usage: None = None + estimated_cost: None = None + actual_cost: None = None + created_at: datetime + updated_at: datetime + completed_at: datetime | None = None + + +class AiHistory(BaseModel): + model_config = ConfigDict(extra="forbid") + + items: list[AiJob] + offset: int + limit: int diff --git a/app/ai/service.py b/app/ai/service.py new file mode 100644 index 0000000000000000000000000000000000000000..fd7dcd0d09b23ba3932e9fa2c995f3b6b4e03e6f --- /dev/null +++ b/app/ai/service.py @@ -0,0 +1,378 @@ +from __future__ import annotations + +from app.ai.schemas import ( + AiCapabilities, + AiGenerateImageRequest, + AiGenerateVideoRequest, + AiGenerationRequest, + AiHistory, + AiJob, + AiModel, + AiOutput, + AiProvider, + AiTool, +) +from app.generation.domain.enums import ( + GenerationModality, + GenerationRequestStatus, +) +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationValidationError, +) +from app.generation.model_registry import GenerationModelView +from app.generation.schemas.requests import ( + FluxGenerationOptions, + GenerationRequestCreate, + GenerationRequestView, + WanGenerationOptions, +) +from app.generation.services.generation_service import GenerationService +from app.projects.schemas import ProjectStatus +from app.projects.services.project_service import ProjectService +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService + + +_OPERATION_MODALITY = { + "generate_image": GenerationModality.IMAGE, + "generate_video": GenerationModality.VIDEO, +} + + +class AiStudioService: + """Capability-driven facade over the existing durable generation domain.""" + + def __init__( + self, + generation: GenerationService, + projects: ProjectService, + assets: CanonicalAssetService, + audit: AuditService, + ) -> None: + self.generation = generation + self.projects = projects + self.assets = assets + self.audit = audit + + def capabilities(self) -> AiCapabilities: + models = self.generation.models.list() + tools = [ + self._tool( + operation="generate_image", + name="Generate Image", + description="Create or transform a canonical image with an available image model.", + models=[item for item in models if item.model.modality is GenerationModality.IMAGE], + input_types=["text", "image"], + output_types=["image"], + ), + self._tool( + operation="generate_video", + name="Generate Video", + description="Create a video from a canonical project image.", + models=[item for item in models if item.model.modality is GenerationModality.VIDEO], + input_types=["text", "image"], + output_types=["video"], + ), + ] + providers = [] + for adapter in self.generation.providers.list(): + provider_models = [item for item in models if item.provider_id == adapter.provider] + operations = sorted({self._operation_for_model(item) for item in provider_models}) + providers.append( + AiProvider( + id=adapter.provider, + display_name=adapter.capabilities.name, + available=bool( + self.generation.ready + and adapter.available + and any(item.available for item in provider_models) + ), + operations=operations, + supports_cancellation=adapter.capabilities.supports_cancellation, + ) + ) + return AiCapabilities( + available=any(tool.available for tool in tools), + categories=["generate"], + tools=tools, + providers=providers, + permissions=[ + "ai:read", + "ai:generate", + "ai:transform", + "ai:analyze", + "ai:create", + ], + ) + + async def create( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str | None, + request_id: str, + payload: AiGenerationRequest, + idempotency_key: str, + ) -> AiJob: + project_id = str(payload.project_id) if payload.project_id else None + source_asset_id = str(payload.source_asset_ids[0]) if payload.source_asset_ids else None + if payload.project_id is not None: + project = await self.projects.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=project_id or "", + ) + if project.status is not ProjectStatus.ACTIVE: + raise GenerationValidationError("AI jobs require an active project.") + if source_asset_id is not None: + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=source_asset_id, + ) + except CanonicalAssetNotFoundError as exc: + raise GenerationValidationError( + "AI source asset was not found in this workspace." + ) from exc + if project_id is not None and asset.project_id != project_id: + raise GenerationValidationError( + "AI source asset must belong to the selected project." + ) + model = self._select_model(payload.operation, payload.model) + request = self._generation_request(payload, model) + created = await self.generation.create( + workspace_id=workspace_id, + user_id=user_id, + payload=request, + idempotency_key=idempotency_key, + project_id=(project_id if payload.output_preferences.attach_to_project else None), + product_surface="ai_studio", + ) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="ai.generation_requested", + entity_type="generation_job", + entity_id=created.job.id, + metadata={ + "operation": payload.operation, + "model": created.model_id, + "project_id": project_id, + }, + ) + return self._job(created) + + async def history( + self, + *, + workspace_id: str, + user_id: str, + offset: int, + limit: int, + ) -> AiHistory: + items = await self.generation.list_requests( + workspace_id, + user_id, + offset=offset, + limit=limit, + product_surface="ai_studio", + ) + return AiHistory( + items=[self._job(item) for item in items], + offset=offset, + limit=limit, + ) + + async def get(self, *, workspace_id: str, user_id: str, generation_id: str) -> AiJob: + request = await self.generation.get_request(workspace_id, user_id, generation_id) + if request.product_surface != "ai_studio": + raise GenerationValidationError("AI generation was not found.") + return self._job(request) + + async def cancel( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str | None, + request_id: str, + generation_id: str, + ) -> AiJob: + request = await self.generation.get_request(workspace_id, user_id, generation_id) + if request.product_surface != "ai_studio": + raise GenerationValidationError("AI generation was not found.") + await self.generation.cancel(workspace_id, user_id, request.job.id) + updated = await self.generation.get_request(workspace_id, user_id, generation_id) + if updated.status is GenerationRequestStatus.CANCELLED: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="ai.generation_cancelled", + entity_type="generation_job", + entity_id=request.job.id, + metadata={"operation": self._operation_for_request(request)}, + ) + return self._job(updated) + + def _select_model(self, operation: str, requested_model: str | None) -> GenerationModelView: + modality = _OPERATION_MODALITY[operation] + matches = [ + model + for model in self.generation.models.list() + if model.model.modality is modality + and model.available + and (requested_model is None or model.model.id == requested_model) + ] + if not matches: + raise GenerationCapabilityUnsupportedError( + "No ready model supports the requested AI operation." + ) + if requested_model is not None and len(matches) != 1: + raise GenerationCapabilityUnsupportedError("The requested AI model is not available.") + return sorted(matches, key=lambda item: (item.provider_id, item.model.id))[0] + + @staticmethod + def _generation_request( + payload: AiGenerationRequest, model: GenerationModelView + ) -> GenerationRequestCreate: + source_asset_id = str(payload.source_asset_ids[0]) if payload.source_asset_ids else None + if isinstance(payload, AiGenerateImageRequest): + options = payload.parameters + return GenerationRequestCreate( + provider=model.provider_id, + model_id=model.model.id, + modality=GenerationModality.IMAGE, + prompt=payload.prompt, + input_asset_id=source_asset_id, + flux=FluxGenerationOptions( + mode_choice=( + "Base (50 steps)" if options.mode == "quality" else "Distilled (4 steps)" + ), + seed=options.seed, + randomize_seed=options.randomize_seed, + width=options.width, + height=options.height, + num_inference_steps=options.steps, + guidance_scale=options.guidance, + prompt_upsampling=options.enhance_prompt, + ), + ) + if isinstance(payload, AiGenerateVideoRequest): + options = payload.parameters + return GenerationRequestCreate( + provider=model.provider_id, + model_id=model.model.id, + modality=GenerationModality.VIDEO, + prompt=payload.prompt, + input_asset_id=source_asset_id, + wan=WanGenerationOptions( + negative_prompt=options.negative_prompt, + duration_seconds=options.duration_seconds, + steps=options.steps, + guidance_scale=options.guidance, + guidance_scale_2=options.secondary_guidance, + seed=options.seed, + randomize_seed=options.randomize_seed, + ), + ) + raise GenerationCapabilityUnsupportedError("AI operation is unsupported.") + + def _tool( + self, + *, + operation: str, + name: str, + description: str, + models: list[GenerationModelView], + input_types: list[str], + output_types: list[str], + ) -> AiTool: + public_models = [self._model(item, operation) for item in models] + return AiTool( + id=operation.replace("_", "-"), + category="generate", + operation=operation, + name=name, + description=description, + input_types=input_types, + output_types=output_types, + required_permission="ai:generate", + supports_asset_input=any(item.model.input_asset_supported for item in models), + available=any(item.available and self.generation.ready for item in models), + models=public_models, + ) + + def _model(self, item: GenerationModelView, operation: str) -> AiModel: + provider = self.generation.providers.get(item.provider_id) + return AiModel( + id=item.model.id, + display_name=item.model.name, + operation=operation, + input_types=(["text", "image"] if item.model.input_asset_supported else ["text"]), + output_types=[item.model.modality.value], + input_schema=dict(item.model.input_schema), + available=bool(self.generation.ready and item.available), + provider_display_name=provider.capabilities.name, + ) + + @staticmethod + def _operation_for_model(item: GenerationModelView) -> str: + return ( + "generate_image" + if item.model.modality is GenerationModality.IMAGE + else "generate_video" + ) + + @staticmethod + def _operation_for_request(item: GenerationRequestView) -> str: + return "generate_image" if item.modality is GenerationModality.IMAGE else "generate_video" + + def _job(self, item: GenerationRequestView) -> AiJob: + operation = self._operation_for_request(item) + status_map = { + GenerationRequestStatus.QUEUED: "queued", + GenerationRequestStatus.SUBMITTING: "processing", + GenerationRequestStatus.RUNNING: "processing", + GenerationRequestStatus.RETRYING: "retrying", + GenerationRequestStatus.SUCCEEDED: "completed", + GenerationRequestStatus.FAILED: "failed", + GenerationRequestStatus.CANCEL_REQUESTED: "cancelling", + GenerationRequestStatus.CANCELLED: "cancelled", + } + provider = self.generation.providers.get(item.provider) + output = None + if item.job.output_asset_id and item.status is GenerationRequestStatus.SUCCEEDED: + output = AiOutput( + asset_id=item.job.output_asset_id, + media_type=("image" if item.modality is GenerationModality.IMAGE else "video"), + ) + retryable_codes = { + "GENERATION_MODEL_UNAVAILABLE", + "GENERATION_SUBMISSION_RETRYING", + "GENERATION_PROVIDER_STATUS_FAILED", + } + return AiJob( + id=item.job.id, + generation_id=item.id, + operation=operation, + status=status_map[item.status], + prompt=item.prompt, + project_id=item.project_id, + source_asset_id=item.input_asset_id, + output=output, + model=item.model_id, + provider_display_name=provider.capabilities.name, + error_code=item.job.error_code, + error_message=item.job.error_message, + retryable=item.job.error_code in retryable_codes, + created_at=item.created_at, + updated_at=item.updated_at, + completed_at=item.completed_at, + ) diff --git a/app/analytics/__init__.py b/app/analytics/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..9fa75452fb7d144d8ba7eaa94dadb40138b91c96 --- /dev/null +++ b/app/analytics/__init__.py @@ -0,0 +1 @@ +"""Workspace analytics domain built on authoritative Social provider data.""" diff --git a/app/analytics/api.py b/app/analytics/api.py new file mode 100644 index 0000000000000000000000000000000000000000..20dffea67a72eb571eede692e1f09e40f7fc62c4 --- /dev/null +++ b/app/analytics/api.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Depends, Header, Query, Request, status + +from app.analytics.schemas import ( + AnalyticsCapabilities, + AnalyticsOverview, + AnalyticsPostList, + AnalyticsQuery, + AnalyticsSyncList, + AnalyticsSyncRequest, + AnalyticsSyncRunView, + AnalyticsTimeseries, +) + +router = APIRouter(prefix="/v1/analytics", tags=["analytics"]) + + +def _identity(request: Request) -> tuple[str, str]: + context = request.state.auth + if not context.workspace_id or not context.user_id: + from fastapi import HTTPException + + raise HTTPException(status_code=403, detail="No active workspace membership.") + return context.workspace_id, context.user_id + + +def _service(request: Request): + service = request.app.state.container.analytics + service.ensure_ready() + return service + + +@router.get("/capabilities", response_model=list[AnalyticsCapabilities]) +async def capabilities(request: Request) -> list[AnalyticsCapabilities]: + return await _service(request).capabilities() + + +@router.get("/overview", response_model=AnalyticsOverview) +async def overview( + request: Request, query: Annotated[AnalyticsQuery, Depends()] +) -> AnalyticsOverview: + workspace_id, _ = _identity(request) + return await _service(request).overview(workspace_id, query) + + +@router.get("/timeseries", response_model=AnalyticsTimeseries) +async def timeseries( + request: Request, query: Annotated[AnalyticsQuery, Depends()] +) -> AnalyticsTimeseries: + workspace_id, _ = _identity(request) + return await _service(request).timeseries(workspace_id, query) + + +@router.get("/platforms", response_model=AnalyticsOverview) +async def platforms( + request: Request, query: Annotated[AnalyticsQuery, Depends()] +) -> AnalyticsOverview: + workspace_id, _ = _identity(request) + return await _service(request).overview(workspace_id, query) + + +@router.get("/platforms/{provider}", response_model=AnalyticsOverview) +async def platform( + request: Request, + provider: str, + query: Annotated[AnalyticsQuery, Depends()], +) -> AnalyticsOverview: + workspace_id, _ = _identity(request) + query.provider = provider + return await _service(request).overview(workspace_id, query) + + +@router.get("/posts", response_model=AnalyticsPostList) +async def posts(request: Request, query: Annotated[AnalyticsQuery, Depends()]) -> AnalyticsPostList: + workspace_id, _ = _identity(request) + return await _service(request).posts(workspace_id, query) + + +@router.get("/posts/{post_id}", response_model=AnalyticsPostList) +async def post( + request: Request, + post_id: str, + query: Annotated[AnalyticsQuery, Depends()], +) -> AnalyticsPostList: + workspace_id, _ = _identity(request) + return await _service(request).post(workspace_id, post_id, query) + + +@router.get("/projects/{project_id}", response_model=AnalyticsOverview) +async def project( + request: Request, + project_id: str, + query: Annotated[AnalyticsQuery, Depends()], +) -> AnalyticsOverview: + workspace_id, user_id = _identity(request) + await request.app.state.container.projects.get( + workspace_id=workspace_id, user_id=user_id, project_id=project_id + ) + query.project_id = project_id + return await _service(request).overview(workspace_id, query) + + +@router.get("/sync-runs", response_model=AnalyticsSyncList) +async def sync_runs( + request: Request, + offset: int = Query(default=0, ge=0, le=100_000), + limit: int = Query(default=50, ge=1, le=500), +) -> AnalyticsSyncList: + workspace_id, _ = _identity(request) + items = await _service(request).repository.list_syncs(workspace_id, offset=offset, limit=limit) + return AnalyticsSyncList( + items=[_service(request)._sync_view(item) for item in items], + offset=offset, + limit=limit, + ) + + +@router.get("/sync-runs/{run_id}", response_model=AnalyticsSyncRunView) +async def sync_run(request: Request, run_id: str) -> AnalyticsSyncRunView: + workspace_id, _ = _identity(request) + return _service(request)._sync_view( + await _service(request).repository.get_sync(workspace_id, run_id) + ) + + +@router.post("/sync-runs/{run_id}/cancel", response_model=AnalyticsSyncRunView) +async def cancel_sync(request: Request, run_id: str) -> AnalyticsSyncRunView: + workspace_id, _ = _identity(request) + run = await _service(request).repository.cancel_sync(workspace_id, run_id) + await _service(request).audit.record( + workspace_id=workspace_id, + event_type="analytics.sync_cancelled", + metadata={"sync_run_id": run.id}, + api_key_id=request.state.auth.api_key_id, + request_id=request.state.request_id, + ) + return _service(request)._sync_view(run) + + +@router.post("/sync", response_model=AnalyticsSyncRunView, status_code=status.HTTP_202_ACCEPTED) +async def sync( + request: Request, + payload: AnalyticsSyncRequest, + idempotency_key: str | None = Header(default=None, alias="Idempotency-Key"), +) -> AnalyticsSyncRunView: + workspace_id, user_id = _identity(request) + if idempotency_key: + payload.idempotency_key = idempotency_key + return await _service(request).create_sync(workspace_id, user_id, payload) diff --git a/app/analytics/errors.py b/app/analytics/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..4bb04742da009aeba5e0b54738b9ad5e6fc3a1e9 --- /dev/null +++ b/app/analytics/errors.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from app.core.exceptions import MediaAPIError + + +class AnalyticsValidationError(MediaAPIError): + code = "ANALYTICS_VALIDATION_ERROR" + status_code = 422 + + +class AnalyticsNotFoundError(MediaAPIError): + code = "ANALYTICS_NOT_FOUND" + status_code = 404 + + +class AnalyticsIdempotencyConflictError(MediaAPIError): + code = "ANALYTICS_IDEMPOTENCY_CONFLICT" + status_code = 409 diff --git a/app/analytics/models.py b/app/analytics/models.py new file mode 100644 index 0000000000000000000000000000000000000000..d918b1ac1d244379cc8f6bacd11195b7e1f2109e --- /dev/null +++ b/app/analytics/models.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from uuid import uuid4 + +from sqlalchemy import ( + JSON, + DateTime, + Float, + ForeignKey, + Index, + Integer, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.social.models import SocialBase + + +def utcnow() -> datetime: + return datetime.now(timezone.utc) + + +def new_id() -> str: + return str(uuid4()) + + +class AnalyticsSyncRun(SocialBase): + __tablename__ = "analytics_sync_runs" + __table_args__ = ( + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_analytics_sync_workspace_idempotency" + ), + Index("ix_analytics_sync_workspace_status", "workspace_id", "status"), + Index("ix_analytics_sync_due", "status", "next_attempt_at"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + project_id: Mapped[str | None] = mapped_column(String(36)) + provider: Mapped[str | None] = mapped_column(String(32)) + date_from: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + date_to: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + timezone: Mapped[str] = mapped_column(String(100), nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + requested_by: Mapped[str | None] = mapped_column(String(120)) + attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + next_attempt_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + error_code: Mapped[str | None] = mapped_column(String(100)) + error_message: Mapped[str | None] = mapped_column(Text) + metrics_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class AnalyticsMetricSnapshot(SocialBase): + __tablename__ = "analytics_metric_snapshots" + __table_args__ = ( + UniqueConstraint( + "workspace_id", + "provider", + "external_object_id", + "metric_name", + "bucket_start", + name="uq_analytics_metric_snapshot", + ), + Index("ix_analytics_metric_snapshots_workspace_bucket", "workspace_id", "bucket_start"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + project_id: Mapped[str | None] = mapped_column(String(36)) + social_post_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE") + ) + social_account_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="CASCADE") + ) + provider: Mapped[str] = mapped_column(String(32), nullable=False) + external_object_id: Mapped[str] = mapped_column(String(255), nullable=False) + metric_name: Mapped[str] = mapped_column(String(64), nullable=False) + metric_value: Mapped[float] = mapped_column(Float, nullable=False) + bucket_start: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + dimensions: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + source: Mapped[str] = mapped_column(String(64), nullable=False, default="provider") + collected_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + provider_updated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class AnalyticsPostMetric(SocialBase): + __tablename__ = "analytics_post_metrics" + __table_args__ = ( + UniqueConstraint( + "workspace_id", + "social_post_target_id", + "metric_date", + name="uq_analytics_post_metric_bucket", + ), + Index("ix_analytics_post_metrics_workspace_date", "workspace_id", "metric_date"), + Index( + "ix_analytics_post_metrics_project_date", "workspace_id", "project_id", "metric_date" + ), + Index("ix_analytics_post_metrics_provider_date", "workspace_id", "provider", "metric_date"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + project_id: Mapped[str | None] = mapped_column(String(36)) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + social_post_target_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_post_targets.id", ondelete="CASCADE"), nullable=False + ) + social_account_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False + ) + provider: Mapped[str] = mapped_column(String(32), nullable=False) + external_post_id: Mapped[str | None] = mapped_column(String(255)) + metric_date: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + views: Mapped[int | None] = mapped_column() + impressions: Mapped[int | None] = mapped_column() + likes: Mapped[int | None] = mapped_column() + comments: Mapped[int | None] = mapped_column() + shares: Mapped[int | None] = mapped_column() + engagement_rate: Mapped[float | None] = mapped_column(Float) + dimensions: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + source: Mapped[str] = mapped_column(String(64), nullable=False, default="provider") + collected_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + provider_updated_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class AnalyticsPlatformMetric(SocialBase): + __tablename__ = "analytics_platform_metrics" + __table_args__ = ( + UniqueConstraint( + "workspace_id", + "provider", + "social_account_id", + "metric_date", + name="uq_analytics_platform_metric_bucket", + ), + Index("ix_analytics_platform_metrics_workspace_date", "workspace_id", "metric_date"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + project_id: Mapped[str | None] = mapped_column(String(36)) + social_account_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False + ) + provider: Mapped[str] = mapped_column(String(32), nullable=False) + metric_date: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + posts_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + views: Mapped[int | None] = mapped_column() + impressions: Mapped[int | None] = mapped_column() + likes: Mapped[int | None] = mapped_column() + comments: Mapped[int | None] = mapped_column() + shares: Mapped[int | None] = mapped_column() + engagement_rate: Mapped[float | None] = mapped_column(Float) + dimensions: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + source: Mapped[str] = mapped_column(String(64), nullable=False, default="provider") + collected_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) diff --git a/app/analytics/repository.py b/app/analytics/repository.py new file mode 100644 index 0000000000000000000000000000000000000000..551f3da5a9a7f002815028762587dcb09ee87f8d --- /dev/null +++ b/app/analytics/repository.py @@ -0,0 +1,385 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import or_, select +from sqlalchemy.exc import IntegrityError + +from app.analytics.errors import ( + AnalyticsIdempotencyConflictError, + AnalyticsNotFoundError, + AnalyticsValidationError, +) +from app.analytics.models import ( + AnalyticsMetricSnapshot, + AnalyticsPlatformMetric, + AnalyticsPostMetric, + AnalyticsSyncRun, +) +from app.social.database import SocialDatabase +from app.social.models import SocialPost, SocialPostTarget + + +class AnalyticsRepository: + def __init__(self, database: SocialDatabase) -> None: + self.database = database + + async def create_sync(self, run: AnalyticsSyncRun) -> AnalyticsSyncRun: + try: + async with self.database.session(run.workspace_id) as session: + session.add(run) + await session.commit() + await session.refresh(run) + return run + except IntegrityError as exc: + existing = await self.get_sync_by_idempotency(run.workspace_id, run.idempotency_key) + if existing is not None: + if ( + existing.date_from != run.date_from + or existing.date_to != run.date_to + or existing.provider != run.provider + or existing.project_id != run.project_id + ): + raise AnalyticsIdempotencyConflictError( + "The analytics idempotency key was already used for a different request." + ) from exc + return existing + raise + + async def get_sync_by_idempotency( + self, workspace_id: str, idempotency_key: str + ) -> AnalyticsSyncRun | None: + async with self.database.session(workspace_id) as session: + return await session.scalar( + select(AnalyticsSyncRun).where( + AnalyticsSyncRun.workspace_id == workspace_id, + AnalyticsSyncRun.idempotency_key == idempotency_key, + ) + ) + + async def get_sync(self, workspace_id: str, run_id: str) -> AnalyticsSyncRun: + async with self.database.session(workspace_id) as session: + run = await session.scalar( + select(AnalyticsSyncRun).where( + AnalyticsSyncRun.id == run_id, + AnalyticsSyncRun.workspace_id == workspace_id, + ) + ) + if run is None: + raise AnalyticsNotFoundError("Analytics sync run was not found.") + return run + + async def list_syncs( + self, workspace_id: str, *, offset: int, limit: int + ) -> list[AnalyticsSyncRun]: + async with self.database.session(workspace_id) as session: + return list( + ( + await session.scalars( + select(AnalyticsSyncRun) + .where(AnalyticsSyncRun.workspace_id == workspace_id) + .order_by(AnalyticsSyncRun.created_at.desc()) + .offset(offset) + .limit(limit) + ) + ).all() + ) + + async def claim_due_syncs(self, *, limit: int = 4) -> list[AnalyticsSyncRun]: + async with self.database.worker_session() as session: + now = datetime.now(timezone.utc) + statement = ( + select(AnalyticsSyncRun) + .where( + AnalyticsSyncRun.status.in_(["queued", "retrying"]), + (AnalyticsSyncRun.next_attempt_at.is_(None)) + | (AnalyticsSyncRun.next_attempt_at <= now), + ) + .order_by(AnalyticsSyncRun.created_at) + .limit(limit) + .with_for_update(skip_locked=True) + ) + runs = list((await session.scalars(statement)).all()) + for run in runs: + run.status = "running" + run.started_at = now + run.attempt_count += 1 + run.updated_at = now + await session.commit() + return runs + + async def update_sync( + self, + run_id: str, + *, + status: str, + metrics_count: int | None = None, + error_code: str | None = None, + error_message: str | None = None, + next_attempt_at: datetime | None = None, + ) -> AnalyticsSyncRun: + async with self.database.worker_session() as session: + run = await session.get(AnalyticsSyncRun, run_id) + if run is None: + raise AnalyticsNotFoundError("Analytics sync run was not found.") + run.status = status + run.error_code = error_code + run.error_message = error_message + run.next_attempt_at = next_attempt_at + if metrics_count is not None: + run.metrics_count = metrics_count + if status in {"succeeded", "partial", "failed", "cancelled"}: + run.completed_at = datetime.now(timezone.utc) + run.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(run) + return run + + async def cancel_sync(self, workspace_id: str, run_id: str) -> AnalyticsSyncRun: + async with self.database.session(workspace_id) as session: + run = await session.scalar( + select(AnalyticsSyncRun).where( + AnalyticsSyncRun.id == run_id, + AnalyticsSyncRun.workspace_id == workspace_id, + ) + ) + if run is None: + raise AnalyticsNotFoundError("Analytics sync run was not found.") + if run.status not in {"queued", "running"}: + raise AnalyticsValidationError( + "Only queued or running analytics synchronization can be cancelled." + ) + run.status = "cancelled" + run.completed_at = datetime.now(timezone.utc) + run.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(run) + return run + + async def published_targets( + self, + workspace_id: str, + *, + project_id: str | None = None, + provider: str | None = None, + date_from: datetime | None = None, + date_to: datetime | None = None, + ) -> list[tuple[SocialPost, SocialPostTarget]]: + async with self.database.session(workspace_id) as session: + statement = ( + select(SocialPost, SocialPostTarget) + .join(SocialPostTarget, SocialPostTarget.social_post_id == SocialPost.id) + .where( + SocialPost.workspace_id == workspace_id, + SocialPostTarget.external_post_id.is_not(None), + SocialPostTarget.status == "published", + ) + ) + if project_id: + statement = statement.where(SocialPost.project_id == project_id) + if provider: + statement = statement.where(SocialPostTarget.provider == provider) + if date_from: + statement = statement.where(SocialPostTarget.published_at >= date_from) + if date_to: + statement = statement.where(SocialPostTarget.published_at <= date_to) + return list((await session.execute(statement)).all()) + + async def upsert_post_metric( + self, + *, + workspace_id: str, + post: SocialPost, + target: SocialPostTarget, + metric: dict[str, object], + collected_at: datetime, + ) -> AnalyticsPostMetric: + bucket = target.published_at or collected_at + async with self.database.session(workspace_id) as session: + existing = await session.scalar( + select(AnalyticsPostMetric).where( + AnalyticsPostMetric.workspace_id == workspace_id, + AnalyticsPostMetric.social_post_target_id == target.id, + AnalyticsPostMetric.metric_date == bucket, + ) + ) + values = dict( + provider=target.provider, + views=metric.get("views"), + impressions=metric.get("impressions"), + likes=metric.get("likes"), + comments=metric.get("comments"), + shares=metric.get("shares"), + engagement_rate=metric.get("engagement_rate"), + dimensions=( + metric.get("raw_metrics") if isinstance(metric.get("raw_metrics"), dict) else {} + ), + source="provider", + collected_at=collected_at, + external_post_id=target.external_post_id, + ) + if existing is None: + existing = AnalyticsPostMetric( + workspace_id=workspace_id, + project_id=post.project_id, + social_post_id=post.id, + social_post_target_id=target.id, + social_account_id=target.social_account_id, + metric_date=bucket, + **values, + ) + session.add(existing) + else: + for key, value in values.items(): + setattr(existing, key, value) + if target.external_post_id: + for metric_name in ( + "views", + "impressions", + "likes", + "comments", + "shares", + "engagement_rate", + ): + metric_value = metric.get(metric_name) + if not isinstance(metric_value, (int, float)) or isinstance(metric_value, bool): + continue + snapshot = await session.scalar( + select(AnalyticsMetricSnapshot).where( + AnalyticsMetricSnapshot.workspace_id == workspace_id, + AnalyticsMetricSnapshot.provider == target.provider, + AnalyticsMetricSnapshot.external_object_id == target.external_post_id, + AnalyticsMetricSnapshot.metric_name == metric_name, + AnalyticsMetricSnapshot.bucket_start == bucket, + ) + ) + if snapshot is None: + session.add( + AnalyticsMetricSnapshot( + workspace_id=workspace_id, + project_id=post.project_id, + social_post_id=post.id, + social_account_id=target.social_account_id, + provider=target.provider, + external_object_id=target.external_post_id, + metric_name=metric_name, + metric_value=float(metric_value), + bucket_start=bucket, + dimensions={}, + source="provider", + collected_at=collected_at, + ) + ) + else: + snapshot.metric_value = float(metric_value) + snapshot.collected_at = collected_at + await session.flush() + day_start = bucket.astimezone(timezone.utc).replace( + hour=0, minute=0, second=0, microsecond=0 + ) + day_end = day_start.replace(hour=23, minute=59, second=59, microsecond=999999) + account_rows = list( + ( + await session.scalars( + select(AnalyticsPostMetric).where( + AnalyticsPostMetric.workspace_id == workspace_id, + AnalyticsPostMetric.social_account_id == target.social_account_id, + AnalyticsPostMetric.metric_date >= day_start, + AnalyticsPostMetric.metric_date <= day_end, + ) + ) + ).all() + ) + platform = await session.scalar( + select(AnalyticsPlatformMetric).where( + AnalyticsPlatformMetric.workspace_id == workspace_id, + AnalyticsPlatformMetric.provider == target.provider, + AnalyticsPlatformMetric.social_account_id == target.social_account_id, + AnalyticsPlatformMetric.metric_date == day_start, + ) + ) + aggregates: dict[str, int | float | None] = {} + for name in ( + "views", + "impressions", + "likes", + "comments", + "shares", + "engagement_rate", + ): + metric_values = [ + getattr(item, name) for item in account_rows if getattr(item, name) is not None + ] + aggregates[name] = ( + ( + sum(metric_values) / len(metric_values) + if name == "engagement_rate" + else sum(metric_values) + ) + if metric_values + else None + ) + if platform is None: + platform = AnalyticsPlatformMetric( + workspace_id=workspace_id, + project_id=post.project_id, + social_account_id=target.social_account_id, + provider=target.provider, + metric_date=day_start, + ) + session.add(platform) + platform.posts_count = len(account_rows) + platform.collected_at = collected_at + for name, value in aggregates.items(): + setattr(platform, name, value) + await session.commit() + await session.refresh(existing) + return existing + + async def metric_rows( + self, + workspace_id: str, + query: dict[str, object], + ) -> list[AnalyticsPostMetric]: + async with self.database.session(workspace_id) as session: + statement = select(AnalyticsPostMetric).where( + AnalyticsPostMetric.workspace_id == workspace_id, + AnalyticsPostMetric.metric_date >= query["date_from"], + AnalyticsPostMetric.metric_date <= query["date_to"], + ) + for field, model_field in ( + ("project_id", AnalyticsPostMetric.project_id), + ("provider", AnalyticsPostMetric.provider), + ("social_account_id", AnalyticsPostMetric.social_account_id), + ("post_id", AnalyticsPostMetric.social_post_id), + ): + value = query.get(field) + if value: + statement = statement.where(model_field == value) + if query.get("search"): + pattern = f"%{str(query['search']).strip()}%" + statement = statement.join( + SocialPost, SocialPost.id == AnalyticsPostMetric.social_post_id + ).where( + or_( + SocialPost.canonical_caption.ilike(pattern), + SocialPost.id.ilike(pattern), + ) + ) + return list((await session.scalars(statement)).all()) + + async def latest_sync( + self, workspace_id: str, *, provider: str | None = None, project_id: str | None = None + ) -> AnalyticsSyncRun | None: + async with self.database.session(workspace_id) as session: + statement = select(AnalyticsSyncRun).where( + AnalyticsSyncRun.workspace_id == workspace_id, + AnalyticsSyncRun.status.in_(["succeeded", "partial"]), + ) + if provider: + statement = statement.where(AnalyticsSyncRun.provider == provider) + if project_id: + statement = statement.where(AnalyticsSyncRun.project_id == project_id) + return await session.scalar( + statement.order_by(AnalyticsSyncRun.completed_at.desc()).limit(1) + ) diff --git a/app/analytics/schemas.py b/app/analytics/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..db61914656b90178a3020843a2f67e18b55250f2 --- /dev/null +++ b/app/analytics/schemas.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field + +MetricName = Literal[ + "views", + "impressions", + "likes", + "comments", + "shares", + "engagement_rate", + "reach", + "clicks", + "saves", + "watch_time", + "completion_rate", +] +Granularity = Literal["day", "week", "month"] + + +class AnalyticsQuery(BaseModel): + model_config = ConfigDict(extra="forbid") + + date_from: datetime | None = None + date_to: datetime | None = None + timezone: str = Field(default="UTC", min_length=1, max_length=100) + provider: str | None = Field(default=None, max_length=32) + project_id: str | None = Field(default=None, max_length=36) + social_account_id: str | None = Field(default=None, max_length=36) + post_id: str | None = Field(default=None, max_length=36) + metric: MetricName | None = None + granularity: Granularity = "day" + search: str | None = Field(default=None, max_length=200) + sort: Literal["date", "views", "engagement_rate", "likes", "comments", "shares"] = "date" + descending: bool = True + offset: int = Field(default=0, ge=0, le=100_000) + limit: int = Field(default=50, ge=1, le=500) + + +class AnalyticsMetric(BaseModel): + model_config = ConfigDict(extra="forbid") + + provider: str + post_id: str | None = None + project_id: str | None = None + metric_date: datetime + views: int | None = None + impressions: int | None = None + likes: int | None = None + comments: int | None = None + shares: int | None = None + engagement_rate: float | None = None + status: Literal["available", "unavailable", "unsupported", "not_synchronized"] = "available" + source: str = "provider" + collected_at: datetime | None = None + + +class AnalyticsFreshness(BaseModel): + status: Literal["fresh", "stale", "not_synchronized", "partial", "failed"] + last_collected_at: datetime | None = None + last_sync_id: str | None = None + reason: str | None = None + + +class AnalyticsOverview(BaseModel): + model_config = ConfigDict(extra="forbid") + + date_from: datetime + date_to: datetime + timezone: str + timezone_source: Literal["requested", "utc_fallback"] = "requested" + totals: dict[str, int | float | None] + platforms: list[dict[str, Any]] + top_posts: list[AnalyticsMetric] + freshness: AnalyticsFreshness + + +class AnalyticsTimeseries(BaseModel): + date_from: datetime + date_to: datetime + timezone: str + granularity: Granularity + metric: str + points: list[dict[str, Any]] + freshness: AnalyticsFreshness + + +class AnalyticsSyncRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + date_from: datetime | None = None + date_to: datetime | None = None + timezone: str = Field(default="UTC", min_length=1, max_length=100) + provider: str | None = Field(default=None, max_length=32) + project_id: str | None = Field(default=None, max_length=36) + idempotency_key: str | None = Field(default=None, min_length=8, max_length=255) + + +class AnalyticsSyncRunView(BaseModel): + model_config = ConfigDict(from_attributes=True, extra="forbid") + + id: str + status: Literal["queued", "running", "succeeded", "partial", "failed", "cancelled"] + provider: str | None = None + project_id: str | None = None + date_from: datetime + date_to: datetime + timezone: str + attempt_count: int + metrics_count: int + error_code: str | None = None + error_message: str | None = None + created_at: datetime + started_at: datetime | None = None + completed_at: datetime | None = None + updated_at: datetime + + +class AnalyticsCapabilities(BaseModel): + provider: str + implementation_status: str + status: Literal["available", "unsupported", "unavailable"] + metrics: list[str] + required_scopes: list[str] = Field(default_factory=list) + + +class AnalyticsSyncList(BaseModel): + items: list[AnalyticsSyncRunView] + offset: int + limit: int + + +class AnalyticsPostList(BaseModel): + items: list[AnalyticsMetric] + offset: int + limit: int + freshness: AnalyticsFreshness diff --git a/app/analytics/service.py b/app/analytics/service.py new file mode 100644 index 0000000000000000000000000000000000000000..26871cc7ccc94de46fe1658385eb7bc06fae284a --- /dev/null +++ b/app/analytics/service.py @@ -0,0 +1,379 @@ +from __future__ import annotations + +from collections import defaultdict +from datetime import datetime, timedelta, timezone +from typing import Any +from zoneinfo import ZoneInfo, ZoneInfoNotFoundError + +from app.analytics.errors import AnalyticsValidationError +from app.analytics.models import AnalyticsSyncRun +from app.analytics.repository import AnalyticsRepository +from app.analytics.schemas import ( + AnalyticsCapabilities, + AnalyticsFreshness, + AnalyticsMetric, + AnalyticsOverview, + AnalyticsPostList, + AnalyticsQuery, + AnalyticsSyncRequest, + AnalyticsSyncRunView, + AnalyticsTimeseries, +) +from app.core.logger import get_logger +from app.social.providers.registry import ProviderRegistry +from app.social.services.analytics_service import AnalyticsService as SocialAnalyticsService +from app.social.services.audit_service import SocialAuditService + +logger = get_logger(__name__) +COMMON_METRICS = ("views", "impressions", "likes", "comments", "shares", "engagement_rate") +PROVIDER_METRICS: dict[str, tuple[str, ...]] = { + "facebook": ("views", "impressions", "likes", "comments", "shares"), + "instagram": ("views", "likes", "comments", "shares"), + "tiktok": ("views", "likes", "comments", "shares"), + "x": ("impressions", "likes", "comments", "shares"), + "youtube": ("views", "likes", "comments"), + "linkedin": ("impressions", "likes", "comments", "shares", "engagement_rate"), +} + + +class AnalyticsDomainService: + """Tenant-scoped analytics aggregation over provider-reported snapshots.""" + + def __init__( + self, + repository: AnalyticsRepository, + providers: ProviderRegistry, + social_analytics: SocialAnalyticsService, + accounts: object, + audit: SocialAuditService, + projects: object | None = None, + ) -> None: + self.repository = repository + self.providers = providers + self.social_analytics = social_analytics + self.accounts = accounts + self.audit = audit + self.projects = projects + self.ready = False + + async def initialize(self, social_ready: bool) -> None: + self.ready = social_ready + + def ensure_ready(self) -> None: + if not self.ready: + from app.social.domain.errors import SocialProviderUnavailableError + + raise SocialProviderUnavailableError( + "Analytics database schema is unavailable. Apply the analytics migration." + ) + + @staticmethod + def _range(query: AnalyticsQuery) -> tuple[datetime, datetime, str]: + end = query.date_to or datetime.now(timezone.utc) + start = query.date_from or end - timedelta(days=30) + if start.tzinfo is None or end.tzinfo is None: + raise AnalyticsValidationError("Analytics dates must include a UTC offset.") + if end <= start: + raise AnalyticsValidationError("date_to must be after date_from.") + if end - start > timedelta(days=366): + raise AnalyticsValidationError("Analytics date ranges are limited to 366 days.") + try: + ZoneInfo(query.timezone) + except ZoneInfoNotFoundError as exc: + raise AnalyticsValidationError("timezone must be a valid IANA timezone.") from exc + return start.astimezone(timezone.utc), end.astimezone(timezone.utc), query.timezone + + @staticmethod + def _metric_view(row: Any) -> AnalyticsMetric: + return AnalyticsMetric( + provider=row.provider, + post_id=row.social_post_id, + project_id=row.project_id, + metric_date=row.metric_date, + views=row.views, + impressions=row.impressions, + likes=row.likes, + comments=row.comments, + shares=row.shares, + engagement_rate=row.engagement_rate, + collected_at=row.collected_at, + source=row.source, + ) + + async def capabilities(self) -> list[AnalyticsCapabilities]: + result: list[AnalyticsCapabilities] = [] + for adapter in self.providers.list(): + capabilities = adapter.capabilities + available = bool(capabilities.analytics) + result.append( + AnalyticsCapabilities( + provider=capabilities.provider.value, + implementation_status=capabilities.implementation_status, + status="available" if available else "unsupported", + metrics=( + list(PROVIDER_METRICS.get(capabilities.provider.value, ())) + if available + else [] + ), + required_scopes=list(capabilities.analytics_required_scopes), + ) + ) + return result + + async def create_sync( + self, workspace_id: str, user_id: str, payload: AnalyticsSyncRequest + ) -> AnalyticsSyncRunView: + query = AnalyticsQuery( + date_from=payload.date_from, + date_to=payload.date_to, + timezone=payload.timezone, + provider=payload.provider, + project_id=payload.project_id, + ) + start, end, tz = self._range(query) + if not payload.idempotency_key: + raise AnalyticsValidationError( + "Idempotency-Key is required for analytics synchronization." + ) + if payload.provider: + self.providers.get(payload.provider) + if payload.project_id and self.projects is not None: + await self.projects.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=payload.project_id, + ) + run = await self.repository.create_sync( + AnalyticsSyncRun( + workspace_id=workspace_id, + project_id=payload.project_id, + provider=payload.provider, + date_from=start, + date_to=end, + timezone=tz, + idempotency_key=payload.idempotency_key, + requested_by=user_id, + ) + ) + await self.audit.record( + workspace_id=workspace_id, + event_type="analytics.sync_requested", + metadata={"sync_run_id": run.id, "provider": payload.provider}, + ) + return self._sync_view(run) + + async def sync_once(self, run: AnalyticsSyncRun) -> AnalyticsSyncRunView: + targets = await self.repository.published_targets( + run.workspace_id, + project_id=run.project_id, + provider=run.provider, + date_from=run.date_from, + date_to=run.date_to, + ) + await self.audit.record( + workspace_id=run.workspace_id, + event_type="analytics.sync_started", + metadata={"sync_run_id": run.id}, + ) + errors = 0 + count = 0 + by_post: dict[str, tuple[Any, list[Any]]] = {} + for post, target in targets: + by_post.setdefault(post.id, (post, []))[1].append(target) + for post, post_targets in by_post.values(): + try: + latest = await self.repository.get_sync(run.workspace_id, run.id) + if latest.status == "cancelled": + return self._sync_view(latest) + result = await self.social_analytics.post(run.workspace_id, post.id) + for target in post_targets: + matches = [ + item + for item in result.get("metrics", []) + if isinstance(item, dict) and item.get("provider") == target.provider + ] + if not matches: + errors += 1 + continue + metric = matches[-1] + await self.repository.upsert_post_metric( + workspace_id=run.workspace_id, + post=post, + target=target, + metric=metric, + collected_at=datetime.now(timezone.utc), + ) + count += 1 + if result.get("unavailable"): + errors += 1 + except Exception as exc: + errors += 1 + logger.warning( + "analytics target synchronization failed", + extra={"sync_run_id": run.id, "post_id": post.id}, + ) + if not getattr(exc, "code", None): + logger.debug("analytics sync exception", exc_info=True) + status = ( + "failed" if targets and count == 0 and errors else "partial" if errors else "succeeded" + ) + if not targets: + status = "succeeded" + updated = await self.repository.update_sync(run.id, status=status, metrics_count=count) + await self.audit.record( + workspace_id=run.workspace_id, + event_type=( + "analytics.sync_completed" if status != "failed" else "analytics.sync_failed" + ), + metadata={"sync_run_id": run.id, "status": status, "metrics_count": count}, + ) + return self._sync_view(updated) + + async def overview(self, workspace_id: str, query: AnalyticsQuery) -> AnalyticsOverview: + start, end, tz = self._range(query) + rows = await self.repository.metric_rows( + workspace_id, + { + "date_from": start, + "date_to": end, + "provider": query.provider, + "project_id": query.project_id, + "social_account_id": query.social_account_id, + "post_id": query.post_id, + "search": query.search, + }, + ) + totals = self._totals(rows) + grouped: dict[str, list[Any]] = defaultdict(list) + for row in rows: + grouped[row.provider].append(row) + platforms = [ + {"provider": provider, "posts": len(items), **self._totals(items)} + for provider, items in sorted(grouped.items()) + ] + top = sorted( + rows, key=lambda row: self._sort_value(row, query.sort), reverse=query.descending + )[:10] + return AnalyticsOverview( + date_from=start, + date_to=end, + timezone=tz, + timezone_source="requested" if query.timezone != "UTC" else "utc_fallback", + totals=totals, + platforms=platforms, + top_posts=[self._metric_view(row) for row in top], + freshness=await self._freshness(workspace_id, query), + ) + + async def timeseries(self, workspace_id: str, query: AnalyticsQuery) -> AnalyticsTimeseries: + start, end, tz_name = self._range(query) + rows = await self.repository.metric_rows( + workspace_id, + { + "date_from": start, + "date_to": end, + "provider": query.provider, + "project_id": query.project_id, + "social_account_id": query.social_account_id, + "post_id": query.post_id, + "search": query.search, + }, + ) + zone = ZoneInfo(tz_name) + points: dict[str, dict[str, Any]] = {} + for row in rows: + local = row.metric_date.astimezone(zone) + if query.granularity == "month": + key = local.strftime("%Y-%m-01") + elif query.granularity == "week": + monday = local.date() - timedelta(days=local.weekday()) + key = monday.isoformat() + else: + key = local.date().isoformat() + item = points.setdefault(key, {"bucket": key, "value": 0, "posts": 0}) + value = getattr(row, query.metric or "views") + if value is not None: + item["value"] += value + item["posts"] += 1 + return AnalyticsTimeseries( + date_from=start, + date_to=end, + timezone=tz_name, + granularity=query.granularity, + metric=query.metric or "views", + points=[points[key] for key in sorted(points)], + freshness=await self._freshness(workspace_id, query), + ) + + async def posts(self, workspace_id: str, query: AnalyticsQuery) -> AnalyticsPostList: + start, end, _ = self._range(query) + rows = await self.repository.metric_rows( + workspace_id, + { + "date_from": start, + "date_to": end, + "provider": query.provider, + "project_id": query.project_id, + "social_account_id": query.social_account_id, + "post_id": query.post_id, + "search": query.search, + }, + ) + rows.sort(key=lambda row: self._sort_value(row, query.sort), reverse=query.descending) + return AnalyticsPostList( + items=[ + self._metric_view(row) for row in rows[query.offset : query.offset + query.limit] + ], + offset=query.offset, + limit=query.limit, + freshness=await self._freshness(workspace_id, query), + ) + + async def post( + self, workspace_id: str, post_id: str, query: AnalyticsQuery + ) -> AnalyticsPostList: + query.post_id = post_id + return await self.posts(workspace_id, query) + + async def _freshness(self, workspace_id: str, query: AnalyticsQuery) -> AnalyticsFreshness: + run = await self.repository.latest_sync( + workspace_id, provider=query.provider, project_id=query.project_id + ) + if run is None: + return AnalyticsFreshness( + status="not_synchronized", reason="No analytics sync has completed." + ) + status = ( + "fresh" + if run.completed_at + and run.completed_at >= datetime.now(timezone.utc) - timedelta(hours=24) + else "stale" + ) + if run.status == "partial": + status = "partial" + return AnalyticsFreshness( + status=status, last_collected_at=run.completed_at, last_sync_id=run.id + ) + + @staticmethod + def _totals(rows: list[Any]) -> dict[str, int | float | None]: + result: dict[str, int | float | None] = {} + for name in COMMON_METRICS: + values = [getattr(row, name) for row in rows if getattr(row, name) is not None] + result[name] = ( + (sum(values) / len(values) if name == "engagement_rate" else sum(values)) + if values + else None + ) + return result + + @staticmethod + def _sort_value(row: Any, field: str) -> float: + if field == "date": + return row.metric_date.timestamp() + value = getattr(row, field, None) + return float(value) if isinstance(value, (int, float)) else 0.0 + + @staticmethod + def _sync_view(run: AnalyticsSyncRun) -> AnalyticsSyncRunView: + return AnalyticsSyncRunView.model_validate(run, from_attributes=True) diff --git a/app/analytics/workers/__init__.py b/app/analytics/workers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5e5f4e70958cfc11c06eb708262120b8091bdaf5 --- /dev/null +++ b/app/analytics/workers/__init__.py @@ -0,0 +1 @@ +"""Durable analytics synchronization workers.""" diff --git a/app/analytics/workers/sync.py b/app/analytics/workers/sync.py new file mode 100644 index 0000000000000000000000000000000000000000..8e44510af92d4ad6541e064eb80b583ad13d3832 --- /dev/null +++ b/app/analytics/workers/sync.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +import asyncio +import random +from datetime import datetime, timedelta, timezone + +from app.analytics.service import AnalyticsDomainService +from app.core.logger import get_logger +from app.social.database import SocialDatabase + +logger = get_logger(__name__) + + +class AnalyticsSyncWorker: + """Claims durable sync runs and executes bounded provider synchronization.""" + + def __init__( + self, + analytics: AnalyticsDomainService, + database: SocialDatabase, + *, + interval_seconds: int, + concurrency: int = 2, + ) -> None: + self.analytics = analytics + self.database = database + self.interval_seconds = max(1, interval_seconds) + self.concurrency = max(1, min(concurrency, 8)) + self._task: asyncio.Task[None] | None = None + self._stop = asyncio.Event() + + async def start(self) -> None: + if self._task is not None or not self.analytics.ready: + return + self._stop.clear() + self._task = asyncio.create_task(self._run(), name="analytics-sync-worker") + + async def stop(self) -> None: + self._stop.set() + if self._task is not None: + self._task.cancel() + await asyncio.gather(self._task, return_exceptions=True) + self._task = None + + async def run_once(self) -> None: + if not self.analytics.ready: + return + async with self.database.worker_boundary(): + runs = await self.analytics.repository.claim_due_syncs(limit=self.concurrency) + await asyncio.gather(*(self._execute(run) for run in runs)) + + async def _execute(self, run) -> None: + try: + latest = await self.analytics.repository.get_sync(run.workspace_id, run.id) + if latest.status == "cancelled": + return + await self.analytics.sync_once(run) + except asyncio.CancelledError: + raise + except Exception as exc: + retryable = run.attempt_count < 3 + if retryable: + delay = min(900, 30 * (2 ** max(0, run.attempt_count - 1))) + random.randint(0, 10) + await self.analytics.repository.update_sync( + run.id, + status="queued", + error_code=getattr(exc, "code", "ANALYTICS_SYNC_RETRY"), + error_message="Analytics synchronization will be retried.", + next_attempt_at=datetime.now(timezone.utc) + timedelta(seconds=delay), + ) + else: + await self.analytics.repository.update_sync( + run.id, + status="failed", + error_code=getattr(exc, "code", "ANALYTICS_SYNC_FAILED"), + error_message="Analytics synchronization failed.", + ) + await self.analytics.audit.record( + workspace_id=run.workspace_id, + event_type="analytics.sync_failed", + metadata={"sync_run_id": run.id, "error_code": getattr(exc, "code", None)}, + ) + logger.warning( + "analytics sync execution failed", + extra={"sync_run_id": run.id, "attempt": run.attempt_count}, + ) + + async def _run(self) -> None: + while not self._stop.is_set(): + try: + await self.run_once() + except asyncio.CancelledError: + raise + except Exception: + logger.exception("analytics sync worker iteration failed") + try: + await asyncio.wait_for(self._stop.wait(), timeout=self.interval_seconds) + except TimeoutError: + pass diff --git a/app/api/api_keys.py b/app/api/api_keys.py index 23596afd51ac30de1798fd11191ffb367405ef93..f989528cf003733aa1609ce0dbe118bce45dbcf1 100644 --- a/app/api/api_keys.py +++ b/app/api/api_keys.py @@ -38,6 +38,9 @@ async def current_auth_context(request: Request) -> AuthContextView: role=context.role, scopes=sorted(context.scopes), expires_at=context.expires_at, + workspace_id=context.workspace_id, + user_id=context.user_id, + membership_role=context.membership_role, ) @@ -80,7 +83,10 @@ async def create_api_key(request: Request, payload: APIKeyCreate) -> APIKeyCreat context = request.state.auth try: record, secret = await request.app.state.container.api_keys.create( - payload, created_by=context.api_key_id + payload, + created_by=context.api_key_id, + workspace_id=context.workspace_id, + user_id=context.user_id, ) except APIKeyConflictError as exc: raise HTTPException(status_code=409, detail=str(exc)) from exc diff --git a/app/api/generation.py b/app/api/generation.py new file mode 100644 index 0000000000000000000000000000000000000000..aaf7f3b75cde0d25a12854c9fa89dc8eb4d7aabd --- /dev/null +++ b/app/api/generation.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Header, HTTPException, Query, Request, status + +from app.generation.schemas.requests import ( + GenerationJobView, + GenerationProviderView, + GenerationRequestCreate, + GenerationRequestView, +) +from app.generation.model_registry import GenerationModelView + +router = APIRouter(prefix="/v1/generation", tags=["generation"]) + + +def _generation(request: Request): + service = request.app.state.container.generation + service.ensure_ready() + return service + + +def _identity(request: Request) -> tuple[str, str]: + context = request.state.auth + if not context.workspace_id or not context.user_id: + # The credential itself is never a workspace. API-key middleware + # resolves the authoritative membership before this route runs. + raise HTTPException(status_code=403, detail="No active workspace membership.") + return context.workspace_id, context.user_id + + +@router.get("/providers", response_model=list[GenerationProviderView]) +async def list_providers(request: Request) -> list[GenerationProviderView]: + return _generation(request).list_providers() + + +@router.get( + "/providers/{provider}", response_model=GenerationProviderView +) +async def get_provider(request: Request, provider: str) -> GenerationProviderView: + return _generation(request).get_provider(provider) + + +@router.get("/models", response_model=list[GenerationModelView]) +async def list_models( + request: Request, + provider: str | None = Query(default=None, min_length=1, max_length=64), +) -> list[GenerationModelView]: + return _generation(request).list_models(provider=provider) + + +@router.get("/providers/{provider}/models/{model_id}", response_model=GenerationModelView) +async def get_model( + request: Request, provider: str, model_id: str +) -> GenerationModelView: + return _generation(request).get_model(provider, model_id) + + +@router.post( + "/requests", + response_model=GenerationRequestView, + status_code=status.HTTP_202_ACCEPTED, +) +async def create_request( + request: Request, + payload: GenerationRequestCreate, + idempotency_key: Annotated[ + str, Header(alias="Idempotency-Key", min_length=8, max_length=255) + ], +) -> GenerationRequestView: + workspace_id, user_id = _identity(request) + return await _generation(request).create( + workspace_id=workspace_id, + user_id=user_id, + payload=payload, + idempotency_key=idempotency_key, + ) + + +@router.get("/requests", response_model=list[GenerationRequestView]) +async def list_requests( + request: Request, + offset: int = Query(default=0, ge=0), + limit: int = Query(default=100, ge=1, le=500), +) -> list[GenerationRequestView]: + workspace_id, user_id = _identity(request) + return await _generation(request).list_requests( + workspace_id, user_id, offset=offset, limit=limit + ) + + +@router.get("/requests/{generation_request_id}", response_model=GenerationRequestView) +async def get_request( + request: Request, generation_request_id: str +) -> GenerationRequestView: + workspace_id, user_id = _identity(request) + return await _generation(request).get_request( + workspace_id, user_id, generation_request_id + ) + + +@router.get("/jobs/{generation_job_id}", response_model=GenerationJobView) +async def get_job(request: Request, generation_job_id: str) -> GenerationJobView: + workspace_id, user_id = _identity(request) + return await _generation(request).get_job( + workspace_id, user_id, generation_job_id + ) + + +@router.post( + "/jobs/{generation_job_id}/cancel", response_model=GenerationJobView +) +async def cancel_job(request: Request, generation_job_id: str) -> GenerationJobView: + workspace_id, user_id = _identity(request) + return await _generation(request).cancel( + workspace_id, user_id, generation_job_id + ) diff --git a/app/api/media.py b/app/api/media.py index 3a175238700d02d642b0c317a8db04a755cc0595..9988fb1cd60936a3128447938950333f70a99d29 100644 --- a/app/api/media.py +++ b/app/api/media.py @@ -7,7 +7,9 @@ import aiofiles from fastapi import APIRouter, Request from fastapi.responses import StreamingResponse +from app.core.exceptions import NotFoundError from app.core.response import SuccessResponse +from app.security.assets import CanonicalAssetNotFoundError from app.services.media_service import MediaProcessor, Operation router = APIRouter(tags=["media"]) @@ -34,7 +36,25 @@ async def stream_file(path: Path) -> AsyncIterator[bytes]: async def download_media(request: Request, request_id: str, filename: str) -> StreamingResponse: request.state.operation = "media.download" container = request.app.state.container - path = container.cleanup.resolve_download(request_id, filename) + # The filesystem locator is not an authorization token. In authenticated + # deployments an output must have been issued by the pipeline to the + # caller's authoritative workspace before it can be downloaded. + if container.settings.auth_enabled: + context = getattr(request.state, "auth", None) + if context is None or not context.workspace_id: + raise NotFoundError("Output file not found") + try: + asset = await container.assets.get_owned( + workspace_id=context.workspace_id, + request_id=request_id, + filename=filename, + ) + path = container.cleanup.resolve_download(request_id, asset.filename) + await container.assets.verify_file(asset, path) + except CanonicalAssetNotFoundError as exc: + raise NotFoundError("Output file not found") from exc + else: + path = container.cleanup.resolve_download(request_id, filename) media_type = container.validator.infer_mime(path) headers = { "Content-Disposition": f'attachment; filename="{path.name}"', diff --git a/app/api/social.py b/app/api/social.py index d2d05e2e3f487f6d28192cf9ca22aa0551c5bc14..5e16b345a08ee37ff208da0f2cb04849b01c8685 100644 --- a/app/api/social.py +++ b/app/api/social.py @@ -1,24 +1,42 @@ from __future__ import annotations +from datetime import datetime from typing import Annotated -from fastapi import APIRouter, Header, Query, Request, Response, status +from fastapi import APIRouter, Header, HTTPException, Query, Request, Response, status from app.social.schemas.accounts import ( SocialAccountConnectRequest, SocialAccountSelectionRequest, SocialAccountView, SocialConnectResponse, - SocialPublishOptionsView, SocialProviderView, + SocialPublishOptionsView, ) from app.social.schemas.assets import ( SocialMediaAssetRegister, SocialMediaAssetView, ) from app.social.schemas.jobs import SocialJobView -from app.social.schemas.posts import SocialPostCreate, SocialPostView -from app.social.schemas.scheduling import SocialScheduleCreate, SocialScheduleView +from app.social.schemas.operations import ( + PublishingBatchView, + PublishingBulkRequest, + PublishingCalendarView, + PublishingContextView, + PublishingQueueView, +) +from app.social.schemas.posts import ( + SocialPostCreate, + SocialPostDuplicateRequest, + SocialPostPatch, + SocialPostValidation, + SocialPostView, +) +from app.social.schemas.scheduling import ( + SocialRescheduleRequest, + SocialScheduleCreate, + SocialScheduleView, +) router = APIRouter(prefix="/v1/social", tags=["social automation"]) @@ -31,9 +49,11 @@ def _social(request: Request): def _identity(request: Request) -> tuple[str, str]: context = request.state.auth - # Phase 1 isolation: the authenticated API-key identity is the tenant - # boundary until the backend gains a native workspace membership model. - return context.api_key_id, context.api_key_id + if not context.workspace_id or not context.user_id: + # Tenant resolution is performed by APIKeyService after credential + # verification. Never fall back to treating an API-key ID as a tenant. + raise HTTPException(status_code=403, detail="No active workspace membership.") + return context.workspace_id, context.user_id async def _audit( @@ -45,11 +65,11 @@ async def _audit( post_id: str | None = None, job_id: str | None = None, ) -> None: - workspace_id, user_id = _identity(request) + workspace_id, _ = _identity(request) await request.app.state.container.social.audit.record( workspace_id=workspace_id, event_type=event_type, - api_key_id=user_id, + api_key_id=request.state.auth.api_key_id, request_id=request.state.request_id, provider=provider, social_account_id=account_id, @@ -63,9 +83,7 @@ async def list_providers(request: Request) -> list[SocialProviderView]: return request.app.state.container.social.accounts.list_providers() -@router.get( - "/providers/{provider}/capabilities", response_model=SocialProviderView -) +@router.get("/providers/{provider}/capabilities", response_model=SocialProviderView) async def provider_capabilities(request: Request, provider: str) -> SocialProviderView: return request.app.state.container.social.accounts.get_provider(provider) @@ -95,9 +113,7 @@ async def list_accounts( limit: int = Query(default=100, ge=1, le=500), ) -> list[SocialAccountView]: workspace_id, _ = _identity(request) - return await _social(request).accounts.list( - workspace_id, offset=offset, limit=limit - ) + return await _social(request).accounts.list(workspace_id, offset=offset, limit=limit) @router.post("/accounts/select", response_model=list[SocialAccountView]) @@ -105,9 +121,7 @@ async def select_discovered_accounts( request: Request, payload: SocialAccountSelectionRequest ) -> list[SocialAccountView]: workspace_id, _ = _identity(request) - selected = await _social(request).accounts.select_discovered( - workspace_id, payload.account_ids - ) + selected = await _social(request).accounts.select_discovered(workspace_id, payload.account_ids) for account in selected: await _audit( request, @@ -132,14 +146,10 @@ async def get_account_publish_options( request: Request, account_id: str ) -> SocialPublishOptionsView: workspace_id, _ = _identity(request) - return await _social(request).publishing.publish_options( - workspace_id, account_id - ) + return await _social(request).publishing.publish_options(workspace_id, account_id) -@router.post( - "/accounts/{provider}/connect", response_model=SocialConnectResponse -) +@router.post("/accounts/{provider}/connect", response_model=SocialConnectResponse) async def connect_account( request: Request, provider: str, payload: SocialAccountConnectRequest ) -> SocialConnectResponse: @@ -162,9 +172,7 @@ async def connect_account( async def oauth_callback( request: Request, provider: str, - state: Annotated[ - str, Query(min_length=32, max_length=255, pattern=r"^[A-Za-z0-9_-]+$") - ], + state: Annotated[str, Query(min_length=32, max_length=255, pattern=r"^[A-Za-z0-9_-]+$")], code: Annotated[str | None, Query(min_length=1, max_length=4096)] = None, error: Annotated[str | None, Query(max_length=128)] = None, ) -> SocialAccountView: @@ -181,9 +189,7 @@ async def oauth_callback( @router.post("/accounts/{account_id}/refresh", response_model=SocialAccountView) async def refresh_account(request: Request, account_id: str) -> SocialAccountView: workspace_id, _ = _identity(request) - result = await _social(request).oauth.refresh( - workspace_id=workspace_id, account_id=account_id - ) + result = await _social(request).oauth.refresh(workspace_id=workspace_id, account_id=account_id) await _audit(request, "SOCIAL_ACCOUNT_REAUTHORIZED", account_id=account_id) return result @@ -202,9 +208,7 @@ async def disconnect_account(request: Request, account_id: str) -> Response: return Response(status_code=status.HTTP_204_NO_CONTENT) -@router.post( - "/posts", response_model=SocialPostView, status_code=status.HTTP_201_CREATED -) +@router.post("/posts", response_model=SocialPostView, status_code=status.HTTP_201_CREATED) async def create_post( request: Request, payload: SocialPostCreate, @@ -220,6 +224,8 @@ async def create_post( idempotency_key=idempotency_key, ) await _audit(request, "SOCIAL_POST_CREATED", post_id=result.id) + if result.publish_mode.value == "draft": + await _audit(request, "publishing.draft_created", post_id=result.id) return result @@ -228,10 +234,18 @@ async def list_posts( request: Request, offset: int = Query(default=0, ge=0), limit: int = Query(default=100, ge=1, le=500), + post_status: str | None = Query(default=None, alias="status", max_length=32), + project_id: str | None = Query(default=None, max_length=36), + search: str | None = Query(default=None, max_length=200), ) -> list[SocialPostView]: workspace_id, _ = _identity(request) return await _social(request).publishing.list( - workspace_id, offset=offset, limit=limit + workspace_id, + offset=offset, + limit=limit, + status=post_status, + project_id=project_id, + search=search, ) @@ -241,6 +255,76 @@ async def get_post(request: Request, post_id: str) -> SocialPostView: return await _social(request).publishing.get(workspace_id, post_id) +@router.patch("/posts/{post_id}", response_model=SocialPostView) +async def update_post(request: Request, post_id: str, payload: SocialPostPatch) -> SocialPostView: + workspace_id, _ = _identity(request) + result = await _social(request).operations.update_draft(workspace_id, post_id, payload) + await _audit(request, "publishing.draft_updated", post_id=post_id) + return result + + +@router.get("/drafts", response_model=list[SocialPostView]) +async def list_drafts( + request: Request, + offset: int = Query(default=0, ge=0), + limit: int = Query(default=100, ge=1, le=500), + search: str | None = Query(default=None, max_length=200), +) -> list[SocialPostView]: + workspace_id, _ = _identity(request) + posts = await _social(request).publishing.list( + workspace_id, + offset=offset, + limit=limit, + search=search, + ) + return [ + post + for post in posts + if post.status.value in {"draft", "ready", "failed"} and post.publish_mode.value == "draft" + ] + + +@router.patch("/drafts/{post_id}", response_model=SocialPostView) +async def update_draft(request: Request, post_id: str, payload: SocialPostPatch) -> SocialPostView: + return await update_post(request, post_id, payload) + + +@router.delete("/drafts/{post_id}", status_code=status.HTTP_204_NO_CONTENT) +async def delete_draft(request: Request, post_id: str) -> Response: + workspace_id, _ = _identity(request) + await _social(request).operations.delete_draft(workspace_id, post_id) + await _audit(request, "publishing.draft_deleted", post_id=post_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + + +@router.post("/posts/{post_id}/duplicate", response_model=SocialPostView) +async def duplicate_post( + request: Request, + post_id: str, + payload: SocialPostDuplicateRequest, + idempotency_key: str = Header(alias="Idempotency-Key", min_length=8, max_length=255), +) -> SocialPostView: + workspace_id, user_id = _identity(request) + result = await _social(request).operations.duplicate( + workspace_id, + user_id, + post_id, + payload, + idempotency_key=idempotency_key, + ) + await _audit(request, "publishing.duplicated", post_id=result.id) + return result + + +@router.post("/posts/{post_id}/validate", response_model=SocialPostValidation) +async def validate_post(request: Request, post_id: str) -> SocialPostValidation: + workspace_id, _ = _identity(request) + result = await _social(request).publishing.validate_post_targets(workspace_id, post_id) + await _audit(request, "SOCIAL_POST_VALIDATED", post_id=post_id) + await _audit(request, "publishing.validated", post_id=post_id) + return result + + @router.delete("/posts/{post_id}", status_code=status.HTTP_204_NO_CONTENT) async def delete_post(request: Request, post_id: str) -> Response: workspace_id, _ = _identity(request) @@ -263,6 +347,7 @@ async def publish_post( jobs = await _social(request).publishing.queue( workspace_id, post_id, idempotency_key=idempotency_key ) + await _audit(request, "SOCIAL_POST_PUBLISH_STARTED", post_id=post_id) return jobs @@ -271,19 +356,134 @@ async def schedule_post( request: Request, post_id: str, payload: SocialScheduleCreate ) -> SocialScheduleView: workspace_id, _ = _identity(request) - await _social(request).publishing.validate_post_targets(workspace_id, post_id) - schedule = await _social(request).scheduling.schedule( - workspace_id, post_id, payload - ) + validation = await _social(request).publishing.validate_post_targets(workspace_id, post_id) + if not validation.valid: + raise HTTPException( + status_code=422, + detail=validation.model_dump(mode="json"), + ) + schedule = await _social(request).scheduling.schedule(workspace_id, post_id, payload) await _audit(request, "SOCIAL_SCHEDULE_CREATED", post_id=post_id) + await _audit(request, "SOCIAL_POST_SCHEDULED", post_id=post_id) + await _audit(request, "publishing.scheduled", post_id=post_id) return schedule +@router.post("/posts/{post_id}/reschedule", response_model=SocialScheduleView) +async def reschedule_post( + request: Request, post_id: str, payload: SocialRescheduleRequest +) -> SocialScheduleView: + workspace_id, _ = _identity(request) + result = await _social(request).operations.reschedule(workspace_id, post_id, payload) + await _audit(request, "publishing.rescheduled", post_id=post_id) + return result + + @router.post("/posts/{post_id}/cancel", response_model=SocialPostView) async def cancel_post(request: Request, post_id: str) -> SocialPostView: workspace_id, _ = _identity(request) result = await _social(request).publishing.cancel(workspace_id, post_id) await _audit(request, "SOCIAL_POST_CANCELLED", post_id=post_id) + await _audit(request, "publishing.cancelled", post_id=post_id) + return result + + +@router.get("/calendar", response_model=PublishingCalendarView) +async def publishing_calendar( + request: Request, + starts_at: datetime = Query(), + ends_at: datetime = Query(), + offset: int = Query(default=0, ge=0), + limit: int = Query(default=100, ge=1, le=500), +) -> PublishingCalendarView: + workspace_id, _ = _identity(request) + return await _social(request).operations.calendar( + workspace_id, + starts_at=starts_at, + ends_at=ends_at, + offset=offset, + limit=limit, + ) + + +@router.get("/publishing-context", response_model=PublishingContextView) +async def publishing_context(request: Request) -> PublishingContextView: + workspace_id, _ = _identity(request) + return PublishingContextView( + timezone=await _social(request).operations.workspace_timezone(workspace_id) + ) + + +@router.get("/queue", response_model=PublishingQueueView) +async def publishing_queue( + request: Request, + offset: int = Query(default=0, ge=0), + limit: int = Query(default=50, ge=1, le=200), + queue_status: str | None = Query(default=None, alias="status", max_length=32), + provider: str | None = Query(default=None, max_length=32), + account_id: str | None = Query(default=None, max_length=120), + project_id: str | None = Query(default=None, max_length=36), + search: str | None = Query(default=None, max_length=200), +) -> PublishingQueueView: + workspace_id, _ = _identity(request) + return await _social(request).operations.queue( + workspace_id, + offset=offset, + limit=limit, + status=queue_status, + provider=provider, + account_id=account_id, + project_id=project_id, + search=search, + ) + + +@router.post( + "/bulk", + response_model=PublishingBatchView, + status_code=status.HTTP_202_ACCEPTED, +) +async def create_publishing_batch( + request: Request, + payload: PublishingBulkRequest, + idempotency_key: str = Header(alias="Idempotency-Key", min_length=8, max_length=255), +) -> PublishingBatchView: + workspace_id, user_id = _identity(request) + result = await _social(request).operations.create_batch( + workspace_id, + user_id, + payload, + idempotency_key=idempotency_key, + ) + await _audit(request, "publishing.bulk_started") + return result + + +@router.post( + "/posts/{post_id}/targets/{target_id}/retry", + response_model=SocialJobView, + status_code=status.HTTP_202_ACCEPTED, +) +async def retry_post_target( + request: Request, + post_id: str, + target_id: str, + idempotency_key: str = Header(alias="Idempotency-Key", min_length=8, max_length=255), +) -> SocialJobView: + workspace_id, _ = _identity(request) + result = await _social(request).publishing.retry_target( + workspace_id, + post_id, + target_id, + idempotency_key=idempotency_key, + ) + await _audit( + request, + "SOCIAL_TARGET_RETRY_QUEUED", + provider=result.provider.value if result.provider else None, + post_id=post_id, + job_id=result.id, + ) return result @@ -294,9 +494,7 @@ async def list_jobs( limit: int = Query(default=100, ge=1, le=500), ) -> list[SocialJobView]: workspace_id, _ = _identity(request) - return await _social(request).jobs.list( - workspace_id, offset=offset, limit=limit - ) + return await _social(request).jobs.list(workspace_id, offset=offset, limit=limit) @router.get("/jobs/{job_id}", response_model=SocialJobView) diff --git a/app/brand/__init__.py b/app/brand/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..51c51d92add70c12906427e4bee22b271ca35303 --- /dev/null +++ b/app/brand/__init__.py @@ -0,0 +1,2 @@ +"""Workspace-scoped brand kits and brand governance.""" + diff --git a/app/brand/api.py b/app/brand/api.py new file mode 100644 index 0000000000000000000000000000000000000000..29510cec3a326e55b37963c5e588ba775ed228c7 --- /dev/null +++ b/app/brand/api.py @@ -0,0 +1,99 @@ +from __future__ import annotations +from uuid import UUID +from fastapi import APIRouter, Request, Response, status +from app.brand.schemas import BrandKitCreate, BrandKitResponse +from app.security.errors import ForbiddenError + +router = APIRouter(prefix="/v1/brand", tags=["brand"]) + +def _identity(request: Request) -> tuple[str, str]: + context = request.state.auth + if not context.workspace_id or not context.user_id: + raise ForbiddenError + return context.workspace_id, context.user_id + +@router.post("", response_model=BrandKitResponse, status_code=status.HTTP_201_CREATED) +async def create_brand_kit(request: Request, payload: BrandKitCreate) -> BrandKitResponse: + workspace_id, user_id = _identity(request) + kit, version = await request.app.state.container.brand_kits.create_brand_kit( + workspace_id=workspace_id, + name=payload.name, + description=payload.description, + user_id=user_id, + ) + latest = await request.app.state.container.brand_kits.get_brand_kit_version( + workspace_id=workspace_id, + brand_kit_id=kit.id, + user_id=user_id, + ) + return BrandKitResponse( + id=kit.id, + name=kit.name, + description=kit.description, + status=kit.status, + is_default=kit.is_default, + active_version_id=latest.id, + created_at=kit.created_at.isoformat(), + updated_at=kit.updated_at.isoformat(), + ) + +@router.get("", response_model=list[BrandKitResponse]) +async def list_brand_kits(request: Request) -> list[BrandKitResponse]: + workspace_id, user_id = _identity(request) + kits = await request.app.state.container.brand_kits.list_brand_kits( + workspace_id=workspace_id, + user_id=user_id, + ) + response: list[BrandKitResponse] = [] + for kit in kits: + try: + latest = await request.app.state.container.brand_kits.get_brand_kit_version( + workspace_id=workspace_id, + brand_kit_id=kit.id, + user_id=user_id, + ) + active_version_id = latest.id + except Exception: + active_version_id = kit.active_version_id + response.append(BrandKitResponse( + id=kit.id, + name=kit.name, + description=kit.description, + status=kit.status, + is_default=kit.is_default, + active_version_id=active_version_id, + created_at=kit.created_at.isoformat(), + updated_at=kit.updated_at.isoformat(), + )) + return response + +@router.patch("/{brand_kit_id}", response_model=BrandKitResponse) +async def update_brand_kit(request: Request, brand_kit_id: UUID, payload: BrandKitCreate) -> BrandKitResponse: + workspace_id, user_id = _identity(request) + version = await request.app.state.container.brand_kits.update_brand_kit( + workspace_id=workspace_id, + brand_kit_id=str(brand_kit_id), + user_id=user_id, + data=payload.initial_version, + ) + kit = await request.app.state.container.brand_kits.get_brand_kit(workspace_id, str(brand_kit_id), user_id=user_id) + return BrandKitResponse( + id=kit.id, + name=kit.name, + description=kit.description, + status=kit.status, + is_default=kit.is_default, + active_version_id=version.id, + created_at=kit.created_at.isoformat(), + updated_at=kit.updated_at.isoformat(), + ) + +@router.delete("/{brand_kit_id}", status_code=status.HTTP_204_NO_CONTENT) +async def delete_brand_kit(request: Request, brand_kit_id: UUID) -> Response: + workspace_id, user_id = _identity(request) + await request.app.state.container.brand_kits.delete_brand_kit( + workspace_id=workspace_id, + brand_kit_id=str(brand_kit_id), + user_id=user_id, + ) + return Response(status_code=status.HTTP_204_NO_CONTENT) diff --git a/app/brand/capabilities.py b/app/brand/capabilities.py new file mode 100644 index 0000000000000000000000000000000000000000..0d97749bf30154ff805178084a48870c5ea1310f --- /dev/null +++ b/app/brand/capabilities.py @@ -0,0 +1,26 @@ +from app.brand.schemas import BrandCapabilities + + +def brand_capabilities(*, watermark_rendering: bool = False) -> BrandCapabilities: + return BrandCapabilities( + watermark_rendering="available" if watermark_rendering else "configuration_only", + supported_asset_roles=[ + "logo_primary", + "logo_secondary", + "watermark", + "favicon", + "social_avatar", + "social_cover", + ], + supported_providers=[ + "facebook", + "instagram", + "tiktok", + "x", + "youtube", + "linkedin", + "telegram", + "whatsapp", + ], + ) + diff --git a/app/brand/errors.py b/app/brand/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..3f47ee638921577a2f54823261400acf897f782c --- /dev/null +++ b/app/brand/errors.py @@ -0,0 +1,9 @@ +from app.core.exceptions import MediaAPIError + +class BrandKitNotFoundError(MediaAPIError): + code = "BRAND_KIT_NOT_FOUND" + status_code = 404 + +class BrandKitValidationError(MediaAPIError): + code = "INVALID_BRAND_KIT" + status_code = 422 diff --git a/app/brand/models.py b/app/brand/models.py new file mode 100644 index 0000000000000000000000000000000000000000..a8d2c456feefe7179cc70413c0525bf053123940 --- /dev/null +++ b/app/brand/models.py @@ -0,0 +1,165 @@ +from __future__ import annotations + +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import ( + JSON, + Boolean, + CheckConstraint, + DateTime, + ForeignKey, + Index, + Integer, + Numeric, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base, utcnow + + +class BrandKit(Base): + __tablename__ = "brand_kits" + __table_args__ = ( + CheckConstraint("status in ('active','archived')", name="ck_brand_kits_status"), + Index("ix_brand_kits_workspace_updated", "workspace_id", "updated_at"), + Index("ix_brand_kits_workspace_default", "workspace_id", "is_default"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + name: Mapped[str] = mapped_column(String(200), nullable=False) + description: Mapped[str | None] = mapped_column(Text) + status: Mapped[str] = mapped_column(String(16), nullable=False, default="active") + is_default: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) + active_version_id: Mapped[str | None] = mapped_column(String(36)) + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class BrandKitVersion(Base): + __tablename__ = "brand_kit_versions" + __table_args__ = ( + UniqueConstraint("brand_kit_id", "version_number", name="uq_brand_kit_version"), + CheckConstraint("status in ('draft','published','archived')", name="ck_brand_version_status"), + Index("ix_brand_versions_kit_created", "brand_kit_id", "created_at"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + brand_kit_id: Mapped[str] = mapped_column( + String(36), ForeignKey("brand_kits.id", ondelete="CASCADE"), nullable=False + ) + version_number: Mapped[int] = mapped_column(Integer, nullable=False) + status: Mapped[str] = mapped_column(String(16), nullable=False, default="draft") + logo_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + favicon_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + primary_color: Mapped[str | None] = mapped_column(String(32)) + secondary_color: Mapped[str | None] = mapped_column(String(32)) + accent_color: Mapped[str | None] = mapped_column(String(32)) + background_color: Mapped[str | None] = mapped_column(String(32)) + text_color: Mapped[str | None] = mapped_column(String(32)) + font_family_primary: Mapped[str | None] = mapped_column(String(255)) + font_family_secondary: Mapped[str | None] = mapped_column(String(255)) + heading_font: Mapped[str | None] = mapped_column(String(255)) + body_font: Mapped[str | None] = mapped_column(String(255)) + voice_style: Mapped[str | None] = mapped_column(String(255)) + tone: Mapped[str | None] = mapped_column(String(255)) + default_cta: Mapped[str | None] = mapped_column(String(255)) + default_hashtags: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=list) + watermark_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + watermark_position: Mapped[str | None] = mapped_column(String(32)) + watermark_opacity: Mapped[float | None] = mapped_column(Numeric(precision=3, scale=2)) + metadata: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + published_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class BrandKitAsset(Base): + __tablename__ = "brand_kit_assets" + __table_args__ = ( + UniqueConstraint("brand_kit_version_id", "role", name="uq_brand_asset_role"), + Index("ix_brand_assets_workspace_asset", "workspace_id", "asset_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + brand_kit_version_id: Mapped[str] = mapped_column( + String(36), ForeignKey("brand_kit_versions.id", ondelete="CASCADE"), nullable=False + ) + asset_id: Mapped[str] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="RESTRICT"), nullable=False + ) + role: Mapped[str] = mapped_column(String(32), nullable=False) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + +class BrandKitPlatformSetting(Base): + __tablename__ = "brand_kit_platform_settings" + __table_args__ = ( + UniqueConstraint( + "brand_kit_version_id", "provider", name="uq_brand_platform_setting" + ), + Index("ix_brand_platform_settings_workspace", "workspace_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + brand_kit_version_id: Mapped[str] = mapped_column( + String(36), ForeignKey("brand_kit_versions.id", ondelete="CASCADE"), nullable=False + ) + provider: Mapped[str] = mapped_column(String(32), nullable=False) + settings_json: Mapped[dict[str, object]] = mapped_column( + "settings", JSON, nullable=False, default=dict + ) + + +class BrandGovernanceSetting(Base): + __tablename__ = "brand_governance_settings" + __table_args__ = (UniqueConstraint("workspace_id", name="uq_brand_governance_workspace"),) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False + ) + require_brand_kit_for_publish: Mapped[bool] = mapped_column( + Boolean, nullable=False, default=False + ) + require_published_brand_version: Mapped[bool] = mapped_column( + Boolean, nullable=False, default=False + ) + allow_user_override_brand_defaults: Mapped[bool] = mapped_column( + Boolean, nullable=False, default=True + ) + updated_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + diff --git a/app/brand/models/__init__.py b/app/brand/models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..9731ddc5aeb3b407ef1d6b62c4ab02f79793e666 --- /dev/null +++ b/app/brand/models/__init__.py @@ -0,0 +1,3 @@ +from app.brand.models.brand import BrandKit, BrandKitVersion + +__all__ = ["BrandKit", "BrandKitVersion"] diff --git a/app/brand/models/brand.py b/app/brand/models/brand.py new file mode 100644 index 0000000000000000000000000000000000000000..b91695b5e76986c496fd15516ff2dfbe314d523d --- /dev/null +++ b/app/brand/models/brand.py @@ -0,0 +1,64 @@ +from __future__ import annotations +from datetime import datetime +from uuid import uuid4 +from sqlalchemy import String, JSON, DateTime, Integer, Index, ForeignKey, UniqueConstraint, Text, Numeric +from sqlalchemy.orm import Mapped, mapped_column +from app.security.models import Base, utcnow + +class BrandKit(Base): + __tablename__ = "brand_kits" + + __table_args__ = ( + Index("ix_brand_kits_workspace_updated", "workspace_id", "updated_at"), + UniqueConstraint("workspace_id", "name", name="uq_brand_kits_workspace_name"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + created_by: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False) + name: Mapped[str] = mapped_column(String(255), nullable=False) + description: Mapped[str | None] = mapped_column(Text) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="draft") # draft, published, archived + active_version_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("brand_kit_versions.id", ondelete="SET NULL")) + + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + archived_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + +class BrandKitVersion(Base): + __tablename__ = "brand_kit_versions" + + __table_args__ = ( + UniqueConstraint("brand_kit_id", "version_number", name="uq_brand_kit_version"), + Index("ix_brand_versions_kit_created", "brand_kit_id", "version_number"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + brand_kit_id: Mapped[str] = mapped_column(String(36), ForeignKey("brand_kits.id", ondelete="CASCADE"), nullable=False) + version_number: Mapped[int] = mapped_column(Integer, nullable=False) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="draft") + + # Branding Fields + logo_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + favicon_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + primary_color: Mapped[str | None] = mapped_column(String(7)) + secondary_color: Mapped[str | None] = mapped_column(String(7)) + accent_color: Mapped[str | None] = mapped_column(String(7)) + background_color: Mapped[str | None] = mapped_column(String(7)) + text_color: Mapped[str | None] = mapped_column(String(7)) + font_family_primary: Mapped[str | None] = mapped_column(String(100)) + font_family_secondary: Mapped[str | None] = mapped_column(String(100)) + heading_font: Mapped[str | None] = mapped_column(String(100)) + body_font: Mapped[str | None] = mapped_column(String(100)) + voice_style: Mapped[str | None] = mapped_column(String(50)) + tone: Mapped[str | None] = mapped_column(String(50)) + default_cta: Mapped[str | None] = mapped_column(String(100)) + default_hashtags: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + watermark_asset_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_assets.id", ondelete="SET NULL")) + watermark_position: Mapped[str | None] = mapped_column(String(20)) + watermark_opacity: Mapped[float | None] = mapped_column(Numeric(precision=3, scale=2)) + metadata: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) + + created_by: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + published_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) diff --git a/app/brand/repositories/brand_repository.py b/app/brand/repositories/brand_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..2e47e17f84149f6b9b5447eabc225eae0faeb481 --- /dev/null +++ b/app/brand/repositories/brand_repository.py @@ -0,0 +1,138 @@ +from __future__ import annotations +from sqlalchemy import select, func +from app.brand.models.brand import BrandKit, BrandKitVersion +from app.brand.errors import BrandKitNotFoundError +from app.security.database import SecurityDatabase +from app.security.models import utcnow + +class BrandKitRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + def _map_data_to_version_args(self, data: dict[str, object]) -> dict[str, object]: + return { + "logo_asset_id": data.get("logo_asset_id"), + "favicon_asset_id": data.get("favicon_asset_id"), + "primary_color": data.get("primary_color"), + "secondary_color": data.get("secondary_color"), + "accent_color": data.get("accent_color"), + "background_color": data.get("background_color"), + "text_color": data.get("text_color"), + "font_family_primary": data.get("font_family_primary"), + "font_family_secondary": data.get("font_family_secondary"), + "heading_font": data.get("heading_font"), + "body_font": data.get("body_font"), + "voice_style": data.get("voice_style"), + "tone": data.get("tone"), + "default_cta": data.get("default_cta"), + "default_hashtags": data.get("default_hashtags", {}), + "watermark_asset_id": data.get("watermark_asset_id"), + "watermark_position": data.get("watermark_position"), + "watermark_opacity": data.get("watermark_opacity"), + "metadata": data.get("metadata", {}), + } + + async def create(self, workspace_id: str, name: str, data: dict[str, object], *, user_id: str) -> tuple[BrandKit, BrandKitVersion]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = BrandKit(workspace_id=workspace_id, name=name) + session.add(kit) + await session.flush() # Populate ID + + version = BrandKitVersion( + brand_kit_id=kit.id, + version_number=1, + created_by=user_id, + **self._map_data_to_version_args(data) + ) + session.add(version) + + await session.commit() + await session.refresh(kit) + await session.refresh(version) + return kit, version + + async def get(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> BrandKit: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + ) + if kit is None: + raise BrandKitNotFoundError(f"Brand Kit {brand_kit_id} not found.") + return kit + + async def list(self, workspace_id: str, *, user_id: str) -> list[BrandKit]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return list((await session.scalars( + select(BrandKit).where(BrandKit.workspace_id == workspace_id) + )).all()) + + async def get_latest_version(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + version = await session.scalar( + select(BrandKitVersion) + .join(BrandKit) + .where(BrandKit.id == brand_kit_id, BrandKit.workspace_id == workspace_id) + .order_by(BrandKitVersion.version_number.desc()) + .limit(1) + ) + if version is None: + raise BrandKitNotFoundError(f"No version found for Brand Kit {brand_kit_id}.") + return version + + async def create_version(self, workspace_id: str, brand_kit_id: str, data: dict[str, object], *, user_id: str) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + # Check existence and lock + kit = await session.scalar( + select(BrandKit).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ).with_for_update() + ) + if kit is None: + raise BrandKitNotFoundError(f"Brand Kit {brand_kit_id} not found.") + + # Get latest version number + latest_version = await session.scalar( + select(func.max(BrandKitVersion.version_number)) + .where(BrandKitVersion.brand_kit_id == brand_kit_id) + ) + + new_version = BrandKitVersion( + brand_kit_id=brand_kit_id, + version_number=(latest_version or 0) + 1, + created_by=user_id, + **self._map_data_to_version_args(data) + ) + session.add(new_version) + kit.updated_at = utcnow() + await session.commit() + await session.refresh(new_version) + return new_version + + async def delete(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ).with_for_update() + ) + if kit is None: + raise BrandKitNotFoundError(f"Brand Kit {brand_kit_id} not found.") + await session.delete(kit) + await session.commit() diff --git a/app/brand/repository.py b/app/brand/repository.py new file mode 100644 index 0000000000000000000000000000000000000000..8ebcdb79f78b7d60f2fa20250d7c42c8610d6096 --- /dev/null +++ b/app/brand/repository.py @@ -0,0 +1,603 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import delete, func, or_, select, update +from sqlalchemy.exc import IntegrityError + +from app.brand.errors import ( + BrandConflictError, + BrandImmutableError, + BrandNotFoundError, + BrandVersionNotFoundError, +) +from app.brand.models import ( + BrandGovernanceSetting, + BrandKit, + BrandKitAsset, + BrandKitPlatformSetting, + BrandKitVersion, +) +from app.projects.models import Project +from app.security.database import SecurityDatabase +from app.security.models import CanonicalMediaAsset + + +class BrandRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def list( + self, + workspace_id: str, + *, + user_id: str, + search: str | None, + status: str, + offset: int, + limit: int, + ) -> tuple[list[BrandKit], int]: + predicates = [ + BrandKit.workspace_id == workspace_id, + BrandKit.status == status, + ] + if search: + term = search.casefold() + predicates.append( + or_( + func.lower(BrandKit.name).contains(term, autoescape=True), + func.lower(BrandKit.description).contains(term, autoescape=True), + ) + ) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + total = int( + await session.scalar( + select(func.count()).select_from(BrandKit).where(*predicates) + ) + or 0 + ) + items = list( + ( + await session.scalars( + select(BrandKit) + .where(*predicates) + .order_by(BrandKit.is_default.desc(), BrandKit.updated_at.desc()) + .offset(offset) + .limit(limit) + ) + ).all() + ) + return items, total + + async def create( + self, kit: BrandKit, version: BrandKitVersion, *, user_id: str + ) -> tuple[BrandKit, BrandKitVersion]: + async with self.database.tenant_session( + workspace_id=kit.workspace_id, user_id=user_id + ) as session: + if kit.is_default: + await session.execute( + update(BrandKit) + .where(BrandKit.workspace_id == kit.workspace_id) + .values(is_default=False) + ) + session.add(kit) + await session.flush() + version.brand_kit_id = kit.id + session.add(version) + try: + await session.commit() + except IntegrityError as exc: + await session.rollback() + raise BrandConflictError("Brand kit creation conflicted with existing data.") from exc + await session.refresh(kit) + await session.refresh(version) + return kit, version + + async def get( + self, workspace_id: str, brand_kit_id: str, *, user_id: str, mutable: bool = False + ) -> BrandKit: + query = select(BrandKit).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + if mutable: + query = query.with_for_update() + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar(query) + if kit is None: + raise BrandNotFoundError("Brand kit was not found in this workspace.") + return kit + + async def update( + self, + workspace_id: str, + brand_kit_id: str, + *, + user_id: str, + fields: dict[str, object], + ) -> BrandKit: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Brand kit was not found in this workspace.") + for key, value in fields.items(): + setattr(kit, key, value) + kit.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(kit) + return kit + + async def archive( + self, workspace_id: str, brand_kit_id: str, *, user_id: str + ) -> BrandKit: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Brand kit was not found in this workspace.") + if kit.is_default: + raise BrandConflictError("The default brand kit must be replaced before archiving.") + kit.status = "archived" + kit.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(kit) + return kit + + async def set_default( + self, workspace_id: str, brand_kit_id: str, *, user_id: str + ) -> BrandKit: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + BrandKit.status == "active", + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Active brand kit was not found in this workspace.") + await session.execute( + update(BrandKit) + .where(BrandKit.workspace_id == workspace_id) + .values(is_default=False) + ) + kit.is_default = True + kit.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(kit) + return kit + + async def versions( + self, workspace_id: str, brand_kit_id: str, *, user_id: str + ) -> list[BrandKitVersion]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + if not await session.scalar( + select(BrandKit.id).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + ): + raise BrandNotFoundError("Brand kit was not found in this workspace.") + return list( + ( + await session.scalars( + select(BrandKitVersion) + .where( + BrandKitVersion.workspace_id == workspace_id, + BrandKitVersion.brand_kit_id == brand_kit_id, + ) + .order_by(BrandKitVersion.version_number.desc()) + ) + ).all() + ) + + async def version( + self, + workspace_id: str, + brand_kit_id: str, + version_id: str, + *, + user_id: str, + ) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + item = await session.scalar( + select(BrandKitVersion).where( + BrandKitVersion.id == version_id, + BrandKitVersion.brand_kit_id == brand_kit_id, + BrandKitVersion.workspace_id == workspace_id, + ) + ) + if item is None: + raise BrandVersionNotFoundError( + "Brand kit version was not found in this workspace." + ) + return item + + async def create_version( + self, version: BrandKitVersion, *, user_id: str + ) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=version.workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == version.brand_kit_id, + BrandKit.workspace_id == version.workspace_id, + BrandKit.status == "active", + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Active brand kit was not found in this workspace.") + version.version_number = ( + int( + await session.scalar( + select(func.max(BrandKitVersion.version_number)).where( + BrandKitVersion.brand_kit_id == version.brand_kit_id + ) + ) + or 0 + ) + + 1 + ) + session.add(version) + await session.commit() + await session.refresh(version) + return version + + async def update_version( + self, + workspace_id: str, + brand_kit_id: str, + version_id: str, + *, + user_id: str, + fields: dict[str, object], + ) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == version_id, + BrandKitVersion.brand_kit_id == brand_kit_id, + BrandKitVersion.workspace_id == workspace_id, + ) + .with_for_update() + ) + if version is None: + raise BrandVersionNotFoundError( + "Brand kit version was not found in this workspace." + ) + if version.status != "draft": + raise BrandImmutableError("Published or archived brand versions are immutable.") + for key, value in fields.items(): + setattr(version, key, value) + await session.commit() + await session.refresh(version) + return version + + async def publish_version( + self, workspace_id: str, brand_kit_id: str, version_id: str, *, user_id: str + ) -> tuple[BrandKit, BrandKitVersion]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + BrandKit.status == "active", + ) + .with_for_update() + ) + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == version_id, + BrandKitVersion.brand_kit_id == brand_kit_id, + BrandKitVersion.workspace_id == workspace_id, + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Active brand kit was not found in this workspace.") + if version is None: + raise BrandVersionNotFoundError( + "Brand kit version was not found in this workspace." + ) + if version.status != "draft": + raise BrandConflictError("Only a draft brand version can be published.") + now = datetime.now(timezone.utc) + version.status = "published" + version.published_at = now + kit.active_version_id = version.id + kit.updated_at = now + await session.commit() + await session.refresh(kit) + await session.refresh(version) + return kit, version + + async def archive_version( + self, workspace_id: str, brand_kit_id: str, version_id: str, *, user_id: str + ) -> BrandKitVersion: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit) + .where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + ) + .with_for_update() + ) + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == version_id, + BrandKitVersion.brand_kit_id == brand_kit_id, + BrandKitVersion.workspace_id == workspace_id, + ) + .with_for_update() + ) + if kit is None or version is None: + raise BrandVersionNotFoundError( + "Brand kit version was not found in this workspace." + ) + if kit.active_version_id == version.id: + raise BrandConflictError("The active published version cannot be archived.") + version.status = "archived" + await session.commit() + await session.refresh(version) + return version + + async def replace_platform_settings( + self, + workspace_id: str, + version_id: str, + settings: list[BrandKitPlatformSetting], + *, + user_id: str, + ) -> None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == version_id, + BrandKitVersion.workspace_id == workspace_id, + ) + .with_for_update() + ) + if version is None: + raise BrandVersionNotFoundError("Brand kit version was not found.") + if version.status != "draft": + raise BrandImmutableError("Published or archived brand versions are immutable.") + await session.execute( + delete(BrandKitPlatformSetting).where( + BrandKitPlatformSetting.brand_kit_version_id == version_id + ) + ) + session.add_all(settings) + await session.commit() + + async def platform_settings( + self, workspace_id: str, version_id: str, *, user_id: str + ) -> list[BrandKitPlatformSetting]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return list( + ( + await session.scalars( + select(BrandKitPlatformSetting).where( + BrandKitPlatformSetting.workspace_id == workspace_id, + BrandKitPlatformSetting.brand_kit_version_id == version_id, + ) + ) + ).all() + ) + + async def assets( + self, workspace_id: str, version_id: str, *, user_id: str + ) -> list[tuple[BrandKitAsset, CanonicalMediaAsset]]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + rows = await session.execute( + select(BrandKitAsset, CanonicalMediaAsset) + .join(CanonicalMediaAsset, CanonicalMediaAsset.id == BrandKitAsset.asset_id) + .where( + BrandKitAsset.workspace_id == workspace_id, + BrandKitAsset.brand_kit_version_id == version_id, + CanonicalMediaAsset.workspace_id == workspace_id, + ) + ) + return list(rows.all()) + + async def attach_asset( + self, association: BrandKitAsset, *, user_id: str + ) -> BrandKitAsset: + async with self.database.tenant_session( + workspace_id=association.workspace_id, user_id=user_id + ) as session: + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == association.brand_kit_version_id, + BrandKitVersion.workspace_id == association.workspace_id, + ) + .with_for_update() + ) + if version is None: + raise BrandVersionNotFoundError("Brand kit version was not found.") + if version.status != "draft": + raise BrandImmutableError("Assets on a published brand version are immutable.") + existing = await session.scalar( + select(BrandKitAsset).where( + BrandKitAsset.brand_kit_version_id + == association.brand_kit_version_id, + BrandKitAsset.role == association.role, + ) + ) + if existing: + existing.asset_id = association.asset_id + await session.commit() + await session.refresh(existing) + return existing + session.add(association) + await session.commit() + await session.refresh(association) + return association + + async def detach_asset( + self, + workspace_id: str, + version_id: str, + asset_id: str, + *, + user_id: str, + ) -> None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + version = await session.scalar( + select(BrandKitVersion) + .where( + BrandKitVersion.id == version_id, + BrandKitVersion.workspace_id == workspace_id, + ) + .with_for_update() + ) + if version is None: + raise BrandVersionNotFoundError("Brand kit version was not found.") + if version.status != "draft": + raise BrandImmutableError("Assets on a published brand version are immutable.") + result = await session.execute( + delete(BrandKitAsset).where( + BrandKitAsset.workspace_id == workspace_id, + BrandKitAsset.brand_kit_version_id == version_id, + BrandKitAsset.asset_id == asset_id, + ) + ) + if not result.rowcount: + raise BrandVersionNotFoundError("Brand asset association was not found.") + await session.commit() + + async def governance( + self, workspace_id: str, *, user_id: str + ) -> BrandGovernanceSetting | None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return await session.scalar( + select(BrandGovernanceSetting).where( + BrandGovernanceSetting.workspace_id == workspace_id + ) + ) + + async def set_governance( + self, + setting: BrandGovernanceSetting, + *, + user_id: str, + ) -> BrandGovernanceSetting: + async with self.database.tenant_session( + workspace_id=setting.workspace_id, user_id=user_id + ) as session: + existing = await session.scalar( + select(BrandGovernanceSetting) + .where(BrandGovernanceSetting.workspace_id == setting.workspace_id) + .with_for_update() + ) + if existing is None: + session.add(setting) + existing = setting + else: + existing.require_brand_kit_for_publish = ( + setting.require_brand_kit_for_publish + ) + existing.require_published_brand_version = ( + setting.require_published_brand_version + ) + existing.allow_user_override_brand_defaults = ( + setting.allow_user_override_brand_defaults + ) + existing.updated_by = setting.updated_by + existing.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(existing) + return existing + + async def apply_to_project( + self, + workspace_id: str, + brand_kit_id: str, + project_id: str, + *, + user_id: str, + ) -> Project: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + kit = await session.scalar( + select(BrandKit).where( + BrandKit.id == brand_kit_id, + BrandKit.workspace_id == workspace_id, + BrandKit.status == "active", + ) + ) + project = await session.scalar( + select(Project) + .where( + Project.id == project_id, + Project.workspace_id == workspace_id, + Project.status == "active", + ) + .with_for_update() + ) + if kit is None: + raise BrandNotFoundError("Active brand kit was not found in this workspace.") + if project is None: + raise BrandNotFoundError("Project was not found in this workspace.") + project.brand_kit_id = kit.id + project.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(project) + return project + diff --git a/app/brand/schemas.py b/app/brand/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..eb0d61726d5f7138efe8889deff5316723b0eeaf --- /dev/null +++ b/app/brand/schemas.py @@ -0,0 +1,20 @@ +from pydantic import BaseModel, Field +from datetime import datetime + +class BrandKitCreate(BaseModel): + name: str = Field(..., min_length=1, max_length=255) + description: str | None = None + initial_version: dict[str, object] # For now, keep as dict for simplicity, can be updated later + +class BrandKitResponse(BaseModel): + id: str + name: str + description: str | None + status: str + is_default: bool + active_version_id: str | None + created_at: datetime + updated_at: datetime + + class Config: + from_attributes = True diff --git a/app/brand/service.py b/app/brand/service.py new file mode 100644 index 0000000000000000000000000000000000000000..67b888ff5647659360183ac34ea6285022557c8a --- /dev/null +++ b/app/brand/service.py @@ -0,0 +1,636 @@ +from __future__ import annotations + +from app.brand.capabilities import brand_capabilities +from app.brand.errors import BrandAssetInvalidError, BrandGovernanceError +from app.brand.models import ( + BrandGovernanceSetting, + BrandKit, + BrandKitAsset, + BrandKitPlatformSetting, + BrandKitVersion, +) +from app.brand.repository import BrandRepository +from app.brand.schemas import ( + BrandApplyProject, + BrandAssetAttach, + BrandAssetView, + BrandCapabilities, + BrandGovernance, + BrandKitCreate, + BrandKitList, + BrandKitUpdate, + BrandKitView, + BrandVersionContent, + BrandVersionCreate, + BrandVersionUpdate, + BrandVersionView, +) +from app.brand.validation import validate_platform_settings +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService + + +class BrandService: + def __init__( + self, + repository: BrandRepository, + assets: CanonicalAssetService, + audit: AuditService, + ) -> None: + self.repository = repository + self.assets = assets + self.audit = audit + + def capabilities(self) -> BrandCapabilities: + # Configuration is authoritative; the existing render compiler does not + # yet accept brand watermark instructions as a first-class contract. + return brand_capabilities(watermark_rendering=False) + + async def list( + self, + *, + workspace_id: str, + user_id: str, + search: str | None, + status: str, + offset: int, + limit: int, + ) -> BrandKitList: + kits, total = await self.repository.list( + workspace_id, + user_id=user_id, + search=search, + status=status, + offset=offset, + limit=limit, + ) + return BrandKitList( + items=[await self._kit_view(kit, user_id=user_id) for kit in kits], + offset=offset, + limit=limit, + total=total, + ) + + async def get( + self, *, workspace_id: str, user_id: str, brand_kit_id: str + ) -> BrandKitView: + return await self._kit_view( + await self.repository.get(workspace_id, brand_kit_id, user_id=user_id), + user_id=user_id, + ) + + async def create( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + payload: BrandKitCreate, + ) -> BrandKitView: + self._validate_content(payload.initial_version) + kit, version = await self.repository.create( + BrandKit( + workspace_id=workspace_id, + name=payload.name, + description=payload.description, + is_default=payload.is_default, + created_by=user_id, + ), + self._version_model( + workspace_id, "", user_id, payload.initial_version, version=1 + ), + user_id=user_id, + ) + await self._replace_platform_settings(version, payload.initial_version, user_id) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.created", + kit.id, + {"is_default": kit.is_default}, + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.version_created", + kit.id, + {"version_id": version.id, "version": version.version}, + ) + return await self.get( + workspace_id=workspace_id, user_id=user_id, brand_kit_id=kit.id + ) + + async def update( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + payload: BrandKitUpdate, + ) -> BrandKitView: + kit = await self.repository.update( + workspace_id, + brand_kit_id, + user_id=user_id, + fields=payload.model_dump(exclude_unset=True), + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.updated", + kit.id, + {"fields": sorted(payload.model_fields_set)}, + ) + return await self._kit_view(kit, user_id=user_id) + + async def delete( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + ) -> None: + kit = await self.repository.archive(workspace_id, brand_kit_id, user_id=user_id) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.deleted", + kit.id, + {"disposition": "archived"}, + ) + + async def versions( + self, *, workspace_id: str, user_id: str, brand_kit_id: str + ) -> list[BrandVersionView]: + return [ + await self._version_view(item, user_id=user_id) + for item in await self.repository.versions( + workspace_id, brand_kit_id, user_id=user_id + ) + ] + + async def version( + self, + *, + workspace_id: str, + user_id: str, + brand_kit_id: str, + version_id: str, + ) -> BrandVersionView: + return await self._version_view( + await self.repository.version( + workspace_id, brand_kit_id, version_id, user_id=user_id + ), + user_id=user_id, + ) + + async def create_version( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + payload: BrandVersionCreate, + ) -> BrandVersionView: + content: BrandVersionContent = payload + if payload.source_version_id: + source = await self.version( + workspace_id=workspace_id, + user_id=user_id, + brand_kit_id=brand_kit_id, + version_id=payload.source_version_id, + ) + if not any( + field in payload.model_fields_set + for field in BrandVersionContent.model_fields + ): + content = BrandVersionContent.model_validate( + source.model_dump( + include=set(BrandVersionContent.model_fields), mode="json" + ) + ) + self._validate_content(content) + version = await self.repository.create_version( + self._version_model(workspace_id, brand_kit_id, user_id, content, version=0), + user_id=user_id, + ) + await self._replace_platform_settings(version, content, user_id) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.version_created", + brand_kit_id, + {"version_id": version.id, "version": version.version}, + ) + return await self._version_view(version, user_id=user_id) + + async def update_version( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + version_id: str, + payload: BrandVersionUpdate, + ) -> BrandVersionView: + self._validate_content(payload) + version = await self.repository.update_version( + workspace_id, + brand_kit_id, + version_id, + user_id=user_id, + fields=self._version_fields(payload), + ) + await self._replace_platform_settings(version, payload, user_id) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.updated", + brand_kit_id, + {"version_id": version.id}, + ) + return await self._version_view(version, user_id=user_id) + + async def publish_version( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + version_id: str, + ) -> BrandKitView: + kit, version = await self.repository.publish_version( + workspace_id, brand_kit_id, version_id, user_id=user_id + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.version_published", + kit.id, + {"version_id": version.id, "version": version.version}, + ) + return await self._kit_view(kit, user_id=user_id) + + async def archive_version( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + version_id: str, + ) -> BrandVersionView: + version = await self.repository.archive_version( + workspace_id, brand_kit_id, version_id, user_id=user_id + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.version_archived", + brand_kit_id, + {"version_id": version.id, "version": version.version}, + ) + return await self._version_view(version, user_id=user_id) + + async def set_default( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + ) -> BrandKitView: + kit = await self.repository.set_default( + workspace_id, brand_kit_id, user_id=user_id + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.default_changed", + kit.id, + {}, + ) + return await self._kit_view(kit, user_id=user_id) + + async def attach_asset( + self, + *, + workspace_id: str, + user_id: str, + brand_kit_id: str, + payload: BrandAssetAttach, + ) -> BrandAssetView: + version = await self.repository.version( + workspace_id, brand_kit_id, payload.version_id, user_id=user_id + ) + if version.status != "draft": + raise BrandAssetInvalidError( + "Assets may only be changed on a draft brand version." + ) + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, user_id=user_id, asset_id=payload.asset_id + ) + except CanonicalAssetNotFoundError as exc: + raise BrandAssetInvalidError( + "Brand asset is not an owned canonical asset." + ) from exc + if not asset.mime_type.startswith("image/"): + raise BrandAssetInvalidError("Brand logo and watermark assets must be images.") + association = await self.repository.attach_asset( + BrandKitAsset( + workspace_id=workspace_id, + brand_kit_version_id=version.id, + asset_id=asset.id, + role=payload.role, + ), + user_id=user_id, + ) + return BrandAssetView( + id=association.id, + version_id=version.id, + asset_id=asset.id, + role=payload.role, + mime_type=asset.mime_type, + filename=asset.filename, + ) + + async def detach_asset( + self, + *, + workspace_id: str, + user_id: str, + brand_kit_id: str, + version_id: str, + asset_id: str, + ) -> None: + await self.repository.version( + workspace_id, brand_kit_id, version_id, user_id=user_id + ) + await self.repository.detach_asset( + workspace_id, version_id, asset_id, user_id=user_id + ) + + async def governance( + self, *, workspace_id: str, user_id: str + ) -> BrandGovernance: + setting = await self.repository.governance(workspace_id, user_id=user_id) + return BrandGovernance.model_validate(setting, from_attributes=True) if setting else BrandGovernance() + + async def set_governance( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + payload: BrandGovernance, + ) -> BrandGovernance: + setting = await self.repository.set_governance( + BrandGovernanceSetting( + workspace_id=workspace_id, + updated_by=user_id, + **payload.model_dump(), + ), + user_id=user_id, + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.updated", + setting.id, + {"governance": True}, + ) + return BrandGovernance.model_validate(setting, from_attributes=True) + + async def apply_to_project( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + brand_kit_id: str, + payload: BrandApplyProject, + ) -> BrandKitView: + await self.repository.apply_to_project( + workspace_id, + brand_kit_id, + payload.project_id, + user_id=user_id, + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "brand.applied_to_project", + brand_kit_id, + {"project_id": payload.project_id}, + ) + return await self.get( + workspace_id=workspace_id, user_id=user_id, brand_kit_id=brand_kit_id + ) + + async def validate_publish( + self, + *, + workspace_id: str, + user_id: str, + project_id: str | None, + brand_kit_id: str | None, + brand_kit_version_id: str | None, + ) -> tuple[str | None, str | None]: + governance = await self.governance(workspace_id=workspace_id, user_id=user_id) + if governance.require_brand_kit_for_publish and not brand_kit_id: + raise BrandGovernanceError("A brand kit is required before publishing.") + if not brand_kit_id: + return None, None + kit = await self.repository.get(workspace_id, brand_kit_id, user_id=user_id) + selected_version = brand_kit_version_id or kit.active_version_id + if governance.require_published_brand_version and not selected_version: + raise BrandGovernanceError("A published brand version is required before publishing.") + if selected_version: + version = await self.repository.version( + workspace_id, kit.id, selected_version, user_id=user_id + ) + if governance.require_published_brand_version and version.status != "published": + raise BrandGovernanceError("Publishing requires a published brand version.") + return kit.id, selected_version + + @staticmethod + def _validate_content(content: BrandVersionContent) -> None: + for item in content.platform_settings: + validate_platform_settings(item.provider, item.settings) + + @staticmethod + def _version_fields(content: BrandVersionContent) -> dict[str, object]: + return { + "colors_json": [item.model_dump(mode="json") for item in content.colors], + "typography_json": content.typography.model_dump(mode="json"), + "voice_json": content.voice.model_dump(mode="json"), + "ctas_json": content.ctas.model_dump(mode="json"), + "hashtags_json": content.hashtags.model_dump(mode="json"), + "watermark_json": content.watermark.model_dump(mode="json"), + "ai_guidance": content.ai_guidance, + "metadata_json": content.metadata, + } + + def _version_model( + self, + workspace_id: str, + brand_kit_id: str, + user_id: str, + content: BrandVersionContent, + *, + version: int, + ) -> BrandKitVersion: + return BrandKitVersion( + workspace_id=workspace_id, + brand_kit_id=brand_kit_id, + version=version, + status="draft", + created_by=user_id, + **self._version_fields(content), + ) + + async def _replace_platform_settings( + self, version: BrandKitVersion, content: BrandVersionContent, user_id: str + ) -> None: + await self.repository.replace_platform_settings( + version.workspace_id, + version.id, + [ + BrandKitPlatformSetting( + workspace_id=version.workspace_id, + brand_kit_version_id=version.id, + provider=item.provider, + settings_json=item.settings, + ) + for item in content.platform_settings + ], + user_id=user_id, + ) + + async def _kit_view(self, kit: BrandKit, *, user_id: str) -> BrandKitView: + active = ( + await self.repository.version( + kit.workspace_id, + kit.id, + kit.active_version_id, + user_id=user_id, + ) + if kit.active_version_id + else None + ) + return BrandKitView( + id=kit.id, + workspace_id=kit.workspace_id, + name=kit.name, + description=kit.description, + status=kit.status, + is_default=kit.is_default, + active_version_id=kit.active_version_id, + active_version=( + await self._version_view(active, user_id=user_id) if active else None + ), + created_by=kit.created_by, + created_at=kit.created_at, + updated_at=kit.updated_at, + ) + + async def _version_view( + self, version: BrandKitVersion, *, user_id: str + ) -> BrandVersionView: + settings = await self.repository.platform_settings( + version.workspace_id, version.id, user_id=user_id + ) + assets = await self.repository.assets( + version.workspace_id, version.id, user_id=user_id + ) + return BrandVersionView( + id=version.id, + brand_kit_id=version.brand_kit_id, + version=version.version, + status=version.status, + colors=version.colors_json, + typography=version.typography_json, + voice=version.voice_json, + ctas=version.ctas_json, + hashtags=version.hashtags_json, + watermark=version.watermark_json, + ai_guidance=version.ai_guidance, + platform_settings=[ + {"provider": item.provider, "settings": item.settings_json} + for item in settings + ], + metadata=version.metadata_json, + assets=[ + BrandAssetView( + id=association.id, + version_id=version.id, + asset_id=asset.id, + role=association.role, + mime_type=asset.mime_type, + filename=asset.filename, + ) + for association, asset in assets + ], + created_by=version.created_by, + created_at=version.created_at, + published_at=version.published_at, + ) + + async def _audit( + self, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + event_type: str, + entity_id: str, + metadata: dict[str, object], + ) -> None: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type=event_type, + entity_type="brand_kit", + entity_id=entity_id, + metadata=metadata, + ) diff --git a/app/brand/services/brand_service.py b/app/brand/services/brand_service.py new file mode 100644 index 0000000000000000000000000000000000000000..1aa9b130d1176be37abcde7bfdbbec6d1fee3f8e --- /dev/null +++ b/app/brand/services/brand_service.py @@ -0,0 +1,64 @@ +from __future__ import annotations +from app.brand.models.brand import BrandKit, BrandKitVersion +from app.brand.repositories.brand_repository import BrandKitRepository +from app.security.assets import CanonicalAssetService +from app.security.audit import AuditService + +class BrandKitService: + def __init__(self, repository: BrandKitRepository, assets: CanonicalAssetService, audit: AuditService) -> None: + self.repository = repository + self.assets = assets + self.audit = audit + + async def _validate_assets(self, workspace_id: str, data: dict[str, object]) -> None: + for asset_key in ["logo_asset_id", "favicon_asset_id", "watermark_asset_id"]: + asset_id = data.get(asset_key) + if asset_id: + asset = await self.assets.get_asset(workspace_id=workspace_id, asset_id=str(asset_id)) + if not asset: + raise ValueError(f"Asset {asset_id} not found or inaccessible in this workspace.") + + async def create_brand_kit(self, workspace_id: str, name: str, data: dict[str, object], *, user_id: str) -> tuple[BrandKit, BrandKitVersion]: + await self._validate_assets(workspace_id, data) + kit, version = await self.repository.create(workspace_id, name, data, user_id=user_id) + await self.audit.log_event( + workspace_id=workspace_id, + actor_user_id=user_id, + event_type="brand_kit.created", + entity_type="brand_kit", + entity_id=kit.id, + metadata={"name": kit.name} + ) + return kit, version + + async def get_brand_kit(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> BrandKit: + return await self.repository.get(workspace_id, brand_kit_id, user_id=user_id) + + async def list_brand_kits(self, workspace_id: str, *, user_id: str) -> list[BrandKit]: + return await self.repository.list(workspace_id, user_id=user_id) + + async def get_brand_kit_version(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> BrandKitVersion: + return await self.repository.get_latest_version(workspace_id, brand_kit_id, user_id=user_id) + + async def update_brand_kit(self, workspace_id: str, brand_kit_id: str, *, user_id: str, data: dict[str, object]) -> BrandKitVersion: + await self._validate_assets(workspace_id, data) + version = await self.repository.create_version(workspace_id, brand_kit_id, data, user_id=user_id) + await self.audit.log_event( + workspace_id=workspace_id, + actor_user_id=user_id, + event_type="brand_kit.version_created", + entity_type="brand_kit_version", + entity_id=version.id, + metadata={"brand_kit_id": brand_kit_id, "version_number": version.version_number} + ) + return version + + async def delete_brand_kit(self, workspace_id: str, brand_kit_id: str, *, user_id: str) -> None: + await self.repository.delete(workspace_id, brand_kit_id, user_id=user_id) + await self.audit.log_event( + workspace_id=workspace_id, + actor_user_id=user_id, + event_type="brand_kit.archived", + entity_type="brand_kit", + entity_id=brand_kit_id + ) diff --git a/app/brand/services/validation_service.py b/app/brand/services/validation_service.py new file mode 100644 index 0000000000000000000000000000000000000000..cfc426e83490ee243b904e69c396f796a9f04986 --- /dev/null +++ b/app/brand/services/validation_service.py @@ -0,0 +1,29 @@ +from __future__ import annotations +from typing import Any +from app.brand.models.brand import BrandKitVersion + +class BrandKitValidationService: + def validate(self, version: BrandKitVersion) -> dict[str, Any]: + """ + Deterministic validation service. + Returns a dict with 'valid' boolean and 'issues' list of dicts: + {'severity': 'error'|'warning', 'field': str, 'reason': str} + """ + issues = [] + + # Required assets check + if not version.logo_asset_id: + issues.append({'severity': 'error', 'field': 'logo_asset_id', 'reason': 'Primary logo is required.'}) + + # Color completeness + if not version.primary_color: + issues.append({'severity': 'warning', 'field': 'primary_color', 'reason': 'Primary color is not set.'}) + + # Font checks + if not version.font_family_primary: + issues.append({'severity': 'warning', 'field': 'font_family_primary', 'reason': 'Primary font is not set.'}) + + return { + 'valid': len([i for i in issues if i['severity'] == 'error']) == 0, + 'issues': issues + } diff --git a/app/brand/validation.py b/app/brand/validation.py new file mode 100644 index 0000000000000000000000000000000000000000..900053ff2a540cca19a24a7e651d3eae74fb91a6 --- /dev/null +++ b/app/brand/validation.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from typing import Any + +from app.brand.errors import BrandValidationError + +_PLATFORM_FIELDS: dict[str, frozenset[str]] = { + "facebook": frozenset({"caption_style", "hashtags", "cta"}), + "instagram": frozenset({"caption_style", "hashtags", "cta", "first_comment"}), + "tiktok": frozenset({"caption_style", "hashtags", "cta"}), + "x": frozenset({"caption_style", "hashtags", "cta"}), + "youtube": frozenset({"title_pattern", "description", "tags", "cta"}), + "linkedin": frozenset({"post_style", "hashtags", "cta"}), + "telegram": frozenset({"caption_style", "cta"}), + "whatsapp": frozenset({"caption_style", "cta"}), +} + + +def validate_platform_settings(provider: str, settings: dict[str, Any]) -> None: + unsupported = set(settings) - _PLATFORM_FIELDS.get(provider, frozenset()) + if unsupported: + raise BrandValidationError( + f"Unsupported {provider} brand defaults: {', '.join(sorted(unsupported))}." + ) + for key, value in settings.items(): + if isinstance(value, str) and len(value) > 4000: + raise BrandValidationError(f"{provider}.{key} exceeds the supported length.") + if isinstance(value, list) and len(value) > 100: + raise BrandValidationError(f"{provider}.{key} contains too many values.") + diff --git a/app/container.py b/app/container.py index aa96332f3056cdb1d23c4b863e8918dbcb9e8d99..98ad0d604e41104e08c7b41838824c2ab5a36f5a 100644 --- a/app/container.py +++ b/app/container.py @@ -2,11 +2,48 @@ from __future__ import annotations from dataclasses import dataclass +from app.ai.service import AiStudioService +from app.brand.repositories.brand_repository import BrandKitRepository +from app.brand.services.brand_service import BrandKitService +from app.analytics.service import AnalyticsDomainService +from app.analytics.workers.sync import AnalyticsSyncWorker +from app.copilot.actions import CopilotActionRegistry +from app.copilot.planner import CopilotPlanner +from app.copilot.repository import CopilotRepository +from app.copilot.service import CopilotService from app.core.config import Settings +from app.generation.model_registry import GenerationModelRegistration, GenerationModelRegistry +from app.generation.providers.flux import ( + FLUX_BASE_MODEL_ID, + FLUX_DISTILLED_MODEL_ID, + FLUX_MODEL_CAPABILITY, + FLUX_PROVIDER_ID, + FLUX_TASK, + FluxProviderAdapter, +) +from app.generation.providers.registry import GenerationProviderRegistry +from app.generation.providers.wan import WAN_MODEL_CAPABILITY, WAN_PROVIDER_ID, WanProviderAdapter +from app.generation.repositories.generation import GenerationRepository +from app.generation.services.generation_service import GenerationService +from app.generation.services.output_ingestion import GenerationOutputIngestor +from app.generation.workers.generation_worker import GenerationWorker +from app.projects.repositories.editor_repository import ProjectEditorRepository +from app.projects.repositories.project_repository import ProjectRepository +from app.projects.repositories.render_repository import ProjectRenderRepository +from app.projects.repositories.collaboration_repository import CollaborationRepository +from app.projects.repositories.approval_repository import ApprovalRepository +from app.projects.services.editor_service import ProjectEditorService +from app.projects.services.project_service import ProjectService +from app.projects.services.collaboration_service import CollaborationService +from app.projects.services.approval_service import ApprovalService +from app.projects.services.render_service import ProjectRenderService +from app.projects.workers.render_worker import ProjectRenderWorker +from app.security.assets import CanonicalAssetService from app.security.audit import AuditService from app.security.database import SecurityDatabase from app.security.rate_limit import APIKeyRateLimiter from app.security.service import APIKeyService +from app.security.tenancy import TenantService from app.services.cleanup import CleanupService from app.services.downloader import Downloader from app.services.ffmpeg_service import FFmpegService @@ -22,6 +59,7 @@ from app.social.oauth.state import OAuthStateService from app.social.providers.registry import build_provider_registry from app.social.repositories.accounts import AccountRepository from app.social.repositories.assets import SocialMediaAssetRepository +from app.social.repositories.batches import PublishingBatchRepository from app.social.repositories.jobs import JobRepository from app.social.repositories.posts import PostRepository from app.social.repositories.tokens import TokenRepository @@ -31,14 +69,20 @@ from app.social.services.audit_service import SocialAuditService from app.social.services.job_service import JobService from app.social.services.media_asset_service import SocialMediaAssetService from app.social.services.oauth_service import OAuthService +from app.social.services.publishing_operations_service import PublishingOperationsService from app.social.services.publishing_service import PublishingService from app.social.services.scheduling_service import SchedulingService from app.social.services.social_service import SocialService from app.social.services.token_service import TokenService from app.templates.executor import OperationExecutor, TemplateExecutor from app.templates.loader import TemplateLoader +from app.templates.marketplace_repository import MarketplaceTemplateRepository +from app.templates.marketplace_service import MarketplaceTemplateService from app.templates.registry import TemplateRegistry from app.templates.validator import TemplateValidator +from app.analytics.repository import AnalyticsRepository +from app.projects.repositories.notification_repository import NotificationRepository +from app.projects.services.notification_service import NotificationService @dataclass(slots=True) @@ -55,26 +99,147 @@ class Container: processor: MediaProcessor template_registry: TemplateRegistry template_executor: TemplateExecutor + template_marketplace: MarketplaceTemplateService security_database: SecurityDatabase + tenants: TenantService + assets: CanonicalAssetService api_keys: APIKeyService rate_limiter: APIKeyRateLimiter audit: AuditService social: SocialService - + ai: AiStudioService + copilot: CopilotService + generation: GenerationService + generation_models: GenerationModelRegistry + generation_worker: GenerationWorker + projects: ProjectService + editor: ProjectEditorService + renders: ProjectRenderService + render_worker: ProjectRenderWorker + analytics: AnalyticsDomainService + analytics_worker: AnalyticsSyncWorker + brand_kits: BrandKitService + collaboration: CollaborationService + approval: ApprovalService + notifications: NotificationService def build_container(settings: Settings) -> Container: - security_database = SecurityDatabase(settings.database_url) - api_keys = APIKeyService(security_database, settings) + security_database = SecurityDatabase( + settings.database_url, auto_migrate=settings.security_auto_migrate + ) + tenants = TenantService(security_database) + assets = CanonicalAssetService(security_database) + api_keys = APIKeyService(security_database, settings, tenants) rate_limiter = APIKeyRateLimiter(security_database) audit = AuditService(security_database) - # Social publishing validates the same file-backed outputs as the normal - # media pipeline, so build those shared services before wiring Social. + brand_repository = BrandKitRepository(security_database) + brand_kits = BrandKitService(brand_repository, assets, audit) + projects = ProjectService(ProjectRepository(security_database), assets, audit, brand_kits) + collaboration = CollaborationService(CollaborationRepository(security_database)) + approval = ApprovalService(ApprovalRepository(security_database), ProjectEditorRepository(security_database)) + notifications = NotificationService(NotificationRepository(security_database)) cleanup = CleanupService(settings) + # Generation output validation shares the established FFprobe/validator + # services. Construct them before the output ingestor, without changing + # existing media or social provider behavior. validator = MediaValidator(settings) + ffprobe = FFprobeService(settings) + wan = WanProviderAdapter.from_settings(settings) + flux = FluxProviderAdapter.from_settings(settings) + generation_providers = GenerationProviderRegistry([wan, flux]) + generation_models = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id=WAN_PROVIDER_ID, + model=WAN_MODEL_CAPABILITY, + configuration_reference="wan-space", + metadata={ + "underlying_model_id": "Wan-AI/Wan2.2-I2V-A14B-Diffusers", + "task": "image-to-video", + "fps": 16, + }, + ), + GenerationModelRegistration( + provider_id=FLUX_PROVIDER_ID, + model=FLUX_MODEL_CAPABILITY, + configuration_reference="flux-space", + metadata={ + "task": FLUX_TASK, + "license": "Apache-2.0", + "models": { + "distilled": FLUX_DISTILLED_MODEL_ID, + "base": FLUX_BASE_MODEL_ID, + }, + }, + ), + ] + ) + generation = GenerationService( + settings=settings, + database=security_database, + assets=assets, + repository=GenerationRepository(security_database), + providers=generation_providers, + models=generation_models, + output_ingestor=GenerationOutputIngestor( + settings=settings, + cleanup=cleanup, + assets=assets, + ffprobe=ffprobe, + validator=validator, + ), + ) + generation_worker = GenerationWorker( + settings=settings, + generation=generation, + assets=assets, + cleanup=cleanup, + audit=audit, + ) + ai = AiStudioService(generation, projects, assets, audit) + # Social publishing validates the same file-backed outputs as the normal + # media pipeline, so build those shared services before wiring Social. downloader = Downloader(settings, validator) ytdlp = YTDLPService(settings, validator) ffmpeg = FFmpegService(settings) - ffprobe = FFprobeService(settings) + editor_repository = ProjectEditorRepository(security_database) + editor = ProjectEditorService( + editor_repository, + assets, + audit, + state_max_bytes=settings.editor_state_max_bytes, + max_tracks=settings.render_max_tracks, + max_clips=settings.render_max_clips, + max_duration_seconds=settings.render_max_duration_seconds, + ) + render_repository = ProjectRenderRepository(security_database) + renders = ProjectRenderService( + settings, render_repository, editor_repository, assets, cleanup, audit + ) + render_worker = ProjectRenderWorker( + settings, render_repository, renders, ffmpeg, cleanup, audit + ) + template_marketplace = MarketplaceTemplateService( + MarketplaceTemplateRepository(security_database), assets, ai, audit + ) + copilot_actions = CopilotActionRegistry( + projects=projects, + assets=assets, + editor=editor, + renders=renders, + ai=ai, + templates=template_marketplace, + ) + copilot = CopilotService( + repository=CopilotRepository(security_database), + planner=CopilotPlanner(), + actions=copilot_actions, + projects=projects, + assets=assets, + editor=editor, + ai=ai, + audit=audit, + ) whisper = WhisperService(settings) social_database = SocialDatabase(settings) providers = build_provider_registry(settings) @@ -101,26 +266,66 @@ def build_container(settings: Settings) -> Container: social_audit, ) social_media_assets = SocialMediaAssetService( - SocialMediaAssetRepository(social_database), cleanup, ffprobe, validator + settings, SocialMediaAssetRepository(social_database), assets, cleanup, ffprobe, validator + ) + publishing_service = PublishingService( + settings, + post_repository, + job_repository, + account_repository, + providers, + social_media_assets, + oauth_service, + brand_kits, + projects, + ) + scheduling_service = SchedulingService(post_repository, social_media_assets) + social_analytics = AnalyticsService( + social_database, account_repository, providers, oauth_service ) social = SocialService( settings=settings, database=social_database, accounts=account_service, oauth=oauth_service, - publishing=PublishingService( - settings, post_repository, job_repository, account_repository, providers, - social_media_assets, oauth_service, + publishing=publishing_service, + operations=PublishingOperationsService( + security_database=security_database, + posts=post_repository, + jobs=job_repository, + accounts=account_repository, + batches=PublishingBatchRepository(social_database), + publishing=publishing_service, + scheduling=scheduling_service, + audit=social_audit, ), - scheduling=SchedulingService(post_repository, social_media_assets), + scheduling=scheduling_service, jobs=JobService(job_repository), media_assets=social_media_assets, - analytics=AnalyticsService(social_database, account_repository, providers, oauth_service), + analytics=social_analytics, audit=social_audit, ) + analytics = AnalyticsDomainService( + AnalyticsRepository(social_database), + providers, + social_analytics, + account_repository, + social_audit, + projects, + ) + brand_repository = BrandKitRepository(security_database) + brand_kits = BrandKitService(brand_repository, assets, audit) + analytics_worker = AnalyticsSyncWorker( + analytics, + social_database, + interval_seconds=settings.social_scheduler_interval_seconds, + ) + copilot_actions.analytics = analytics + copilot_actions.publishing = social.publishing + copilot_actions.scheduling = social.scheduling resolver = InputResolver(settings, cleanup, downloader, ytdlp, validator) processor = MediaProcessor( - settings, resolver, cleanup, validator, ffmpeg, ffprobe, ytdlp, whisper + settings, resolver, cleanup, validator, ffmpeg, ffprobe, ytdlp, whisper, assets ) operation_executor = OperationExecutor(processor) template_validator = TemplateValidator(operation_executor.supported_operations) @@ -141,9 +346,27 @@ def build_container(settings: Settings) -> Container: processor=processor, template_registry=template_registry, template_executor=template_executor, + template_marketplace=template_marketplace, security_database=security_database, + tenants=tenants, + assets=assets, api_keys=api_keys, rate_limiter=rate_limiter, audit=audit, social=social, + ai=ai, + copilot=copilot, + generation=generation, + generation_models=generation.models, + generation_worker=generation_worker, + projects=projects, + editor=editor, + renders=renders, + render_worker=render_worker, + analytics=analytics, + analytics_worker=analytics_worker, + brand_kits=brand_kits, + collaboration=collaboration, + approval=approval, + notifications=notifications, ) diff --git a/app/copilot/__init__.py b/app/copilot/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a716a0d6b213f6d1d2d05e34c3a34c5bc3c0fc58 --- /dev/null +++ b/app/copilot/__init__.py @@ -0,0 +1 @@ +"""Validated orchestration over MediaRouter's existing product services.""" diff --git a/app/copilot/actions.py b/app/copilot/actions.py new file mode 100644 index 0000000000000000000000000000000000000000..bd5f2a7c5c77373a6d157d2daad3aff443c91d5e --- /dev/null +++ b/app/copilot/actions.py @@ -0,0 +1,853 @@ +from __future__ import annotations + +from dataclasses import dataclass +from uuid import uuid4 + +from app.ai.schemas import AiGenerateImageRequest, AiGenerateVideoRequest +from app.ai.service import AiStudioService +from app.analytics.schemas import AnalyticsQuery, AnalyticsSyncRequest +from app.analytics.service import AnalyticsDomainService +from app.copilot.errors import ( + CopilotCapabilityError, + CopilotInvalidRequestError, + CopilotPermissionError, +) +from app.copilot.schemas import ( + AiGenerateImageAction, + AiGenerateVideoAction, + AnalyticsOverviewAction, + AnalyticsSyncAction, + AssetSelectAction, + CopilotAction, + CopilotActionCapability, + CopilotActionResult, + EditorAddClipAction, + EditorDeleteClipAction, + EditorRenderAction, + EditorSetDurationAction, + EditorSplitClipAction, + ProjectOpenAction, + PublishingCancelAction, + PublishingCreatePostAction, + PublishingPublishAction, + PublishingScheduleAction, + PublishingValidateAction, + TemplateApplyAction, + TemplateCreateProjectAction, + TemplateGetAction, + TemplateSearchAction, +) +from app.projects.editor_schemas import ( + AudioClip, + ClipTransform, + EditorDocument, + EditorSaveRequest, + EffectClip, + MediaClip, + ProjectRenderCreate, + SourceClip, + Track, +) +from app.projects.services.editor_service import ProjectEditorService +from app.projects.services.project_service import ProjectService +from app.projects.services.render_service import ProjectRenderService +from app.security.assets import CanonicalAssetService +from app.social.services.publishing_service import PublishingService +from app.social.services.scheduling_service import SchedulingService +from app.templates.marketplace_schemas import TemplateApply, TemplateInstantiate +from app.templates.marketplace_service import MarketplaceTemplateService + + +@dataclass(frozen=True, slots=True) +class CopilotActionDefinition: + type: str + description: str + required_permission: str + required_capability: str + destructive: bool + external_side_effect: bool + requires_confirmation: bool + audit_event: str + + +ACTION_DEFINITIONS = ( + CopilotActionDefinition( + "project.open", + "Open a project.", + "projects:read", + "project.open", + False, + False, + False, + "copilot.project_opened", + ), + CopilotActionDefinition( + "asset.select", + "Open a canonical asset.", + "assets:read", + "asset.select", + False, + False, + False, + "copilot.asset_selected", + ), + CopilotActionDefinition( + "ai.generate_image", + "Submit image generation.", + "ai:generate", + "ai.generate_image", + False, + False, + True, + "copilot.ai_generation_requested", + ), + CopilotActionDefinition( + "ai.generate_video", + "Submit video generation.", + "ai:generate", + "ai.generate_video", + False, + False, + True, + "copilot.ai_generation_requested", + ), + CopilotActionDefinition( + "editor.split_clip", + "Split a timeline clip.", + "projects:update", + "editor.split_clip", + False, + False, + False, + "copilot.editor_updated", + ), + CopilotActionDefinition( + "editor.delete_clip", + "Delete a timeline clip.", + "projects:update", + "editor.delete_clip", + True, + False, + True, + "copilot.editor_updated", + ), + CopilotActionDefinition( + "editor.set_duration", + "Set a clip duration.", + "projects:update", + "editor.set_duration", + False, + False, + False, + "copilot.editor_updated", + ), + CopilotActionDefinition( + "editor.add_clip", + "Add an asset to the timeline.", + "projects:update", + "editor.add_clip", + False, + False, + False, + "copilot.editor_updated", + ), + CopilotActionDefinition( + "editor.render", + "Submit a project render.", + "projects:update", + "editor.render", + False, + False, + True, + "copilot.render_requested", + ), + CopilotActionDefinition( + "template.search", + "Search visible marketplace templates.", + "templates:read", + "template.search", + False, + False, + False, + "copilot.template_searched", + ), + CopilotActionDefinition( + "template.get", + "Inspect a visible marketplace template.", + "templates:read", + "template.get", + False, + False, + False, + "copilot.template_opened", + ), + CopilotActionDefinition( + "template.apply", + "Apply a template to an existing project.", + "templates:apply", + "template.apply", + True, + False, + True, + "copilot.template_applied", + ), + CopilotActionDefinition( + "template.create_project", + "Create a project from a template.", + "templates:apply", + "template.create_project", + False, + False, + True, + "copilot.template_project_created", + ), + CopilotActionDefinition( + "publishing.validate", + "Validate every social publishing target.", + "social:posts:write", + "publishing.validate", + False, + False, + False, + "copilot.publishing_validated", + ), + CopilotActionDefinition( + "publishing.create_post", + "Create a typed canonical social post draft.", + "social:posts:write", + "publishing.create_post", + False, + False, + False, + "copilot.publishing_draft_created", + ), + CopilotActionDefinition( + "publishing.schedule", + "Schedule an existing social post.", + "social:schedules:write", + "publishing.schedule", + False, + True, + True, + "copilot.publishing_scheduled", + ), + CopilotActionDefinition( + "publishing.publish", + "Publish an existing social post to its selected accounts.", + "social:posts:publish", + "publishing.publish", + False, + True, + True, + "copilot.publishing_started", + ), + CopilotActionDefinition( + "publishing.cancel", + "Cancel eligible publishing targets or request in-flight cancellation.", + "social:posts:write", + "publishing.cancel", + False, + True, + True, + "copilot.publishing_cancelled", + ), + CopilotActionDefinition( + "analytics.overview", + "Read authoritative analytics insights.", + "analytics:read", + "analytics.overview", + False, + False, + False, + "copilot.analytics_viewed", + ), + CopilotActionDefinition( + "analytics.sync", + "Queue authoritative analytics synchronization.", + "analytics:sync", + "analytics.sync", + False, + False, + True, + "copilot.analytics_sync_requested", + ), +) + + +class CopilotActionRegistry: + def __init__( + self, + *, + projects: ProjectService, + assets: CanonicalAssetService, + editor: ProjectEditorService, + renders: ProjectRenderService, + ai: AiStudioService, + templates: MarketplaceTemplateService, + publishing: PublishingService | None = None, + scheduling: SchedulingService | None = None, + analytics: AnalyticsDomainService | None = None, + ) -> None: + self.projects = projects + self.assets = assets + self.editor = editor + self.renders = renders + self.ai = ai + self.templates = templates + self.publishing = publishing + self.scheduling = scheduling + self.analytics = analytics + self.definitions = {item.type: item for item in ACTION_DEFINITIONS} + + def validate(self, action: CopilotAction) -> CopilotActionDefinition: + definition = self.definitions.get(action.type) + if definition is None: + raise CopilotInvalidRequestError("Copilot action type is not registered.") + expected = { + "required_permission": definition.required_permission, + "required_capability": definition.required_capability, + "destructive": definition.destructive, + "external_side_effect": definition.external_side_effect, + "requires_confirmation": definition.requires_confirmation, + } + if any(getattr(action, field) != value for field, value in expected.items()): + raise CopilotInvalidRequestError( + "Copilot action policy metadata does not match the registered action." + ) + return definition + + def validate_plan(self, actions: list[CopilotAction]) -> None: + for action in actions: + self.validate(action) + + def capabilities(self, available: set[str]) -> list[CopilotActionCapability]: + return [ + CopilotActionCapability( + type=item.type, + description=item.description, + required_permission=item.required_permission, + required_capability=item.required_capability, + destructive=item.destructive, + external_side_effect=item.external_side_effect, + requires_confirmation=item.requires_confirmation, + available=item.required_capability in available, + ) + for item in ACTION_DEFINITIONS + ] + + async def execute( + self, + action: CopilotAction, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + run_id: str, + permissions: frozenset[str], + available_capabilities: set[str], + ) -> CopilotActionResult: + definition = self.validate(action) + if definition.required_permission not in permissions and "admin" not in permissions: + raise CopilotPermissionError( + f"Permission '{definition.required_permission}' is required." + ) + if definition.required_capability not in available_capabilities: + raise CopilotCapabilityError( + f"Capability '{definition.required_capability}' is unavailable." + ) + if isinstance(action, EditorRenderAction) and not ( + "jobs:create" in permissions or "admin" in permissions + ): + raise CopilotPermissionError("Permission 'jobs:create' is required.") + if isinstance(action, TemplateApplyAction) and not ( + "projects:update" in permissions or "admin" in permissions + ): + raise CopilotPermissionError("Permission 'projects:update' is required.") + if isinstance(action, TemplateCreateProjectAction) and not ( + "projects:create" in permissions or "admin" in permissions + ): + raise CopilotPermissionError("Permission 'projects:create' is required.") + if ( + isinstance(action, (TemplateApplyAction, TemplateCreateProjectAction)) + and any( + binding.asset_id is not None for binding in action.arguments.slot_bindings.values() + ) + and not ("assets:read" in permissions or "admin" in permissions) + ): + raise CopilotPermissionError("Permission 'assets:read' is required.") + if isinstance( + action, + (AiGenerateVideoAction, EditorAddClipAction), + ) and not ("assets:read" in permissions or "admin" in permissions): + raise CopilotPermissionError("Permission 'assets:read' is required.") + if ( + isinstance(action, AiGenerateImageAction) + and action.arguments.source_asset_id is not None + and not ("assets:read" in permissions or "admin" in permissions) + ): + raise CopilotPermissionError("Permission 'assets:read' is required.") + if isinstance(action, ProjectOpenAction): + project = await self.projects.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(action.arguments.project_id), + ) + return self._success(action, "Project is ready to open.", "project", project.id) + if isinstance(action, AssetSelectAction): + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=str(action.arguments.asset_id), + ) + if action.arguments.project_id is not None and asset.project_id != str( + action.arguments.project_id + ): + raise CopilotInvalidRequestError( + "The selected asset does not belong to the selected project." + ) + return self._success(action, "Asset is ready to open.", "asset", asset.id) + if isinstance( + action, + ( + PublishingValidateAction, + PublishingCreatePostAction, + PublishingScheduleAction, + PublishingPublishAction, + PublishingCancelAction, + ), + ): + if self.publishing is None or self.scheduling is None: + raise CopilotCapabilityError("Publishing orchestration is unavailable.") + if isinstance(action, PublishingValidateAction): + result = await self.publishing.validate_post_targets( + workspace_id, str(action.arguments.post_id) + ) + return self._success( + action, + ( + "Publishing targets are valid." + if result.valid + else "Publishing validation found issues." + ), + "social_post", + result.post_id, + ) + if isinstance(action, PublishingCreatePostAction): + result = await self.publishing.create( + workspace_id=workspace_id, + user_id=user_id, + payload=action.arguments.post, + idempotency_key=f"copilot:{run_id}:{action.id}", + ) + return self._success( + action, "Publishing draft was created.", "social_post", result.id + ) + if isinstance(action, PublishingScheduleAction): + validation = await self.publishing.validate_post_targets( + workspace_id, str(action.arguments.post_id) + ) + if not validation.valid: + raise CopilotInvalidRequestError( + "Publishing validation must pass before scheduling." + ) + result = await self.scheduling.schedule( + workspace_id, str(action.arguments.post_id), action.arguments.schedule + ) + return self._success( + action, "Publishing was scheduled.", "social_schedule", result.id + ) + if isinstance(action, PublishingPublishAction): + jobs = await self.publishing.queue( + workspace_id, + str(action.arguments.post_id), + idempotency_key=f"copilot:{run_id}:{action.id}", + ) + return self._success( + action, + f"Queued {len(jobs)} publishing targets.", + "social_post", + str(action.arguments.post_id), + ) + result = await self.publishing.cancel(workspace_id, str(action.arguments.post_id)) + return self._success( + action, + "Cancellation was applied to eligible publishing targets.", + "social_post", + result.id, + ) + if isinstance(action, (AnalyticsOverviewAction, AnalyticsSyncAction)): + if self.analytics is None: + raise CopilotCapabilityError("Analytics orchestration is unavailable.") + if isinstance(action, AnalyticsOverviewAction): + result = await self.analytics.overview( + workspace_id, + AnalyticsQuery( + project_id=( + str(action.arguments.project_id) + if action.arguments.project_id + else None + ), + provider=action.arguments.provider, + metric=action.arguments.metric, + timezone=action.arguments.timezone, + sort=action.arguments.metric, + ), + ) + freshness = result.freshness.status + return self._success( + action, + f"Analytics overview is {freshness}; unavailable metrics were not inferred.", + "analytics_overview", + ( + str(action.arguments.project_id) + if action.arguments.project_id + else workspace_id + ), + ) + result = await self.analytics.create_sync( + workspace_id, + user_id, + AnalyticsSyncRequest( + project_id=( + str(action.arguments.project_id) if action.arguments.project_id else None + ), + provider=action.arguments.provider, + timezone=action.arguments.timezone, + idempotency_key=f"copilot:{run_id}:{action.id}", + ), + ) + return self._success( + action, "Analytics synchronization was queued.", "analytics_sync", result.id + ) + if isinstance(action, TemplateSearchAction): + result = await self.templates.list( + workspace_id=workspace_id, + user_id=user_id, + search=action.arguments.query, + category=action.arguments.category, + aspect_ratio=None, + min_duration_ms=None, + max_duration_ms=None, + media_type=None, + visibility=None, + status=None, + capability=None, + available_only=False, + offset=0, + limit=24, + ) + return self._success( + action, + f"Found {result.total} visible templates.", + "template_search", + action.arguments.query, + ) + if isinstance(action, TemplateGetAction): + template = await self.templates.get( + workspace_id=workspace_id, + user_id=user_id, + template_id=str(action.arguments.template_id), + ) + return self._success(action, "Template is ready to open.", "template", template.id) + if isinstance(action, TemplateApplyAction): + result = await self.templates.apply( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + template_id=str(action.arguments.template_id), + payload=TemplateApply( + project_id=action.arguments.project_id, + template_version_id=action.arguments.template_version_id, + slot_bindings=action.arguments.slot_bindings, + ), + idempotency_key=f"copilot:{run_id}:{action.id}", + instantiate=False, + ) + return self._success( + action, "Template was applied to the project.", "project", result.project_id + ) + if isinstance(action, TemplateCreateProjectAction): + result = await self.templates.apply( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + template_id=str(action.arguments.template_id), + payload=TemplateInstantiate( + project_name=action.arguments.project_name, + template_version_id=action.arguments.template_version_id, + slot_bindings=action.arguments.slot_bindings, + ), + idempotency_key=f"copilot:{run_id}:{action.id}", + instantiate=True, + ) + return self._success( + action, "Project was created from the template.", "project", result.project_id + ) + if isinstance(action, AiGenerateImageAction): + job = await self.ai.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + idempotency_key=f"copilot:{run_id}:{action.id}", + payload=AiGenerateImageRequest( + operation="generate_image", + prompt=action.arguments.prompt, + model=action.arguments.model, + project_id=action.arguments.project_id, + source_asset_ids=( + [action.arguments.source_asset_id] + if action.arguments.source_asset_id + else [] + ), + ), + ) + return self._success( + action, "Image generation was submitted.", "ai_generation", job.generation_id + ) + if isinstance(action, AiGenerateVideoAction): + job = await self.ai.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + idempotency_key=f"copilot:{run_id}:{action.id}", + payload=AiGenerateVideoRequest( + operation="generate_video", + prompt=action.arguments.prompt, + model=action.arguments.model, + project_id=action.arguments.project_id, + source_asset_ids=[action.arguments.source_asset_id], + ), + ) + return self._success( + action, "Video generation was submitted.", "ai_generation", job.generation_id + ) + if isinstance(action, EditorRenderAction): + render = await self.renders.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(action.arguments.project_id), + idempotency_key=f"copilot:{run_id}:{action.id}", + payload=ProjectRenderCreate(editor_revision=action.arguments.expected_revision), + ) + return self._success(action, "Render was submitted.", "project_render", render.id) + if isinstance( + action, + ( + EditorSplitClipAction, + EditorDeleteClipAction, + EditorSetDurationAction, + EditorAddClipAction, + ), + ): + revision = await self._execute_editor_action( + action, + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + ) + return self._success( + action, + f"Editor revision {revision} was saved.", + "project_editor_revision", + str(revision), + ) + raise CopilotInvalidRequestError("Copilot action type is not registered.") + + async def _execute_editor_action( + self, + action, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + ) -> int: + project_id = str(action.arguments.project_id) + current = await self.editor.get( + workspace_id=workspace_id, user_id=user_id, project_id=project_id + ) + if current.revision != action.arguments.expected_revision: + raise CopilotInvalidRequestError( + "The editor changed after this plan was created. Create a new plan." + ) + document = current.state.model_copy(deep=True) + if isinstance(action, EditorSplitClipAction): + self._split(document, action.arguments.clip_id, action.arguments.at_ms) + elif isinstance(action, EditorDeleteClipAction): + self._delete(document, action.arguments.clip_id) + elif isinstance(action, EditorSetDurationAction): + self._set_duration(document, action.arguments.clip_id, action.arguments.duration_ms) + elif isinstance(action, EditorAddClipAction): + await self._add_clip( + document, + workspace_id=workspace_id, + user_id=user_id, + project_id=project_id, + asset_id=str(action.arguments.asset_id), + duration_ms=action.arguments.duration_ms, + ) + validated = EditorDocument.model_validate(document.model_dump(by_alias=True)) + saved = await self.editor.save( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=project_id, + payload=EditorSaveRequest( + expected_revision=current.revision, + schema_version=validated.schema_version, + state=validated, + ), + ) + return saved.revision + + @staticmethod + def _find_clip(document: EditorDocument, clip_id: str): + for track in document.timeline.tracks: + for index, clip in enumerate(track.clips): + if clip.id == clip_id: + return track, index, clip + raise CopilotInvalidRequestError("The selected clip no longer exists.") + + def _split(self, document: EditorDocument, clip_id: str, at_ms: int) -> None: + track, index, clip = self._find_clip(document, clip_id) + relative = at_ms - clip.start_ms + if relative <= 0 or relative >= clip.duration_ms: + raise CopilotInvalidRequestError("The split point must fall inside the selected clip.") + right = clip.model_copy(deep=True) + right.id = str(uuid4()) + right.start_ms = at_ms + right.duration_ms = clip.duration_ms - relative + clip.duration_ms = relative + if isinstance(clip, SourceClip) and isinstance(right, SourceClip): + right.source_start_ms = clip.source_start_ms + relative + right.source_duration_ms = right.duration_ms + clip.source_duration_ms = clip.duration_ms + track.clips.insert(index + 1, right) + + def _delete(self, document: EditorDocument, clip_id: str) -> None: + track, index, _ = self._find_clip(document, clip_id) + track.clips.pop(index) + document.timeline.transitions = [ + item + for item in document.timeline.transitions + if item.from_clip_id != clip_id and item.to_clip_id != clip_id + ] + for candidate in document.timeline.tracks: + candidate.clips = [ + clip + for clip in candidate.clips + if not (isinstance(clip, EffectClip) and clip.target_clip_id == clip_id) + ] + + def _set_duration(self, document: EditorDocument, clip_id: str, duration_ms: int) -> None: + _, _, clip = self._find_clip(document, clip_id) + if isinstance(clip, SourceClip) and not ( + isinstance(clip, MediaClip) and clip.media_type == "image" + ): + raise CopilotInvalidRequestError( + "Only image or non-source clips can be extended without media analysis." + ) + clip.duration_ms = duration_ms + if isinstance(clip, SourceClip): + clip.source_duration_ms = duration_ms + + async def _add_clip( + self, + document: EditorDocument, + *, + workspace_id: str, + user_id: str, + project_id: str, + asset_id: str, + duration_ms: int, + ) -> None: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, user_id=user_id, asset_id=asset_id + ) + if asset.project_id != project_id: + raise CopilotInvalidRequestError("The selected asset is not attached to this project.") + mime = asset.mime_type + if mime.startswith("audio/"): + track_type = "audio" + kind = "audio" + elif mime.startswith("video/"): + track_type = "video" + kind = "video" + elif mime.startswith("image/"): + track_type = "video" + kind = "image" + else: + raise CopilotInvalidRequestError("This asset type cannot be added to the timeline.") + track = next( + (item for item in document.timeline.tracks if item.type == track_type), + None, + ) + if track is None: + track = Track( + id=str(uuid4()), + type=track_type, + name="Copilot media", + order=len(document.timeline.tracks), + muted=False, + locked=False, + visible=True, + clips=[], + ) + document.timeline.tracks.append(track) + metadata = asset.metadata_json or {} + source_duration = metadata.get("duration_ms") + if source_duration is None and isinstance(metadata.get("duration"), (int, float)): + source_duration = round(float(metadata["duration"]) * 1_000) + clip_duration = duration_ms if kind == "image" else source_duration + if not isinstance(clip_duration, int) or clip_duration <= 0: + raise CopilotInvalidRequestError("The asset has no validated duration metadata.") + common = { + "id": str(uuid4()), + "trackId": track.id, + "label": asset.filename[:500], + "startMs": document.duration_ms(), + "durationMs": clip_duration, + "visible": True, + "opacity": 1, + "metadata": {}, + "assetId": asset.id, + "sourceStartMs": 0, + "sourceDurationMs": clip_duration, + } + if kind == "audio": + track.clips.append(AudioClip(**common, kind="audio", volume=1, fadeInMs=0, fadeOutMs=0)) + else: + track.clips.append( + MediaClip( + **common, + kind="media", + mediaType=kind, + transform=ClipTransform(x=0, y=0, scaleX=1, scaleY=1, rotation=0), + volume=1, + ) + ) + + @staticmethod + def _success( + action: CopilotAction, + summary: str, + resource_type: str, + resource_id: str, + ) -> CopilotActionResult: + return CopilotActionResult( + action_id=action.id, + action_type=action.type, + status="completed", + summary=summary, + resource_type=resource_type, + resource_id=resource_id, + ) diff --git a/app/copilot/api.py b/app/copilot/api.py new file mode 100644 index 0000000000000000000000000000000000000000..92839eb97baddc1d9ea7d07d6a811dbeec7c5223 --- /dev/null +++ b/app/copilot/api.py @@ -0,0 +1,121 @@ +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Header, Path, Query, Request, status + +from app.copilot.schemas import ( + CopilotCapabilities, + CopilotContextInput, + CopilotExecuteRequest, + CopilotRun, + CopilotRunCreate, + CopilotRunList, +) +from app.security.errors import ForbiddenError + +router = APIRouter(prefix="/v1/copilot", tags=["copilot"]) + + +def _identity(request: Request): + context = request.state.auth + if not context.workspace_id or not context.user_id: + raise ForbiddenError + return context + + +@router.post("/capabilities", response_model=CopilotCapabilities) +async def capabilities(request: Request, context: CopilotContextInput) -> CopilotCapabilities: + identity = _identity(request) + return await request.app.state.container.copilot.capabilities( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + context=context, + ) + + +@router.post("/runs", response_model=CopilotRun, status_code=status.HTTP_201_CREATED) +async def create_run( + request: Request, + payload: CopilotRunCreate, + idempotency_key: Annotated[str, Header(alias="Idempotency-Key", min_length=8, max_length=255)], +) -> CopilotRun: + identity = _identity(request) + return await request.app.state.container.copilot.create_run( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + payload=payload, + idempotency_key=idempotency_key, + ) + + +@router.get("/runs", response_model=CopilotRunList) +async def list_runs( + request: Request, + offset: Annotated[int, Query(ge=0)] = 0, + limit: Annotated[int, Query(ge=1, le=100)] = 25, +) -> CopilotRunList: + identity = _identity(request) + return await request.app.state.container.copilot.list( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + offset=offset, + limit=limit, + ) + + +@router.get("/history", response_model=CopilotRunList) +async def history( + request: Request, + offset: Annotated[int, Query(ge=0)] = 0, + limit: Annotated[int, Query(ge=1, le=100)] = 25, +) -> CopilotRunList: + return await list_runs(request, offset, limit) + + +@router.get("/runs/{run_id}", response_model=CopilotRun) +async def get_run( + request: Request, + run_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> CopilotRun: + identity = _identity(request) + return await request.app.state.container.copilot.get( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + run_id=run_id, + ) + + +@router.post("/runs/{run_id}/execute", response_model=CopilotRun) +async def execute_run( + request: Request, + payload: CopilotExecuteRequest, + run_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> CopilotRun: + identity = _identity(request) + return await request.app.state.container.copilot.execute( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + run_id=run_id, + payload=payload, + permissions=identity.scopes, + ) + + +@router.post("/runs/{run_id}/cancel", response_model=CopilotRun) +async def cancel_run( + request: Request, + run_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> CopilotRun: + identity = _identity(request) + return await request.app.state.container.copilot.cancel( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + run_id=run_id, + ) diff --git a/app/copilot/errors.py b/app/copilot/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..b36fd4667182e81d0f56ac0f48d7d0a9c47bfe90 --- /dev/null +++ b/app/copilot/errors.py @@ -0,0 +1,36 @@ +from app.core.exceptions import MediaAPIError + + +class CopilotInvalidRequestError(MediaAPIError): + code = "COPILOT_INVALID_REQUEST" + status_code = 422 + + +class CopilotUnsupportedIntentError(MediaAPIError): + code = "COPILOT_UNSUPPORTED_INTENT" + status_code = 422 + + +class CopilotRunNotFoundError(MediaAPIError): + code = "COPILOT_RUN_NOT_FOUND" + status_code = 404 + + +class CopilotRunConflictError(MediaAPIError): + code = "COPILOT_RUN_CONFLICT" + status_code = 409 + + +class CopilotConfirmationRequiredError(MediaAPIError): + code = "COPILOT_CONFIRMATION_REQUIRED" + status_code = 409 + + +class CopilotPermissionError(MediaAPIError): + code = "COPILOT_ACTION_FORBIDDEN" + status_code = 403 + + +class CopilotCapabilityError(MediaAPIError): + code = "COPILOT_CAPABILITY_UNAVAILABLE" + status_code = 422 diff --git a/app/copilot/models.py b/app/copilot/models.py new file mode 100644 index 0000000000000000000000000000000000000000..cfb23e8d081151ad6bbe49f4b447a15cdbbdc735 --- /dev/null +++ b/app/copilot/models.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import ( + JSON, + CheckConstraint, + DateTime, + ForeignKey, + Index, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base, utcnow + + +class CopilotRunRecord(Base): + __tablename__ = "copilot_runs" + __table_args__ = ( + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_copilot_run_workspace_idempotency" + ), + CheckConstraint( + "status in ('plan_ready','blocked','executing','completed','partial','failed','cancelled')", + name="ck_copilot_run_status", + ), + Index("ix_copilot_runs_workspace_created", "workspace_id", "created_at"), + Index("ix_copilot_runs_workspace_status", "workspace_id", "status"), + Index("ix_copilot_runs_project_created", "project_id", "created_at"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + project_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT") + ) + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + request_fingerprint: Mapped[str] = mapped_column(String(64), nullable=False) + request_text: Mapped[str] = mapped_column(Text, nullable=False) + context_json: Mapped[dict[str, object]] = mapped_column("context", JSON, nullable=False) + plan_json: Mapped[dict[str, object]] = mapped_column("plan", JSON, nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False) + current_action_id: Mapped[str | None] = mapped_column(String(64)) + results_json: Mapped[list[dict[str, object]]] = mapped_column( + "results", JSON, nullable=False, default=list + ) + summary: Mapped[str | None] = mapped_column(Text) + error_code: Mapped[str | None] = mapped_column(String(100)) + error_message: Mapped[str | None] = mapped_column(Text) + confirmed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) diff --git a/app/copilot/planner.py b/app/copilot/planner.py new file mode 100644 index 0000000000000000000000000000000000000000..68d3cdf0fe4e8b4710321807fedaf1d02ef131ef --- /dev/null +++ b/app/copilot/planner.py @@ -0,0 +1,627 @@ +from __future__ import annotations + +import re +from uuid import uuid4 + +from app.copilot.schemas import ( + AiGenerateImageAction, + AiGenerateImageArguments, + AiGenerateVideoAction, + AiGenerateVideoArguments, + AnalyticsOverviewAction, + AnalyticsOverviewArguments, + AnalyticsSyncAction, + AnalyticsSyncArguments, + AssetSelectAction, + AssetSelectArguments, + ShareProjectAction, + ShareProjectArguments, + CopilotContext, + CopilotPlan, + EditorAddClipAction, + EditorAddClipArguments, + EditorDeleteClipAction, + EditorDeleteClipArguments, + EditorRenderAction, + EditorRenderArguments, + EditorSetDurationAction, + EditorSetDurationArguments, + EditorSplitClipAction, + EditorSplitClipArguments, + ProjectOpenAction, + ProjectOpenArguments, + PublishingCancelAction, + PublishingCancelArguments, + PublishingPublishAction, + PublishingPublishArguments, + PublishingValidateAction, + PublishingValidateArguments, + TemplateApplyAction, + TemplateApplyArguments, + TemplateCreateProjectAction, + TemplateCreateProjectArguments, + TemplateGetAction, + TemplateGetArguments, + TemplateSearchAction, + TemplateSearchArguments, +) + + +class CopilotPlanner: + """Closed deterministic planner used until a text tool-calling model exists.""" + + def plan(self, request: str, context: CopilotContext) -> CopilotPlan: + normalized = " ".join(request.strip().split()) + lowered = normalized.casefold() + action_id = str(uuid4()) + project_id = context.project_id + selected_asset = context.selected_asset_ids[0] if context.selected_asset_ids else None + selected_clip = context.selected_clip_ids[0] if context.selected_clip_ids else None + revision = context.editor_summary.revision if context.editor_summary else None + + unavailable = ( + None + if "template" in lowered + else self._requested_unavailable_capability(lowered, context) + ) + if unavailable: + return self._blocked( + normalized, + f"This request requires {unavailable}, but that capability is not available.", + unsupported=[unavailable], + ) + + template_id = self._template_id(lowered) + publishing_post_id = self._template_id(lowered) + if "analytics" in lowered and re.search(r"\b(sync|refresh|update)\b", lowered): + action = AnalyticsSyncAction( + id=action_id, + type="analytics.sync", + arguments=AnalyticsSyncArguments(project_id=project_id), + reason="Queue a durable synchronization through authorized provider adapters.", + requires_confirmation=True, + destructive=False, + external_side_effect=False, + required_permission="analytics:sync", + required_capability="analytics.sync", + ) + return self._plan( + "Sync analytics", + "Queue an idempotent authoritative analytics synchronization.", + [action], + ) + if re.search(r"\b(analytics|performance|performing|insights)\b", lowered): + metric = next( + ( + item + for item in ( + "views", + "impressions", + "likes", + "comments", + "shares", + "engagement_rate", + ) + if item.replace("_", " ") in lowered + ), + "views", + ) + action = AnalyticsOverviewAction( + id=action_id, + type="analytics.overview", + arguments=AnalyticsOverviewArguments(project_id=project_id, metric=metric), + reason="Read synchronized provider metrics without inferring unavailable values.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="analytics:read", + required_capability="analytics.overview", + ) + return self._plan( + "Review analytics", + "Read authoritative analytics and report freshness explicitly.", + [action], + ) + if publishing_post_id and re.search(r"\b(validate|check)\b.*\b(publish|post)\b", lowered): + action = PublishingValidateAction( + id=action_id, + type="publishing.validate", + arguments=PublishingValidateArguments(post_id=publishing_post_id), + reason="Validate the existing canonical post and every selected provider target.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="social:posts:write", + required_capability="publishing.validate", + ) + return self._plan( + "Validate publishing", + "Run authoritative per-target publishing validation.", + [action], + ) + if publishing_post_id and re.search(r"\b(publish|send)\b", lowered): + action = PublishingPublishAction( + id=action_id, + type="publishing.publish", + arguments=PublishingPublishArguments(post_id=publishing_post_id), + reason="Publish the explicitly identified canonical post to its already selected accounts.", + requires_confirmation=True, + destructive=False, + external_side_effect=True, + required_permission="social:posts:publish", + required_capability="publishing.publish", + ) + return self._plan( + "Publish social post", + "Validate and queue external publishing only after explicit confirmation.", + [action], + ) + if publishing_post_id and re.search(r"\bcancel\b.*\b(publish|post)\b", lowered): + action = PublishingCancelAction( + id=action_id, + type="publishing.cancel", + arguments=PublishingCancelArguments(post_id=publishing_post_id), + reason="Cancel eligible jobs and request cancellation for in-flight provider work.", + requires_confirmation=True, + destructive=False, + external_side_effect=True, + required_permission="social:posts:write", + required_capability="publishing.cancel", + ) + return self._plan( + "Cancel publishing", + "Apply truthful cancellation semantics after confirmation.", + [action], + ) + if "template" in lowered and re.search(r"\b(find|search|browse)\b", lowered): + query = re.sub( + r"(?i)\b(find|search|browse|for|me|a|an|template|templates)\b", " ", normalized + ) + query = " ".join(query.split()) or normalized + category = next( + ( + item + for item in ( + "business", + "marketing", + "education", + "podcast", + "gaming", + "news", + "social", + "youtube", + "tiktok", + "instagram", + "product", + "personal", + ) + if item in lowered + ), + None, + ) + action = TemplateSearchAction( + id=action_id, + type="template.search", + arguments=TemplateSearchArguments(query=query, category=category), + reason="Search the authoritative visible template catalog.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="templates:read", + required_capability="template.search", + ) + return self._plan( + "Search templates", + "Search the versioned marketplace catalog using the current workspace context.", + [action], + ) + + if "share" in lowered and "project" in lowered and re.search(r"\b(with)\b", lowered): + # Example parsing for demo purposes + project_id = "..." # simplified + user_id = "..." # simplified + role = "viewer" + action = ShareProjectAction( + id=action_id, + type="project.share", + arguments=ShareProjectArguments(project_id=project_id, user_id=user_id, role=role), + reason="Share project with user.", + requires_confirmation=True, + destructive=False, + external_side_effect=True, + required_permission="projects:share", + required_capability="project.share", + ) + return self._plan("Share project", "Share project with another user.", [action]) + + if template_id and "template" in lowered and re.search(r"\b(open|show|inspect)\b", lowered): + action = TemplateGetAction( + id=action_id, + type="template.get", + arguments=TemplateGetArguments(template_id=template_id), + reason="Inspect the selected authoritative template.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="templates:read", + required_capability="template.get", + ) + return self._plan("Open template", "Open the selected template.", [action]) + + if template_id and "template" in lowered and re.search(r"\b(use|apply)\b", lowered): + if project_id is None: + return self._blocked( + normalized, + "Applying a template requires a selected project.", + missing=["project"], + ) + action = TemplateApplyAction( + id=action_id, + type="template.apply", + arguments=TemplateApplyArguments( + template_id=template_id, + project_id=project_id, + slot_bindings={}, + ), + reason="Apply the selected versioned template to the current project.", + requires_confirmation=True, + destructive=True, + external_side_effect=False, + required_permission="templates:apply", + required_capability="template.apply", + ) + return self._plan( + "Apply template", + "Validate requirements and replace the current authoritative editor state after confirmation.", + [action], + ) + + if template_id and "template" in lowered and "create project" in lowered: + name_match = re.search(r'\bnamed\s+["“]?([^"”]+?)["”]?\s*$', normalized, re.IGNORECASE) + if name_match is None: + return self._blocked( + normalized, + "Creating a project from a template requires an explicit project name.", + missing=["project name"], + ) + action = TemplateCreateProjectAction( + id=action_id, + type="template.create_project", + arguments=TemplateCreateProjectArguments( + template_id=template_id, + project_name=name_match.group(1).strip(), + slot_bindings={}, + ), + reason="Create an editable project from the selected template.", + requires_confirmation=True, + destructive=False, + external_side_effect=False, + required_permission="templates:apply", + required_capability="template.create_project", + ) + return self._plan( + "Create project from template", + "Validate requirements and create a new editable project after confirmation.", + [action], + ) + + if re.search(r"\b(open|show)\b.*\bproject\b", lowered): + if project_id is None: + return self._blocked(normalized, "Select a project first.", missing=["project"]) + action = ProjectOpenAction( + id=action_id, + type="project.open", + arguments=ProjectOpenArguments(project_id=project_id), + reason="Open the selected project.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="projects:read", + required_capability="project.open", + ) + return self._plan("Open project", "Open the selected project.", [action]) + + if re.search(r"\b(open|select|show)\b.*\basset\b", lowered): + if selected_asset is None: + return self._blocked(normalized, "Select an asset first.", missing=["asset"]) + action = AssetSelectAction( + id=action_id, + type="asset.select", + arguments=AssetSelectArguments(asset_id=selected_asset, project_id=project_id), + reason="Open the selected canonical asset.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="assets:read", + required_capability="asset.select", + ) + return self._plan("Open asset", "Open the selected asset.", [action]) + + if "generate" in lowered and ("video" in lowered or "animate" in lowered): + if "ai.generate_video" not in context.available_capabilities: + return self._blocked( + normalized, + "Video generation is not currently available.", + unsupported=["ai.generate_video"], + ) + if selected_asset is None: + return self._blocked( + normalized, + "Video generation requires a selected source image.", + missing=["source image asset"], + ) + prompt = self._generation_prompt(normalized, "video") + action = AiGenerateVideoAction( + id=action_id, + type="ai.generate_video", + arguments=AiGenerateVideoArguments( + prompt=prompt, + project_id=project_id, + source_asset_id=selected_asset, + ), + reason="Submit a real image-to-video generation job.", + requires_confirmation=True, + destructive=False, + external_side_effect=False, + required_permission="ai:generate", + required_capability="ai.generate_video", + ) + return self._plan( + "Generate video", + "Generate a video from the selected image using AI Studio.", + [action], + ) + + if "generate" in lowered and any( + word in lowered for word in ("image", "thumbnail", "artwork") + ): + if "ai.generate_image" not in context.available_capabilities: + return self._blocked( + normalized, + "Image generation is not currently available.", + unsupported=["ai.generate_image"], + ) + prompt = self._generation_prompt(normalized, "image") + action = AiGenerateImageAction( + id=action_id, + type="ai.generate_image", + arguments=AiGenerateImageArguments( + prompt=prompt, + project_id=project_id, + source_asset_id=selected_asset, + ), + reason="Submit a real image generation job.", + requires_confirmation=True, + destructive=False, + external_side_effect=False, + required_permission="ai:generate", + required_capability="ai.generate_image", + ) + return self._plan( + "Generate image", + "Generate an image using the currently available AI Studio model.", + [action], + ) + + if re.search(r"\b(render|export)\b", lowered): + if project_id is None or revision is None: + return self._blocked( + normalized, + "Rendering requires a project with saved editor state.", + missing=["saved editor state"], + ) + action = EditorRenderAction( + id=action_id, + type="editor.render", + arguments=EditorRenderArguments(project_id=project_id, expected_revision=revision), + reason="Submit the current authoritative editor revision for rendering.", + requires_confirmation=True, + destructive=False, + external_side_effect=False, + required_permission="projects:update", + required_capability="editor.render", + ) + return self._plan( + "Render project", + "Validate and submit the current editor revision to the existing render pipeline.", + [action], + ) + + split_match = re.search( + r"\bsplit\b.*?\b(?:at\s+)?(\d+(?:\.\d+)?)\s*(seconds?|secs?|s)\b", + lowered, + ) + if split_match: + missing = self._editor_missing(project_id, revision, selected_clip) + if missing: + return self._blocked( + normalized, "Select a saved timeline clip first.", missing=missing + ) + action = EditorSplitClipAction( + id=action_id, + type="editor.split_clip", + arguments=EditorSplitClipArguments( + project_id=project_id, + clip_id=selected_clip, + at_ms=round(float(split_match.group(1)) * 1_000), + expected_revision=revision, + ), + reason="Split the selected clip at the requested timeline time.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="projects:update", + required_capability="editor.split_clip", + ) + return self._plan( + "Split selected clip", + "Save a revision-safe split through the authoritative editor service.", + [action], + ) + + if re.search(r"\b(delete|remove)\b.*\b(selected\s+)?clip\b", lowered): + missing = self._editor_missing(project_id, revision, selected_clip) + if missing: + return self._blocked( + normalized, "Select a saved timeline clip first.", missing=missing + ) + action = EditorDeleteClipAction( + id=action_id, + type="editor.delete_clip", + arguments=EditorDeleteClipArguments( + project_id=project_id, + clip_id=selected_clip, + expected_revision=revision, + ), + reason="Remove the selected clip from the authoritative timeline.", + requires_confirmation=True, + destructive=True, + external_side_effect=False, + required_permission="projects:update", + required_capability="editor.delete_clip", + ) + return self._plan( + "Delete selected clip", + "Delete the selected clip after explicit confirmation.", + [action], + ) + + duration_match = re.search( + r"\b(?:last|duration|make)\b.*?(\d+(?:\.\d+)?)\s*(seconds?|secs?|s)\b", + lowered, + ) + if duration_match and ("clip" in lowered or "image" in lowered): + missing = self._editor_missing(project_id, revision, selected_clip) + if missing: + return self._blocked( + normalized, "Select a saved timeline clip first.", missing=missing + ) + action = EditorSetDurationAction( + id=action_id, + type="editor.set_duration", + arguments=EditorSetDurationArguments( + project_id=project_id, + clip_id=selected_clip, + duration_ms=round(float(duration_match.group(1)) * 1_000), + expected_revision=revision, + ), + reason="Set the selected clip duration.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="projects:update", + required_capability="editor.set_duration", + ) + return self._plan( + "Update clip duration", + "Save the requested duration through the authoritative editor service.", + [action], + ) + + if re.search(r"\badd\b.*\b(asset|image|video|audio)\b.*\b(timeline|editor)\b", lowered): + if project_id is None or revision is None or selected_asset is None: + return self._blocked( + normalized, + "Adding media requires a project, saved editor state, and selected asset.", + missing=["project", "saved editor state", "asset"], + ) + action = EditorAddClipAction( + id=action_id, + type="editor.add_clip", + arguments=EditorAddClipArguments( + project_id=project_id, + asset_id=selected_asset, + expected_revision=revision, + ), + reason="Add the selected canonical asset to the timeline.", + requires_confirmation=False, + destructive=False, + external_side_effect=False, + required_permission="projects:update", + required_capability="editor.add_clip", + ) + return self._plan( + "Add asset to timeline", + "Insert the selected project asset through the authoritative editor service.", + [action], + ) + + return self._blocked( + normalized, + "This request does not map to a currently registered Copilot action.", + unsupported=["natural-language intent"], + ) + + @staticmethod + def _editor_missing(project_id, revision, selected_clip) -> list[str]: + missing = [] + if project_id is None: + missing.append("project") + if revision is None: + missing.append("saved editor state") + if selected_clip is None: + missing.append("selected clip") + return missing + + @staticmethod + def _generation_prompt(request: str, media_word: str) -> str: + stripped = re.sub( + rf"(?i)^\s*(please\s+)?generate\s+(an?\s+)?{media_word}\s*(of|for|with|:)?\s*", + "", + request, + ).strip() + return stripped or request + + @staticmethod + def _template_id(request: str) -> str | None: + match = re.search( + r"\b[0-9a-f]{8}-[0-9a-f]{4}-[1-5][0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}\b", + request, + re.IGNORECASE, + ) + return match.group(0) if match else None + + @staticmethod + def _requested_unavailable_capability(request: str, context: CopilotContext) -> str | None: + requested = { + "transcrib": "ai.transcribe", + "upscal": "ai.upscale", + "remove background": "ai.remove_background", + "voice": "ai.generate_voice", + "music": "ai.generate_music", + "caption": "ai.transcribe", + "tiktok": "ai.transcribe", + "highlight": "ai.analyze", + } + for fragment, capability in requested.items(): + if fragment in request and capability not in context.available_capabilities: + return capability + return None + + @staticmethod + def _plan(intent: str, explanation: str, actions: list) -> CopilotPlan: + return CopilotPlan( + intent=intent, + explanation=explanation, + actions=actions, + executable=True, + requires_confirmation=any(action.requires_confirmation for action in actions), + ) + + @staticmethod + def _blocked( + intent: str, + explanation: str, + *, + missing: list[str] | None = None, + unsupported: list[str] | None = None, + ) -> CopilotPlan: + return CopilotPlan( + intent=intent[:200], + explanation=explanation, + actions=[], + missing_information=missing or [], + unsupported_capabilities=unsupported or [], + executable=False, + requires_confirmation=False, + ) diff --git a/app/copilot/repository.py b/app/copilot/repository.py new file mode 100644 index 0000000000000000000000000000000000000000..c089960fc99f36a5e95ab07da674a4dcc7fbf43a --- /dev/null +++ b/app/copilot/repository.py @@ -0,0 +1,196 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError + +from app.copilot.errors import CopilotRunConflictError, CopilotRunNotFoundError +from app.copilot.models import CopilotRunRecord +from app.security.database import SecurityDatabase + + +class CopilotRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def create(self, record: CopilotRunRecord) -> tuple[CopilotRunRecord, bool]: + try: + async with self.database.tenant_session( + workspace_id=record.workspace_id, user_id=record.user_id + ) as session: + session.add(record) + await session.commit() + await session.refresh(record) + return record, True + except IntegrityError: + existing = await self.get_by_idempotency( + record.workspace_id, record.user_id, record.idempotency_key + ) + if existing is None: + raise + if existing.request_fingerprint != record.request_fingerprint: + raise CopilotRunConflictError( + "Idempotency-Key is already associated with another Copilot request." + ) + return existing, False + + async def get_by_idempotency( + self, workspace_id: str, user_id: str, key: str + ) -> CopilotRunRecord | None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return await session.scalar( + select(CopilotRunRecord).where( + CopilotRunRecord.workspace_id == workspace_id, + CopilotRunRecord.idempotency_key == key, + ) + ) + + async def get( + self, workspace_id: str, user_id: str, run_id: str, *, lock: bool = False + ) -> CopilotRunRecord: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + statement = select(CopilotRunRecord).where( + CopilotRunRecord.id == run_id, + CopilotRunRecord.workspace_id == workspace_id, + ) + if lock: + statement = statement.with_for_update() + record = await session.scalar(statement) + if record is None: + raise CopilotRunNotFoundError("Copilot run was not found in this workspace.") + return record + + async def list( + self, workspace_id: str, user_id: str, *, offset: int, limit: int + ) -> list[CopilotRunRecord]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return list( + ( + await session.scalars( + select(CopilotRunRecord) + .where(CopilotRunRecord.workspace_id == workspace_id) + .order_by(CopilotRunRecord.created_at.desc()) + .offset(offset) + .limit(limit) + ) + ).all() + ) + + async def claim_execution( + self, + workspace_id: str, + user_id: str, + run_id: str, + *, + confirmed: bool, + ) -> CopilotRunRecord: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + record = await session.scalar( + select(CopilotRunRecord) + .where( + CopilotRunRecord.id == run_id, + CopilotRunRecord.workspace_id == workspace_id, + ) + .with_for_update() + ) + if record is None: + raise CopilotRunNotFoundError("Copilot run was not found in this workspace.") + if record.status != "plan_ready": + raise CopilotRunConflictError("Only a plan-ready Copilot run can be executed.") + now = datetime.now(timezone.utc) + record.status = "executing" + record.updated_at = now + if confirmed and record.confirmed_at is None: + record.confirmed_at = now + await session.commit() + await session.refresh(record) + return record + + async def cancel_before_execution( + self, + workspace_id: str, + user_id: str, + run_id: str, + ) -> tuple[CopilotRunRecord, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + record = await session.scalar( + select(CopilotRunRecord) + .where( + CopilotRunRecord.id == run_id, + CopilotRunRecord.workspace_id == workspace_id, + ) + .with_for_update() + ) + if record is None: + raise CopilotRunNotFoundError("Copilot run was not found in this workspace.") + if record.status in {"completed", "partial", "failed", "cancelled"}: + return record, False + if record.status == "executing": + raise CopilotRunConflictError( + "This action batch is already executing; cancel its durable child job directly." + ) + now = datetime.now(timezone.utc) + record.status = "cancelled" + record.current_action_id = None + record.summary = "Copilot run cancelled before execution." + record.updated_at = now + record.completed_at = now + await session.commit() + await session.refresh(record) + return record, True + + async def update( + self, + workspace_id: str, + user_id: str, + run_id: str, + *, + status: str, + current_action_id: str | None = None, + results: list[dict[str, object]] | None = None, + summary: str | None = None, + error_code: str | None = None, + error_message: str | None = None, + confirmed: bool = False, + terminal: bool = False, + ) -> CopilotRunRecord: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + record = await session.scalar( + select(CopilotRunRecord) + .where( + CopilotRunRecord.id == run_id, + CopilotRunRecord.workspace_id == workspace_id, + ) + .with_for_update() + ) + if record is None: + raise CopilotRunNotFoundError("Copilot run was not found in this workspace.") + now = datetime.now(timezone.utc) + record.status = status + record.current_action_id = current_action_id + if results is not None: + record.results_json = results + record.summary = summary + record.error_code = error_code + record.error_message = error_message + if confirmed and record.confirmed_at is None: + record.confirmed_at = now + record.updated_at = now + if terminal: + record.completed_at = now + await session.commit() + await session.refresh(record) + return record diff --git a/app/copilot/schemas.py b/app/copilot/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..4b4e0b15c95bb9839ac19f00881a5e1fb44f34e6 --- /dev/null +++ b/app/copilot/schemas.py @@ -0,0 +1,439 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Annotated, Literal +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from app.social.schemas.posts import SocialPostCreate +from app.social.schemas.scheduling import SocialScheduleCreate +from app.templates.marketplace_schemas import SlotBinding + +CopilotRunStatus = Literal[ + "plan_ready", + "blocked", + "executing", + "completed", + "partial", + "failed", + "cancelled", +] +CopilotActionStatus = Literal["pending", "running", "completed", "failed", "cancelled"] + + +class CopilotEditorSummary(BaseModel): + model_config = ConfigDict(extra="forbid") + + revision: int = Field(ge=1) + duration_ms: int = Field(ge=0) + track_count: int = Field(ge=0, le=32) + clip_count: int = Field(ge=0, le=500) + + +class CopilotContextInput(BaseModel): + model_config = ConfigDict(extra="forbid") + + project_id: UUID | None = None + selected_asset_ids: list[UUID] = Field(default_factory=list, max_length=20) + selected_clip_ids: list[str] = Field(default_factory=list, max_length=20) + active_tool: str | None = Field(default=None, max_length=100) + editor_summary: CopilotEditorSummary | None = None + + +class CopilotContext(CopilotContextInput): + workspace_id: str + available_capabilities: list[str] = Field(default_factory=list, max_length=100) + + +class ProjectOpenArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + + +class AssetSelectArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + asset_id: UUID + project_id: UUID | None = None + + +class AiGenerateImageArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + prompt: str = Field(min_length=1, max_length=4_000) + project_id: UUID | None = None + source_asset_id: UUID | None = None + model: str | None = Field(default=None, max_length=255) + + +class AiGenerateVideoArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + prompt: str = Field(min_length=1, max_length=4_000) + project_id: UUID | None = None + source_asset_id: UUID + model: str | None = Field(default=None, max_length=255) + + +class EditorSplitClipArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + clip_id: str = Field(min_length=1, max_length=128) + at_ms: int = Field(gt=0) + expected_revision: int = Field(ge=1) + + +class EditorDeleteClipArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + clip_id: str = Field(min_length=1, max_length=128) + expected_revision: int = Field(ge=1) + + +class EditorSetDurationArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + clip_id: str = Field(min_length=1, max_length=128) + duration_ms: int = Field(gt=0) + expected_revision: int = Field(ge=1) + + +class EditorAddClipArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + asset_id: UUID + expected_revision: int = Field(ge=1) + duration_ms: int = Field(default=5_000, gt=0) + + +class EditorRenderArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID + expected_revision: int = Field(ge=1) + +class TemplateSearchArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + query: str = Field(min_length=1, max_length=200) + category: str | None = Field(default=None, max_length=50) + + +class TemplateGetArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + template_id: UUID + + +class TemplateApplyArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + template_id: UUID + project_id: UUID + template_version_id: UUID | None = None + slot_bindings: dict[str, SlotBinding] = Field(default_factory=dict, max_length=100) + + +class TemplateCreateProjectArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + template_id: UUID + project_name: str = Field(min_length=1, max_length=200) + template_version_id: UUID | None = None + slot_bindings: dict[str, SlotBinding] = Field(default_factory=dict, max_length=100) + + +class PublishingValidateArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + post_id: UUID + + +class PublishingCreatePostArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + post: SocialPostCreate + + +class PublishingScheduleArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + post_id: UUID + schedule: SocialScheduleCreate + + +class PublishingPublishArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + post_id: UUID + + +class PublishingCancelArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + post_id: UUID + + +class AnalyticsOverviewArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID | None = None + provider: str | None = Field(default=None, max_length=32) + metric: Literal["views", "impressions", "likes", "comments", "shares", "engagement_rate"] = ( + "views" + ) + timezone: str = Field(default="UTC", min_length=1, max_length=100) + + +class AnalyticsSyncArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: UUID | None = None + provider: str | None = Field(default=None, max_length=32) + timezone: str = Field(default="UTC", min_length=1, max_length=100) + + +class _ActionBase(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str = Field(min_length=1, max_length=64) + reason: str = Field(min_length=1, max_length=500) + requires_confirmation: bool + destructive: bool + external_side_effect: bool + required_permission: str = Field(min_length=1, max_length=100) + required_capability: str = Field(min_length=1, max_length=100) + + +class ProjectOpenAction(_ActionBase): + type: Literal["project.open"] + arguments: ProjectOpenArguments + + +class AssetSelectAction(_ActionBase): + type: Literal["asset.select"] + arguments: AssetSelectArguments + + +class AiGenerateImageAction(_ActionBase): + type: Literal["ai.generate_image"] + arguments: AiGenerateImageArguments + + +class AiGenerateVideoAction(_ActionBase): + type: Literal["ai.generate_video"] + arguments: AiGenerateVideoArguments + + +class EditorSplitClipAction(_ActionBase): + type: Literal["editor.split_clip"] + arguments: EditorSplitClipArguments + + +class EditorDeleteClipAction(_ActionBase): + type: Literal["editor.delete_clip"] + arguments: EditorDeleteClipArguments + + +class EditorSetDurationAction(_ActionBase): + type: Literal["editor.set_duration"] + arguments: EditorSetDurationArguments + + +class EditorAddClipAction(_ActionBase): + type: Literal["editor.add_clip"] + arguments: EditorAddClipArguments + + +class EditorRenderAction(_ActionBase): + type: Literal["editor.render"] + arguments: EditorRenderArguments + + +class ShareProjectArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: str + user_id: str + role: str + +class SubmitForReviewArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: str + workflow_id: str + +class ProjectActivitySummarizeArguments(BaseModel): + model_config = ConfigDict(extra="forbid") + project_id: str + +class ShareProjectAction(_ActionBase): + type: Literal["project.share"] + arguments: ShareProjectArguments + +class SubmitForReviewAction(_ActionBase): + type: Literal["review.submit"] + arguments: SubmitForReviewArguments + +class ProjectActivitySummarizeAction(_ActionBase): + type: Literal["project_activity.summarize"] + arguments: ProjectActivitySummarizeArguments + + +class TemplateSearchAction(_ActionBase): + type: Literal["template.search"] + arguments: TemplateSearchArguments + + +class TemplateGetAction(_ActionBase): + type: Literal["template.get"] + arguments: TemplateGetArguments + + +class TemplateApplyAction(_ActionBase): + type: Literal["template.apply"] + arguments: TemplateApplyArguments + + +class TemplateCreateProjectAction(_ActionBase): + type: Literal["template.create_project"] + arguments: TemplateCreateProjectArguments + + +class PublishingValidateAction(_ActionBase): + type: Literal["publishing.validate"] + arguments: PublishingValidateArguments + + +class PublishingCreatePostAction(_ActionBase): + type: Literal["publishing.create_post"] + arguments: PublishingCreatePostArguments + + +class PublishingScheduleAction(_ActionBase): + type: Literal["publishing.schedule"] + arguments: PublishingScheduleArguments + + +class PublishingPublishAction(_ActionBase): + type: Literal["publishing.publish"] + arguments: PublishingPublishArguments + + +class PublishingCancelAction(_ActionBase): + type: Literal["publishing.cancel"] + arguments: PublishingCancelArguments + + +class AnalyticsOverviewAction(_ActionBase): + type: Literal["analytics.overview"] + arguments: AnalyticsOverviewArguments + + +class AnalyticsSyncAction(_ActionBase): + type: Literal["analytics.sync"] + arguments: AnalyticsSyncArguments + + +CopilotAction = Annotated[ + ProjectOpenAction + | AssetSelectAction + | AiGenerateImageAction + | AiGenerateVideoAction + | EditorSplitClipAction + | EditorDeleteClipAction + | EditorSetDurationAction + | EditorAddClipAction + | EditorRenderAction + | ShareProjectAction + | SubmitForReviewAction + | ProjectActivitySummarizeAction + | TemplateSearchAction + | TemplateGetAction + | TemplateApplyAction + | TemplateCreateProjectAction + | PublishingValidateAction + | PublishingCreatePostAction + | PublishingScheduleAction + | PublishingPublishAction + | PublishingCancelAction + | AnalyticsOverviewAction + | AnalyticsSyncAction, + Field(discriminator="type"), +] + + +class CopilotPlan(BaseModel): + model_config = ConfigDict(extra="forbid") + + intent: str = Field(min_length=1, max_length=200) + explanation: str = Field(min_length=1, max_length=1_000) + actions: list[CopilotAction] = Field(default_factory=list, max_length=20) + missing_information: list[str] = Field(default_factory=list, max_length=20) + unsupported_capabilities: list[str] = Field(default_factory=list, max_length=20) + executable: bool + requires_confirmation: bool + + @model_validator(mode="after") + def validate_execution(self) -> "CopilotPlan": + if self.executable and not self.actions: + raise ValueError("Executable plans require at least one action") + if self.requires_confirmation != any( + action.requires_confirmation for action in self.actions + ): + raise ValueError("Plan confirmation state must match its actions") + return self + + +class CopilotActionResult(BaseModel): + model_config = ConfigDict(extra="forbid") + + action_id: str + action_type: str + status: CopilotActionStatus + summary: str = Field(max_length=1_000) + resource_type: str | None = Field(default=None, max_length=100) + resource_id: str | None = Field(default=None, max_length=255) + retryable: bool = False + error_code: str | None = Field(default=None, max_length=100) + + +class CopilotRunCreate(BaseModel): + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + request: str = Field(min_length=1, max_length=4_000) + context: CopilotContextInput = Field(default_factory=CopilotContextInput) + + +class CopilotExecuteRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + confirmed: bool = False + + +class CopilotRun(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + project_id: str | None + status: CopilotRunStatus + request: str + context: CopilotContext + plan: CopilotPlan + current_action_id: str | None + results: list[CopilotActionResult] + summary: str | None + error_code: str | None + error_message: str | None + created_at: datetime + updated_at: datetime + completed_at: datetime | None + + +class CopilotRunList(BaseModel): + items: list[CopilotRun] + offset: int + limit: int + + +class CopilotActionCapability(BaseModel): + model_config = ConfigDict(extra="forbid") + type: str + description: str + required_permission: str + required_capability: str + destructive: bool + external_side_effect: bool + requires_confirmation: bool + available: bool + + +class CopilotCapabilities(BaseModel): + model_config = ConfigDict(extra="forbid") + available: bool + planner: Literal["deterministic"] + actions: list[CopilotActionCapability] + permissions: list[str] diff --git a/app/copilot/service.py b/app/copilot/service.py new file mode 100644 index 0000000000000000000000000000000000000000..4e95f813965d262428d7ec4e598c3ddb8af91da0 --- /dev/null +++ b/app/copilot/service.py @@ -0,0 +1,499 @@ +from __future__ import annotations + +import hashlib +import json +import time + +from app.ai.service import AiStudioService +from app.copilot.actions import CopilotActionRegistry +from app.copilot.errors import ( + CopilotConfirmationRequiredError, + CopilotInvalidRequestError, + CopilotRunConflictError, +) +from app.copilot.models import CopilotRunRecord +from app.copilot.planner import CopilotPlanner +from app.copilot.repository import CopilotRepository +from app.copilot.schemas import ( + CopilotCapabilities, + CopilotContext, + CopilotContextInput, + CopilotEditorSummary, + CopilotExecuteRequest, + CopilotPlan, + CopilotRun, + CopilotRunCreate, + CopilotRunList, +) +from app.core.exceptions import MediaAPIError +from app.core.logger import get_logger +from app.projects.errors import ProjectEditorNotFoundError +from app.projects.services.editor_service import ProjectEditorService +from app.projects.services.project_service import ProjectService +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService + +logger = get_logger(__name__) + + +class CopilotService: + def __init__( + self, + *, + repository: CopilotRepository, + planner: CopilotPlanner, + actions: CopilotActionRegistry, + projects: ProjectService, + assets: CanonicalAssetService, + editor: ProjectEditorService, + ai: AiStudioService, + audit: AuditService, + ) -> None: + self.repository = repository + self.planner = planner + self.actions = actions + self.projects = projects + self.assets = assets + self.editor = editor + self.ai = ai + self.audit = audit + + async def capabilities( + self, *, workspace_id: str, user_id: str, context: CopilotContextInput + ) -> CopilotCapabilities: + bounded = await self.build_context( + workspace_id=workspace_id, user_id=user_id, supplied=context + ) + available = set(bounded.available_capabilities) + return CopilotCapabilities( + available=True, + planner="deterministic", + actions=self.actions.capabilities(available), + permissions=["copilot:read", "copilot:execute"], + ) + + async def build_context( + self, + *, + workspace_id: str, + user_id: str, + supplied: CopilotContextInput, + ) -> CopilotContext: + project_id = str(supplied.project_id) if supplied.project_id else None + capabilities = { + "project.open", + "asset.select", + "template.search", + "template.get", + "template.apply", + "template.create_project", + "publishing.validate", + "publishing.create_post", + "publishing.schedule", + "publishing.publish", + "publishing.cancel", + "analytics.overview", + "analytics.sync", + } + editor_summary = None + editor_document = None + if project_id: + await self.projects.get( + workspace_id=workspace_id, user_id=user_id, project_id=project_id + ) + try: + editor = await self.editor.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=project_id, + ) + editor_document = editor.state + editor_summary = CopilotEditorSummary( + revision=editor.revision, + duration_ms=editor.state.duration_ms(), + track_count=len(editor.state.timeline.tracks), + clip_count=sum(len(track.clips) for track in editor.state.timeline.tracks), + ) + capabilities.update( + { + "editor.split_clip", + "editor.delete_clip", + "editor.set_duration", + "editor.add_clip", + "editor.render", + } + ) + except ProjectEditorNotFoundError: + editor_summary = None + selected_assets = [] + for asset_id in supplied.selected_asset_ids: + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=str(asset_id), + ) + except CanonicalAssetNotFoundError as exc: + raise CopilotInvalidRequestError( + "A selected asset is not available in this workspace." + ) from exc + if project_id and asset.project_id != project_id: + raise CopilotInvalidRequestError( + "Every selected asset must belong to the selected project." + ) + selected_assets.append(asset.id) + selected_clips = list(dict.fromkeys(supplied.selected_clip_ids)) + if selected_clips: + if editor_document is None: + raise CopilotInvalidRequestError("Selected clips require saved editor state.") + known = {clip.id for track in editor_document.timeline.tracks for clip in track.clips} + if any(clip_id not in known for clip_id in selected_clips): + raise CopilotInvalidRequestError( + "A selected clip is not present in the authoritative editor state." + ) + ai_capabilities = self.ai.capabilities() + for tool in ai_capabilities.tools: + if tool.available: + capabilities.add(f"ai.{tool.operation.removeprefix('generate_')}") + capabilities.add(f"ai.{tool.operation}") + return CopilotContext( + workspace_id=workspace_id, + project_id=supplied.project_id, + selected_asset_ids=selected_assets, + selected_clip_ids=selected_clips, + active_tool=supplied.active_tool, + editor_summary=editor_summary, + available_capabilities=sorted(capabilities), + ) + + async def create_run( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + payload: CopilotRunCreate, + idempotency_key: str, + ) -> CopilotRun: + key = idempotency_key.strip() + if not key or len(key) > 255: + raise CopilotInvalidRequestError("A bounded Idempotency-Key is required.") + context = await self.build_context( + workspace_id=workspace_id, user_id=user_id, supplied=payload.context + ) + fingerprint = hashlib.sha256( + json.dumps( + { + "request": payload.request, + "context": context.model_dump(mode="json"), + }, + sort_keys=True, + separators=(",", ":"), + ).encode() + ).hexdigest() + plan = self.planner.plan(payload.request, context) + self.actions.validate_plan(plan.actions) + record = CopilotRunRecord( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(context.project_id) if context.project_id else None, + idempotency_key=key, + request_fingerprint=fingerprint, + request_text=payload.request, + context_json=context.model_dump(mode="json"), + plan_json=plan.model_dump(mode="json"), + status="plan_ready" if plan.executable else "blocked", + results_json=[], + ) + created, is_new = await self.repository.create(record) + if is_new: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.run_started", + entity_type="copilot_run", + entity_id=created.id, + metadata={"project_id": created.project_id}, + ) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.plan_created", + entity_type="copilot_run", + entity_id=created.id, + metadata={ + "action_count": len(plan.actions), + "requires_confirmation": plan.requires_confirmation, + "executable": plan.executable, + }, + ) + logger.info( + "copilot plan created", + extra={ + "run_id": created.id, + "project_id": created.project_id, + "status": created.status, + "action_count": len(plan.actions), + }, + ) + return self._response(created) + + async def execute( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + run_id: str, + payload: CopilotExecuteRequest, + permissions: frozenset[str], + ) -> CopilotRun: + record = await self.repository.get(workspace_id, user_id, run_id) + if record.status not in {"plan_ready"}: + raise CopilotRunConflictError("Only a plan-ready Copilot run can be executed.") + plan = CopilotPlan.model_validate(record.plan_json) + self.actions.validate_plan(plan.actions) + if not plan.executable: + raise CopilotRunConflictError("This Copilot plan is not executable.") + if plan.requires_confirmation and not payload.confirmed: + raise CopilotConfirmationRequiredError( + "Explicit confirmation is required before this plan can run." + ) + record = await self.repository.claim_execution( + workspace_id, + user_id, + run_id, + confirmed=payload.confirmed, + ) + results = [] + started = time.monotonic() + available = set(CopilotContext.model_validate(record.context_json).available_capabilities) + for action in plan.actions: + await self.repository.update( + workspace_id, + user_id, + run_id, + status="executing", + current_action_id=action.id, + results=[item.model_dump(mode="json") for item in results], + ) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.action_started", + entity_type="copilot_run", + entity_id=run_id, + metadata={"action_id": action.id, "action_type": action.type}, + ) + try: + result = await self.actions.execute( + action, + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + run_id=run_id, + permissions=permissions, + available_capabilities=available, + ) + results.append(result) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.action_completed", + entity_type="copilot_run", + entity_id=run_id, + metadata={ + "action_id": action.id, + "action_type": action.type, + "resource_type": result.resource_type, + "resource_id": result.resource_id, + }, + ) + except Exception as exc: + if isinstance(exc, MediaAPIError): + error_code = exc.code + error_message = exc.message + else: + error_code = "COPILOT_ACTION_FAILED" + error_message = "The Copilot action failed unexpectedly." + logger.error( + "copilot action failed unexpectedly", + extra={ + "run_id": run_id, + "project_id": record.project_id, + "action_id": action.id, + "action_type": action.type, + "error_category": error_code, + }, + ) + results.append( + self._failed_result( + action.id, + action.type, + error_message, + error_code, + ) + ) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.action_failed", + entity_type="copilot_run", + entity_id=run_id, + metadata={ + "action_id": action.id, + "action_type": action.type, + "error_code": error_code, + }, + ) + break + completed = sum(item.status == "completed" for item in results) + failed = sum(item.status == "failed" for item in results) + if failed and completed: + status = "partial" + summary = ( + f"{completed} of {len(plan.actions)} actions completed; " + "the remaining workflow stopped after a failure." + ) + event = "copilot.run_failed" + elif failed: + status = "failed" + summary = "The Copilot action failed before the workflow completed." + event = "copilot.run_failed" + else: + status = "completed" + summary = f"{completed} action{'s' if completed != 1 else ''} completed." + event = "copilot.run_completed" + updated = await self.repository.update( + workspace_id, + user_id, + run_id, + status=status, + current_action_id=None, + results=[item.model_dump(mode="json") for item in results], + summary=summary, + error_code=results[-1].error_code if failed else None, + error_message=results[-1].summary if failed else None, + terminal=True, + ) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type=event, + entity_type="copilot_run", + entity_id=run_id, + metadata={ + "status": status, + "completed_actions": completed, + "failed_actions": failed, + }, + ) + logger.info( + "copilot run finished", + extra={ + "run_id": run_id, + "project_id": record.project_id, + "status": status, + "duration_ms": round((time.monotonic() - started) * 1_000), + }, + ) + return self._response(updated) + + async def cancel( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + run_id: str, + ) -> CopilotRun: + updated, changed = await self.repository.cancel_before_execution( + workspace_id, + user_id, + run_id, + ) + if not changed: + return self._response(updated) + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="copilot.run_cancelled", + entity_type="copilot_run", + entity_id=run_id, + metadata={"project_id": updated.project_id}, + ) + return self._response(updated) + + async def get(self, *, workspace_id: str, user_id: str, run_id: str) -> CopilotRun: + return self._response(await self.repository.get(workspace_id, user_id, run_id)) + + async def list( + self, + *, + workspace_id: str, + user_id: str, + offset: int, + limit: int, + ) -> CopilotRunList: + records = await self.repository.list(workspace_id, user_id, offset=offset, limit=limit) + return CopilotRunList( + items=[self._response(record) for record in records], + offset=offset, + limit=limit, + ) + + @staticmethod + def _failed_result(action_id: str, action_type: str, message: str, error_code: str): + from app.copilot.schemas import CopilotActionResult + + return CopilotActionResult( + action_id=action_id, + action_type=action_type, + status="failed", + summary=message[:1_000], + error_code=error_code, + retryable=isinstance(error_code, str) + and error_code + in { + "GENERATION_PROVIDER_UNAVAILABLE", + "RATE_LIMIT_EXCEEDED", + "PROJECT_EDITOR_REVISION_CONFLICT", + }, + ) + + @staticmethod + def _response(record: CopilotRunRecord) -> CopilotRun: + return CopilotRun( + id=record.id, + project_id=record.project_id, + status=record.status, + request=record.request_text, + context=CopilotContext.model_validate(record.context_json), + plan=CopilotPlan.model_validate(record.plan_json), + current_action_id=record.current_action_id, + results=record.results_json or [], + summary=record.summary, + error_code=record.error_code, + error_message=record.error_message, + created_at=record.created_at, + updated_at=record.updated_at, + completed_at=record.completed_at, + ) diff --git a/app/core/config.py b/app/core/config.py index 83efec8067753c38d31fa6be4e8b20f8f1b45e8b..6ed30dfe2642df269178a7cb137749d3f02e406b 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -5,7 +5,7 @@ from functools import lru_cache from pathlib import Path from urllib.parse import urlparse -from pydantic import Field, SecretStr, field_validator +from pydantic import Field, SecretStr, field_validator, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict DEFAULT_TEMPLATE_DIR = Path(__file__).resolve().parents[1] / "templates" / "categories" @@ -20,8 +20,10 @@ class Settings(BaseSettings): app_name: str = "MediaRouter" app_version: str = "1.0.0" + app_environment: str = "development" host: str = "0.0.0.0" port: int = 7860 + cors_allowed_origins: str = "" temp_dir: Path = Path("./temp") output_dir: Path = Path("./outputs") template_dir: Path = DEFAULT_TEMPLATE_DIR @@ -40,6 +42,11 @@ class Settings(BaseSettings): ffprobe_binary: str = "ffprobe" auth_enabled: bool = True database_url: str = "sqlite+aiosqlite:///./data/mediarouter.db" + # Security/tenant schema creation is automatic only for SQLite local + # development. PostgreSQL deployments must apply SQL migrations explicitly. + security_auto_migrate: bool = False + security_database_role: str = "" + security_enforce_rls: bool = True auth_role_scopes: dict[str, list[str]] = Field(default_factory=dict) auth_bootstrap_key_hash: str = "" auth_bootstrap_key_prefix: str = "" @@ -49,9 +56,7 @@ class Settings(BaseSettings): auth_default_requests_per_minute: int = Field(default=100, ge=1, le=1_000_000) auth_default_concurrent_jobs: int = Field(default=10, ge=1, le=10_000) auth_default_uploads_per_hour: int = Field(default=20, ge=1, le=1_000_000) - auth_default_processing_bytes_per_day: int = Field( - default=107_374_182_400, ge=1_048_576 - ) + auth_default_processing_bytes_per_day: int = Field(default=107_374_182_400, ge=1_048_576) auth_trust_proxy_headers: bool = True mcp_stdio_api_key: SecretStr | None = None # Social Automation foundation. The existing database remains the default @@ -59,6 +64,12 @@ class Settings(BaseSettings): # dedicated async SQLAlchemy URL and apply the SQL migration out-of-band. social_enabled: bool = True social_database_url: str = "" + # API requests use SOCIAL_DATABASE_URL with a non-BYPASSRLS role. Workers + # use a separate, backend-only connection with the trusted role below. + social_worker_database_url: str = "" + social_tenant_database_role: str = "" + social_worker_database_role: str = "" + social_enforce_rls: bool = True social_auto_migrate: bool = False social_worker_enabled: bool = True social_scheduler_interval_seconds: int = Field(default=30, ge=5, le=3600) @@ -94,9 +105,7 @@ class Settings(BaseSettings): # Direct Post requires TikTok Content Posting approval and an audited app. # Keep it fail-closed until an operator has confirmed that access. tiktok_direct_post_enabled: bool = False - tiktok_upload_chunk_bytes: int = Field( - default=10_000_000, ge=5_000_000, le=64_000_000 - ) + tiktok_upload_chunk_bytes: int = Field(default=10_000_000, ge=5_000_000, le=64_000_000) tiktok_request_timeout_seconds: float = Field(default=60.0, gt=0, le=600) tiktok_processing_poll_seconds: int = Field(default=30, ge=5, le=3600) linkedin_client_id: str = "" @@ -107,15 +116,9 @@ class Settings(BaseSettings): # the operator has confirmed the application products and write scopes in # LinkedIn Developer Portal. linkedin_publishing_enabled: bool = False - linkedin_request_timeout_seconds: float = Field( - default=60.0, gt=0, le=600 - ) - linkedin_media_processing_poll_seconds: int = Field( - default=5, ge=1, le=300 - ) - linkedin_media_processing_timeout_seconds: int = Field( - default=600, ge=30, le=3600 - ) + linkedin_request_timeout_seconds: float = Field(default=60.0, gt=0, le=600) + linkedin_media_processing_poll_seconds: int = Field(default=5, ge=1, le=300) + linkedin_media_processing_timeout_seconds: int = Field(default=600, ge=30, le=3600) x_client_id: str = "" x_client_secret: SecretStr | None = None # Exact OAuth 2.0 callback registered for the confidential X Web App. @@ -132,6 +135,45 @@ class Settings(BaseSettings): whatsapp_client_id: str = "" whatsapp_client_secret: SecretStr | None = None social_oauth_redirect_base_url: str = "" + # Provider-neutral generation runtime. Providers remain optional and their + # worker endpoints/tokens stay server-side only. + generation_enabled: bool = True + generation_job_retry_limit: int = Field(default=3, ge=0, le=20) + # Shared remote-generation worker transport defaults. These do not enable + # a provider and intentionally contain no worker URL or credentials. + ai_worker_connect_timeout_seconds: float = Field(default=10.0, gt=0, le=300) + ai_worker_request_timeout_seconds: float = Field(default=60.0, gt=0, le=3600) + ai_worker_read_timeout_seconds: float = Field(default=300.0, gt=0, le=7200) + ai_worker_max_retries: int = Field(default=3, ge=0, le=10) + ai_worker_retry_backoff_seconds: float = Field(default=0.5, ge=0, le=60) + # WAN is optional. These values remain backend-only and are deliberately + # not validated at Settings construction time: a bad optional worker + # configuration must leave WAN unavailable without preventing unrelated + # MediaRouter services from starting. + wan_space_url: str = "" + wan_space_token: SecretStr | None = None + # Optional authenticated FLUX.2 Klein worker. Invalid configuration keeps + # FLUX unavailable without affecting startup or other providers. + flux_space_url: str = "" + flux_space_token: SecretStr | None = None + generation_worker_enabled: bool = True + generation_worker_interval_seconds: float = Field(default=5.0, ge=0.5, le=3600) + generation_worker_poll_backoff_seconds: float = Field(default=2.0, ge=0.5, le=300) + generation_worker_batch_size: int = Field(default=8, ge=1, le=100) + generation_job_stale_after_seconds: int = Field(default=900, ge=60, le=86_400) + # Content Studio persistence/render safety limits. Rendering is optional + # infrastructure; disabling its worker never prevents API startup. + editor_state_max_bytes: int = Field(default=1_048_576, ge=16_384, le=16_777_216) + render_worker_enabled: bool = True + render_worker_interval_seconds: float = Field(default=2.0, ge=0.5, le=3600) + render_job_stale_after_seconds: int = Field(default=900, ge=60, le=86_400) + render_job_timeout_seconds: int = Field(default=7200, ge=60, le=86_400) + render_job_retry_limit: int = Field(default=2, ge=0, le=10) + render_max_active_jobs_per_project: int = Field(default=1, ge=1, le=10) + render_max_tracks: int = Field(default=32, ge=1, le=256) + render_max_clips: int = Field(default=500, ge=1, le=10_000) + render_max_duration_seconds: int = Field(default=3600, ge=1, le=21_600) + render_max_input_bytes: int = Field(default=4_294_967_296, ge=1_048_576) @field_validator("whisper_model") @classmethod @@ -150,6 +192,107 @@ class Settings(BaseSettings): raise ValueError(f"LOG_LEVEL must be one of: {', '.join(sorted(allowed))}") return normalized + @field_validator("app_environment") + @classmethod + def normalize_app_environment(cls, value: str) -> str: + normalized = value.strip().lower() + if normalized not in {"development", "test", "production"}: + raise ValueError("APP_ENVIRONMENT must be development, test, or production") + return normalized + + @property + def allowed_cors_origins(self) -> tuple[str, ...]: + """Return validated, normalized origins for Starlette CORS middleware.""" + + origins: list[str] = [] + for configured in self.cors_allowed_origins.split(","): + origin = configured.strip().rstrip("/") + if not origin: + continue + parsed = urlparse(origin) + if ( + "*" in origin + or parsed.scheme not in {"http", "https"} + or not parsed.netloc + or parsed.username + or parsed.password + or parsed.path + or parsed.params + or parsed.query + or parsed.fragment + ): + raise ValueError( + "CORS_ALLOWED_ORIGINS must contain comma-separated HTTP(S) origins" + ) + if origin not in origins: + origins.append(origin) + return tuple(origins) + + @model_validator(mode="after") + def validate_production_contract(self) -> Settings: + """Fail clearly when the Docker production boundary is unsafe. + + Local development and tests retain the established SQLite defaults. + The production Dockerfile sets ``APP_ENVIRONMENT=production``, making + external PostgreSQL, RLS, explicit migrations, authentication, and an + exact frontend CORS origin mandatory at process import/startup. + """ + + origins = self.allowed_cors_origins + if self.app_environment != "production": + return self + + errors: list[str] = [] + if not self._is_external_postgres(self.database_url): + errors.append("DATABASE_URL must use an external PostgreSQL database in production") + if self.security_auto_migrate: + errors.append("SECURITY_AUTO_MIGRATE must be false in production") + if not self.security_enforce_rls: + errors.append("SECURITY_ENFORCE_RLS must be true in production") + if not self.security_database_role.strip(): + errors.append("SECURITY_DATABASE_ROLE is required in production") + if not self.auth_enabled: + errors.append("AUTH_ENABLED must be true in production") + if not origins: + errors.append("CORS_ALLOWED_ORIGINS must include the HTTPS Vercel frontend origin") + elif any(urlparse(origin).scheme != "https" for origin in origins): + errors.append("CORS_ALLOWED_ORIGINS must use HTTPS in production") + + if self.social_auto_migrate: + errors.append("SOCIAL_AUTO_MIGRATE must be false in production") + if self.social_enabled: + if not self.social_database_url.strip() or not self._is_external_postgres( + self.social_database_url + ): + errors.append( + "SOCIAL_DATABASE_URL must use an explicit external PostgreSQL tenant connection" + ) + if not self.social_tenant_database_role.strip(): + errors.append("SOCIAL_TENANT_DATABASE_ROLE is required when social is enabled") + if not self._is_external_postgres(self.social_worker_database_url): + errors.append( + "SOCIAL_WORKER_DATABASE_URL must use an external PostgreSQL worker connection" + ) + if not self.social_worker_database_role.strip(): + errors.append("SOCIAL_WORKER_DATABASE_ROLE is required when social is enabled") + if not self.social_enforce_rls: + errors.append("SOCIAL_ENFORCE_RLS must be true when social is enabled") + + if errors: + raise ValueError("Invalid production configuration: " + "; ".join(errors)) + return self + + @staticmethod + def _is_external_postgres(value: str) -> bool: + configured = value.strip() + if not configured: + return False + parsed = urlparse(configured) + if parsed.scheme not in {"postgres", "postgresql", "postgresql+asyncpg"}: + return False + hostname = (parsed.hostname or "").lower() + return bool(hostname and hostname not in {"localhost", "127.0.0.1", "::1"}) + def ensure_directories(self) -> None: self.temp_dir.mkdir(parents=True, exist_ok=True) self.output_dir.mkdir(parents=True, exist_ok=True) @@ -245,9 +388,7 @@ class Settings(BaseSettings): or (parsed.scheme == "http" and parsed.hostname in local_hosts) ) ): - raise ValueError( - "X_REDIRECT_URI must be the HTTPS MediaRouter X callback URI" - ) + raise ValueError("X_REDIRECT_URI must be the HTTPS MediaRouter X callback URI") return normalized @field_validator("linkedin_redirect_uri") diff --git a/app/core/database_url.py b/app/core/database_url.py new file mode 100644 index 0000000000000000000000000000000000000000..6e171d297af91368257964f2ec85763e6e673410 --- /dev/null +++ b/app/core/database_url.py @@ -0,0 +1,18 @@ +from __future__ import annotations + + +def normalize_async_database_url(value: str) -> str: + """Use the installed asyncpg dialect for bare PostgreSQL URLs. + + Operators frequently supply a standard ``postgresql://`` URL. SQLAlchemy + otherwise selects synchronous psycopg2 for that form, which fails inside + this async application and can tempt deployments to add an unnecessary + synchronous driver. Explicit dialect URLs remain unchanged. + """ + + normalized = value.strip() + if normalized.startswith("postgresql://"): + return "postgresql+asyncpg://" + normalized.removeprefix("postgresql://") + if normalized.startswith("postgres://"): + return "postgresql+asyncpg://" + normalized.removeprefix("postgres://") + return normalized diff --git a/app/generation/__init__.py b/app/generation/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..94b4fd70653028ca21c5bfae00fb639928d2b383 --- /dev/null +++ b/app/generation/__init__.py @@ -0,0 +1,5 @@ +"""Provider-neutral, durable generation orchestration. + +This package contains provider-neutral generation orchestration plus audited, +optional WAN and FLUX adapters. +""" diff --git a/app/generation/domain/__init__.py b/app/generation/domain/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..51a8936b77ef9b8772f1ee64b19866f35116bf00 --- /dev/null +++ b/app/generation/domain/__init__.py @@ -0,0 +1,23 @@ +"""Generation-domain value types, capabilities, and errors.""" + +from app.generation.domain.enums import ( + GenerationJobStatus, + GenerationModality, + GenerationRequestStatus, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerHealthStatus, + WorkerJobStatus, + WorkerReadinessStatus, +) + +__all__ = [ + "GenerationJobStatus", + "GenerationModality", + "GenerationRequestStatus", + "WorkerCancellationStatus", + "WorkerErrorCategory", + "WorkerHealthStatus", + "WorkerJobStatus", + "WorkerReadinessStatus", +] diff --git a/app/generation/domain/capabilities.py b/app/generation/domain/capabilities.py new file mode 100644 index 0000000000000000000000000000000000000000..5c4ddffec41bed9197562c556f8d63e6bcbb7b10 --- /dev/null +++ b/app/generation/domain/capabilities.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +from pydantic import BaseModel, ConfigDict, Field + +from app.generation.domain.enums import GenerationModality + + +class GenerationModelCapability(BaseModel): + """An explicitly implemented model contract. + + ``input_schema`` is descriptive only; request validation remains owned by + the adapter's typed Pydantic input model. It lets future transports build + capability-driven UIs without accepting arbitrary provider payloads. + """ + + model_config = ConfigDict(extra="forbid") + + id: str = Field(min_length=1, max_length=255) + name: str = Field(min_length=1, max_length=255) + modality: GenerationModality + input_asset_supported: bool = False + input_schema: dict[str, object] = Field(default_factory=dict) + + +class GenerationProviderCapabilities(BaseModel): + """Public, non-secret provider discovery metadata.""" + + model_config = ConfigDict(extra="forbid") + + provider: str = Field(pattern=r"^[a-z][a-z0-9_-]{0,63}$") + name: str = Field(min_length=1, max_length=255) + models: list[GenerationModelCapability] = Field(default_factory=list) + implementation_status: str = Field(default="foundation", max_length=64) + supports_cancellation: bool = False + supports_status_reconciliation: bool = False diff --git a/app/generation/domain/enums.py b/app/generation/domain/enums.py new file mode 100644 index 0000000000000000000000000000000000000000..54404d153bfbc9f3d6f52a4cbaf6be68b1c4eca4 --- /dev/null +++ b/app/generation/domain/enums.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +from app.core.enums import StrEnum + + +class GenerationModality(StrEnum): + """Output modalities supported by the foundation. + + The list is intentionally small. An adapter may only advertise a value + after it implements that modality end to end. + """ + + IMAGE = "image" + VIDEO = "video" + + +class GenerationRequestStatus(StrEnum): + QUEUED = "queued" + SUBMITTING = "submitting" + RUNNING = "running" + RETRYING = "retrying" + SUCCEEDED = "succeeded" + FAILED = "failed" + CANCEL_REQUESTED = "cancel_requested" + CANCELLED = "cancelled" + + +class GenerationJobStatus(StrEnum): + QUEUED = "queued" + SUBMITTING = "submitting" + RUNNING = "running" + RETRYING = "retrying" + SUCCEEDED = "succeeded" + FAILED = "failed" + CANCEL_REQUESTED = "cancel_requested" + CANCELLED = "cancelled" + + +class WorkerHealthStatus(StrEnum): + HEALTHY = "healthy" + STARTING = "starting" + UNAVAILABLE = "unavailable" + UNHEALTHY = "unhealthy" + UNKNOWN = "unknown" + + +class WorkerReadinessStatus(StrEnum): + READY = "ready" + STARTING = "starting" + UNAVAILABLE = "unavailable" + UNKNOWN = "unknown" + + +class WorkerJobStatus(StrEnum): + QUEUED = "queued" + RUNNING = "running" + COMPLETED = "completed" + FAILED = "failed" + CANCELLED = "cancelled" + UNKNOWN = "unknown" + + +class WorkerCancellationStatus(StrEnum): + REQUESTED = "requested" + CANCELLED = "cancelled" + UNSUPPORTED = "unsupported" + FAILED = "failed" + + +class WorkerErrorCategory(StrEnum): + INVALID_REQUEST = "invalid_request" + AUTHENTICATION_ERROR = "authentication_error" + AUTHORIZATION_ERROR = "authorization_error" + WORKER_UNAVAILABLE = "worker_unavailable" + WORKER_NOT_READY = "worker_not_ready" + TIMEOUT = "timeout" + RATE_LIMITED = "rate_limited" + PROVIDER_ERROR = "provider_error" + INFERENCE_ERROR = "inference_error" + OUTPUT_ERROR = "output_error" + CANCELLATION_ERROR = "cancellation_error" + UNKNOWN_ERROR = "unknown_error" + + +TERMINAL_GENERATION_JOB_STATUSES = frozenset( + { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + } +) diff --git a/app/generation/domain/errors.py b/app/generation/domain/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..d8f0d108bb049ff8f5fb3197fc219f48a2d7f42b --- /dev/null +++ b/app/generation/domain/errors.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +from app.core.exceptions import MediaAPIError +from app.generation.domain.enums import WorkerErrorCategory + + +class GenerationError(MediaAPIError): + code = "GENERATION_ERROR" + status_code = 400 + + +class GenerationProviderUnavailableError(GenerationError): + code = "GENERATION_PROVIDER_UNAVAILABLE" + status_code = 503 + + +class GenerationCapabilityUnsupportedError(GenerationError): + code = "GENERATION_CAPABILITY_UNSUPPORTED" + status_code = 422 + + +class GenerationValidationError(GenerationError): + code = "GENERATION_VALIDATION_ERROR" + status_code = 422 + + +class GenerationRequestNotFoundError(GenerationError): + code = "GENERATION_REQUEST_NOT_FOUND" + status_code = 404 + + +class GenerationJobNotFoundError(GenerationError): + code = "GENERATION_JOB_NOT_FOUND" + status_code = 404 + + +class GenerationIdempotencyConflictError(GenerationError): + code = "GENERATION_IDEMPOTENCY_CONFLICT" + status_code = 409 + + +class GenerationTransitionError(GenerationError): + code = "GENERATION_INVALID_STATE_TRANSITION" + status_code = 409 + + +class GenerationInputAssetNotFoundError(GenerationError): + code = "GENERATION_INPUT_ASSET_NOT_FOUND" + status_code = 404 + + +class GenerationProviderJobConflictError(GenerationError): + """A provider job ID is already bound to a different logical job.""" + + code = "GENERATION_PROVIDER_JOB_CONFLICT" + status_code = 409 + + +class GenerationOutputConflictError(GenerationError): + code = "GENERATION_OUTPUT_CONFLICT" + status_code = 409 + + +class GenerationWorkerError(GenerationError): + """Safe normalised failure returned by a remote generation worker.""" + + code = "GENERATION_WORKER_ERROR" + status_code = 502 + + _CLIENT_STATUS_BY_CATEGORY = { + WorkerErrorCategory.INVALID_REQUEST: 422, + WorkerErrorCategory.WORKER_UNAVAILABLE: 503, + WorkerErrorCategory.WORKER_NOT_READY: 503, + WorkerErrorCategory.TIMEOUT: 504, + WorkerErrorCategory.RATE_LIMITED: 503, + } + + def __init__( + self, + *, + category: WorkerErrorCategory, + message: str, + retryable: bool, + http_status: int | None = None, + ) -> None: + super().__init__(message) + self.category = category + self.retryable = retryable + self.http_status = http_status + self.status_code = self._CLIENT_STATUS_BY_CATEGORY.get(category, 502) + + +class GenerationCancellationError(GenerationWorkerError): + code = "GENERATION_CANCELLATION_FAILED" + + def __init__(self, message: str = "Generation cancellation could not be confirmed.") -> None: + super().__init__( + category=WorkerErrorCategory.CANCELLATION_ERROR, + message=message, + retryable=False, + ) + + +class GenerationOutputError(GenerationWorkerError): + code = "GENERATION_OUTPUT_INVALID" + + def __init__(self, message: str = "Generation worker output is invalid.") -> None: + super().__init__( + category=WorkerErrorCategory.OUTPUT_ERROR, + message=message, + retryable=False, + ) diff --git a/app/generation/domain/retry.py b/app/generation/domain/retry.py new file mode 100644 index 0000000000000000000000000000000000000000..3afcc30e22e0c56b1ddacd7370e01ad568743be3 --- /dev/null +++ b/app/generation/domain/retry.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from app.generation.domain.enums import WorkerErrorCategory + + +RETRYABLE_HTTP_STATUS_CODES = frozenset({429, 502, 503, 504}) +RETRYABLE_CATEGORIES = frozenset( + { + WorkerErrorCategory.WORKER_UNAVAILABLE, + WorkerErrorCategory.WORKER_NOT_READY, + WorkerErrorCategory.TIMEOUT, + WorkerErrorCategory.RATE_LIMITED, + } +) + + +@dataclass(frozen=True, slots=True) +class GenerationRetryDecision: + retryable: bool + delay_seconds: float = 0.0 + + +class GenerationRetryPolicy: + """Bounded retry classification for safe remote-worker operations.""" + + def __init__(self, *, max_retries: int, backoff_seconds: float) -> None: + self.max_retries = max(0, max_retries) + self.backoff_seconds = max(0.0, backoff_seconds) + + def decide( + self, + *, + category: WorkerErrorCategory, + http_status: int | None, + retry_number: int, + idempotent: bool, + ) -> GenerationRetryDecision: + """Return a retry decision after one failed request attempt. + + ``retry_number`` is zero for the first possible retry. HTTP 500 and + unknown/programming failures deliberately do not receive a retry. + """ + + if not idempotent or retry_number >= self.max_retries: + return GenerationRetryDecision(False) + if http_status in RETRYABLE_HTTP_STATUS_CODES or category in RETRYABLE_CATEGORIES: + return GenerationRetryDecision( + True, + delay_seconds=self.backoff_seconds * (2**retry_number), + ) + return GenerationRetryDecision(False) diff --git a/app/generation/domain/runtime.py b/app/generation/domain/runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..d7b9beb2b8b5a99eb238f7fd49828b99efcf3b9c --- /dev/null +++ b/app/generation/domain/runtime.py @@ -0,0 +1,317 @@ +from __future__ import annotations + +import math +import re +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerHealthStatus, + WorkerJobStatus, + WorkerReadinessStatus, +) + +_EXTERNAL_ID = re.compile(r"^[A-Za-z0-9._:-]{1,255}$") +_SHA256 = re.compile(r"^[0-9a-f]{64}$") +_SAFE_FILENAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,254}$") +_BEARER = re.compile(r"(?i)\bbearer\s+[A-Za-z0-9._~+/=-]+") +_HTTP_URL = re.compile(r"(?i)\bhttps?://[^\s\"'<>]+") +_ASSIGNED_SECRET = re.compile( + r"(?i)(access[_-]?token|refresh[_-]?token|id[_-]?token|client[_-]?secret|" + r"authorization|api[_-]?key|password|secret|credential)" + r"([\"']?\s*[:=]\s*[\"']?)([^\"'\s,&}]+)" +) +_SENSITIVE_METADATA_PARTS = frozenset( + { + "access_token", + "refresh_token", + "id_token", + "token", + "secret", + "authorization", + "cookie", + "password", + "api_key", + "credential", + "url", + "uri", + } +) +_MAX_METADATA_DEPTH = 8 +_MAX_METADATA_ITEMS = 256 +_MAX_METADATA_STRING_LENGTH = 8_192 +_DROP = object() + + +def _safe_metadata(value: Any, *, depth: int = 0) -> Any: + """Return bounded JSON-safe metadata with credential-like content removed.""" + + if depth > _MAX_METADATA_DEPTH: + return _DROP + if value is None or isinstance(value, bool) or isinstance(value, int): + return value + if isinstance(value, float): + return value if math.isfinite(value) else _DROP + if isinstance(value, str): + without_bearer = _BEARER.sub("Bearer [REDACTED]", value) + without_urls = _HTTP_URL.sub("[REDACTED_URL]", without_bearer) + redacted = _ASSIGNED_SECRET.sub( + lambda match: f"{match.group(1)}{match.group(2)}[REDACTED]", + without_urls, + ) + return redacted[:_MAX_METADATA_STRING_LENGTH] + if isinstance(value, dict): + result: dict[str, object] = {} + for key, item in list(value.items())[:_MAX_METADATA_ITEMS]: + if not isinstance(key, str) or len(key) > 255: + continue + normalized_key = key.lower().replace("-", "_") + if any(part in normalized_key for part in _SENSITIVE_METADATA_PARTS): + continue + cleaned = _safe_metadata(item, depth=depth + 1) + if cleaned is not _DROP: + result[key] = cleaned + return result + if isinstance(value, (list, tuple)): + result: list[object] = [] + for item in list(value)[:_MAX_METADATA_ITEMS]: + cleaned = _safe_metadata(item, depth=depth + 1) + if cleaned is not _DROP: + result.append(cleaned) + return result + return _DROP + + +def safe_worker_metadata(value: Any) -> Any: + """Remove secret-like fields and non-JSON values before persistence. + + Worker metadata is useful for diagnostics, but it is never a credential + store. This helper is deliberately conservative and is applied both at + worker-transport parsing and at persistence boundaries. + """ + + cleaned = _safe_metadata(value) + return {} if cleaned is _DROP and isinstance(value, dict) else cleaned + + +def _safe_metadata_dict(value: Any) -> dict[str, object]: + cleaned = safe_worker_metadata(value) + if not isinstance(cleaned, dict): + raise ValueError("worker metadata must be a JSON object") + return cleaned + + +class WorkerModelInfo(BaseModel): + """One worker-discovered model; never an operator-configured endpoint.""" + + model_config = ConfigDict(extra="forbid") + + id: str = Field(min_length=1, max_length=255) + name: str = Field(min_length=1, max_length=255) + media_types: list[GenerationModality] = Field(default_factory=list) + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + +class WorkerInfo(BaseModel): + """Verified non-secret identity returned by a configured worker.""" + + model_config = ConfigDict(extra="forbid") + + id: str = Field(min_length=1, max_length=255) + name: str = Field(min_length=1, max_length=255) + media_types: list[GenerationModality] = Field(default_factory=list) + models: list[WorkerModelInfo] = Field(default_factory=list) + status: WorkerHealthStatus = WorkerHealthStatus.UNKNOWN + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("models") + @classmethod + def unique_models(cls, values: list[WorkerModelInfo]) -> list[WorkerModelInfo]: + if len({value.id for value in values}) != len(values): + raise ValueError("worker metadata contains duplicate model IDs") + return values + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + +class WorkerHealth(BaseModel): + """Normalised liveness result. It never implies model availability.""" + + model_config = ConfigDict(extra="forbid") + + status: WorkerHealthStatus + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + +class WorkerReadiness(BaseModel): + """Normalised inference readiness result.""" + + model_config = ConfigDict(extra="forbid") + + status: WorkerReadinessStatus + model_loaded: bool = False + model_ids: list[str] = Field(default_factory=list) + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("model_ids") + @classmethod + def valid_model_ids(cls, values: list[str]) -> list[str]: + if len(values) != len(set(values)): + raise ValueError("worker readiness model_ids must be unique") + for value in values: + if not value or len(value) > 255 or any(character.isspace() for character in value): + raise ValueError("worker readiness model_ids contain an invalid value") + return values + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + +class WorkerOutput(BaseModel): + """A worker-issued descriptor, never a filesystem path or arbitrary URL.""" + + model_config = ConfigDict(extra="forbid") + + output_type: GenerationModality + mime_type: str = Field( + min_length=3, + max_length=255, + pattern=r"^[a-z0-9!#$&^_.+-]+/[a-z0-9!#$&^_.+-]+$", + ) + provider_output_id: str = Field(min_length=1, max_length=255) + # A worker may expose a download endpoint only under its configured origin. + # RemoteWorkerClient rejects absolute URLs, traversal, queries, and fragments. + download_path: str = Field(min_length=2, max_length=2048) + filename: str | None = Field(default=None, max_length=255) + sha256: str | None = Field(default=None, max_length=64) + byte_size: int | None = Field(default=None, ge=0) + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("provider_output_id") + @classmethod + def valid_provider_output_id(cls, value: str) -> str: + if _EXTERNAL_ID.fullmatch(value) is None: + raise ValueError("provider_output_id contains unsupported characters") + return value + + @field_validator("download_path") + @classmethod + def valid_download_path(cls, value: str) -> str: + if ( + not value.startswith("/") + or "//" in value + or "\\" in value + or "%" in value + or "?" in value + or "#" in value + or any(part in {"", ".", ".."} for part in value.split("/")[1:]) + ): + raise ValueError("download_path must be a safe absolute worker-relative path") + return value + + @field_validator("filename") + @classmethod + def valid_filename(cls, value: str | None) -> str | None: + if value is not None and _SAFE_FILENAME.fullmatch(value) is None: + raise ValueError("filename must not contain a path") + return value + + @field_validator("sha256") + @classmethod + def valid_sha256(cls, value: str | None) -> str | None: + if value is not None and _SHA256.fullmatch(value) is None: + raise ValueError("sha256 must be a lowercase SHA-256 hex digest") + return value + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + @model_validator(mode="after") + def media_type_matches_output_type(self) -> "WorkerOutput": + if not self.mime_type.startswith(f"{self.output_type.value}/"): + raise ValueError("output MIME type does not match its generation modality") + return self + + +class WorkerJob(BaseModel): + """Provider-neutral remote job representation.""" + + model_config = ConfigDict(extra="forbid") + + external_job_id: str = Field(min_length=1, max_length=255) + status: WorkerJobStatus + output: WorkerOutput | None = None + error_category: WorkerErrorCategory | None = None + error_code: str | None = Field(default=None, max_length=100) + error_message: str | None = Field(default=None, max_length=500) + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("external_job_id") + @classmethod + def valid_external_job_id(cls, value: str) -> str: + if _EXTERNAL_ID.fullmatch(value) is None: + raise ValueError("external_job_id contains unsupported characters") + return value + + @field_validator("error_code") + @classmethod + def valid_error_code(cls, value: str | None) -> str | None: + if value is not None and _EXTERNAL_ID.fullmatch(value) is None: + raise ValueError("worker error_code contains unsupported characters") + return value + + @field_validator("error_message", mode="before") + @classmethod + def clean_error_message(cls, value: Any) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise ValueError("worker error_message must be a string") + cleaned = safe_worker_metadata(value) + if not isinstance(cleaned, str): # Defensive: strings are JSON-safe. + raise ValueError("worker error_message is invalid") + return cleaned + + @model_validator(mode="after") + def completed_job_requires_output(self) -> "WorkerJob": + if self.status is WorkerJobStatus.COMPLETED and self.output is None: + raise ValueError("completed worker jobs must include an output descriptor") + return self + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) + + +class WorkerCancellationResult(BaseModel): + model_config = ConfigDict(extra="forbid") + + status: WorkerCancellationStatus + metadata: dict[str, object] = Field(default_factory=dict) + + @field_validator("metadata", mode="before") + @classmethod + def clean_metadata(cls, value: Any) -> dict[str, object]: + return _safe_metadata_dict(value) diff --git a/app/generation/domain/state_machine.py b/app/generation/domain/state_machine.py new file mode 100644 index 0000000000000000000000000000000000000000..5ae2b4948672cd4ac77a5053e5b13fe8ef51f9e8 --- /dev/null +++ b/app/generation/domain/state_machine.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +from app.generation.domain.enums import GenerationJobStatus +from app.generation.domain.errors import GenerationTransitionError + + +ALLOWED_TRANSITIONS: dict[GenerationJobStatus, frozenset[GenerationJobStatus]] = { + GenerationJobStatus.QUEUED: frozenset( + { + GenerationJobStatus.SUBMITTING, + # A trusted remote worker can acknowledge a queued/running or + # already-completed job before the next reconciliation pass. The + # repository only permits these transitions after an opaque worker + # job ID has been bound; public callers cannot make them. + GenerationJobStatus.RUNNING, + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.CANCEL_REQUESTED, + GenerationJobStatus.CANCELLED, + GenerationJobStatus.FAILED, + } + ), + GenerationJobStatus.SUBMITTING: frozenset( + { + GenerationJobStatus.RUNNING, + GenerationJobStatus.QUEUED, + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.RETRYING, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCEL_REQUESTED, + GenerationJobStatus.CANCELLED, + } + ), + GenerationJobStatus.RUNNING: frozenset( + { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.RETRYING, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCEL_REQUESTED, + GenerationJobStatus.CANCELLED, + } + ), + GenerationJobStatus.RETRYING: frozenset( + { + GenerationJobStatus.SUBMITTING, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + } + ), + GenerationJobStatus.CANCEL_REQUESTED: frozenset( + { + GenerationJobStatus.CANCELLED, + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + } + ), + GenerationJobStatus.SUCCEEDED: frozenset(), + GenerationJobStatus.FAILED: frozenset(), + GenerationJobStatus.CANCELLED: frozenset(), +} + + +def validate_transition( + current: str | GenerationJobStatus, target: str | GenerationJobStatus +) -> GenerationJobStatus: + source = GenerationJobStatus(current) + destination = GenerationJobStatus(target) + if destination not in ALLOWED_TRANSITIONS[source]: + raise GenerationTransitionError( + f"Cannot transition generation job from {source.value} to {destination.value}." + ) + return destination diff --git a/app/generation/model_registry.py b/app/generation/model_registry.py new file mode 100644 index 0000000000000000000000000000000000000000..1c3de2ae03bda0ec6aecfa2751e02b685dc9bfec --- /dev/null +++ b/app/generation/model_registry.py @@ -0,0 +1,133 @@ +from __future__ import annotations + +from collections.abc import Iterable + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from app.generation.domain.capabilities import GenerationModelCapability +from app.generation.domain.enums import WorkerReadinessStatus +from app.generation.domain.errors import GenerationCapabilityUnsupportedError +from app.generation.domain.runtime import WorkerInfo, WorkerReadiness, safe_worker_metadata + + +class GenerationModelRegistration(BaseModel): + """Trusted, server-owned model configuration metadata. + + ``configuration_reference`` names a server configuration record; it is + deliberately not a URL, credential, or client-selectable worker value. + """ + + model_config = ConfigDict(extra="forbid") + + provider_id: str = Field(pattern=r"^[a-z][a-z0-9_-]{0,63}$") + model: GenerationModelCapability + configuration_reference: str = Field( + min_length=1, max_length=128, pattern=r"^[a-z][a-z0-9_.-]{0,127}$" + ) + metadata: dict[str, object] = Field(default_factory=dict) + available: bool = False + + @model_validator(mode="after") + def no_initial_availability_claim(self) -> "GenerationModelRegistration": + if self.available: + raise ValueError("models may only become available after readiness verification") + return self + + +class GenerationModelView(BaseModel): + model_config = ConfigDict(extra="forbid") + + provider_id: str + model: GenerationModelCapability + configuration_reference: str + metadata: dict[str, object] = Field(default_factory=dict) + available: bool = False + + +class GenerationModelRegistry: + """Provider-neutral registry with availability proven by worker readiness.""" + + def __init__(self, entries: Iterable[GenerationModelRegistration] | None = None) -> None: + self._entries: dict[tuple[str, str], GenerationModelRegistration] = {} + self._availability: dict[tuple[str, str], bool] = {} + for entry in entries or (): + self.register(entry) + + def register(self, entry: GenerationModelRegistration) -> None: + key = (entry.provider_id, entry.model.id) + if key in self._entries: + raise ValueError( + f"Duplicate generation model '{entry.model.id}' for '{entry.provider_id}'." + ) + # Never retain accidental credentials in configuration metadata. + safe_metadata = safe_worker_metadata(entry.metadata) + assert isinstance(safe_metadata, dict) + self._entries[key] = entry.model_copy( + update={"metadata": safe_metadata, "available": False} + ) + self._availability[key] = False + + def list(self, *, provider_id: str | None = None) -> list[GenerationModelView]: + items = ( + (key, entry) + for key, entry in self._entries.items() + if provider_id is None or key[0] == provider_id + ) + return [self._view(key, entry) for key, entry in sorted(items)] + + def get(self, provider_id: str, model_id: str) -> GenerationModelView: + key = (provider_id, model_id) + try: + return self._view(key, self._entries[key]) + except KeyError as exc: + raise GenerationCapabilityUnsupportedError( + f"{provider_id} does not support model '{model_id}'." + ) from exc + + def verify_readiness( + self, + *, + provider_id: str, + worker_info: WorkerInfo, + readiness: WorkerReadiness, + provider_configured: bool, + ) -> list[GenerationModelView]: + """Update only models proven by matching info plus readiness. + + Liveness by itself is intentionally insufficient: the worker must be + configured, report `ready`, report a loaded model, explicitly list the + registered model ID in readiness, and discover that model with a + matching output modality through `/v1/info`. + """ + + available_ids = set(readiness.model_ids) + discovered_models = {model.id: model for model in worker_info.models} + for key, entry in self._entries.items(): + if key[0] != provider_id: + continue + discovered = discovered_models.get(entry.model.id) + self._availability[key] = bool( + provider_configured + and readiness.status is WorkerReadinessStatus.READY + and readiness.model_loaded + and entry.model.id in available_ids + and discovered is not None + and entry.model.modality in discovered.media_types + ) + return self.list(provider_id=provider_id) + + def mark_unavailable(self, provider_id: str) -> None: + for key in self._availability: + if key[0] == provider_id: + self._availability[key] = False + + def _view( + self, key: tuple[str, str], entry: GenerationModelRegistration + ) -> GenerationModelView: + return GenerationModelView( + provider_id=entry.provider_id, + model=entry.model, + configuration_reference=entry.configuration_reference, + metadata=entry.metadata, + available=self._availability.get(key, False), + ) diff --git a/app/generation/models.py b/app/generation/models.py new file mode 100644 index 0000000000000000000000000000000000000000..e73bf819101cadd716dc0f0664c1ba07fccdeeca --- /dev/null +++ b/app/generation/models.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from uuid import uuid4 + +from sqlalchemy import ( + DateTime, + ForeignKey, + Index, + Integer, + JSON, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base + + +def utcnow() -> datetime: + return datetime.now(timezone.utc) + + +def new_id() -> str: + return str(uuid4()) + + +class GenerationRequest(Base): + """Immutable validated intent, scoped to one authoritative workspace.""" + + __tablename__ = "generation_requests" + __table_args__ = ( + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_generation_request_workspace_idempotency" + ), + Index("ix_generation_requests_workspace_created", "workspace_id", "created_at"), + Index("ix_generation_requests_workspace_status", "workspace_id", "status"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + created_by_user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + provider: Mapped[str] = mapped_column(String(64), nullable=False) + model_id: Mapped[str] = mapped_column(String(255), nullable=False) + modality: Mapped[str] = mapped_column(String(32), nullable=False) + input_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="RESTRICT") + ) + project_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT") + ) + product_surface: Mapped[str] = mapped_column(String(32), nullable=False, default="generation") + # This JSON contains only adapter-validated, non-secret input. Its + # request fingerprint is authoritative for idempotency conflict checks. + spec_json: Mapped[dict[str, object]] = mapped_column("spec", JSON, nullable=False, default=dict) + request_fingerprint: Mapped[str] = mapped_column(String(64), nullable=False) + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class GenerationJob(Base): + """One durable execution for one request; retries keep this identity.""" + + __tablename__ = "generation_jobs" + __table_args__ = ( + UniqueConstraint("generation_request_id", name="uq_generation_job_request"), + # An external worker job is an opaque provider identity, not a + # workspace-scoped client value. Binding it once prevents a worker + # status/output from being attached to another tenant's logical job. + UniqueConstraint("provider", "external_job_id", name="uq_generation_job_provider_external"), + Index("ix_generation_jobs_workspace_status", "workspace_id", "status"), + Index("ix_generation_jobs_next_attempt", "status", "next_attempt_at"), + Index("ix_generation_jobs_external", "provider", "external_job_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + generation_request_id: Mapped[str] = mapped_column( + String(36), + ForeignKey("generation_requests.id", ondelete="CASCADE"), + nullable=False, + ) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + provider: Mapped[str] = mapped_column(String(64), nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + max_attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=3) + next_attempt_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + external_job_id: Mapped[str | None] = mapped_column(String(255)) + # Safe metadata only. Bearer-like upload URLs or credentials belong in a + # future secret store, never in this row or externally serialised views. + provider_metadata_json: Mapped[dict[str, object]] = mapped_column( + "provider_metadata", JSON, nullable=False, default=dict + ) + output_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="RESTRICT") + ) + error_code: Mapped[str | None] = mapped_column(String(100)) + error_message: Mapped[str | None] = mapped_column(Text) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class GenerationJobAttempt(Base): + """Auditable retry history without creating new logical jobs.""" + + __tablename__ = "generation_job_attempts" + __table_args__ = ( + UniqueConstraint( + "generation_job_id", "attempt_number", name="uq_generation_job_attempt_number" + ), + Index("ix_generation_job_attempts_job", "generation_job_id", "attempt_number"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + generation_job_id: Mapped[str] = mapped_column( + String(36), ForeignKey("generation_jobs.id", ondelete="CASCADE"), nullable=False + ) + attempt_number: Mapped[int] = mapped_column(Integer, nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False) + brand_kit_version_id: Mapped[str | None] = mapped_column(String(36), nullable=True) + external_job_id: Mapped[str | None] = mapped_column(String(128)) + error_message: Mapped[str | None] = mapped_column(Text) + provider_request_id: Mapped[str | None] = mapped_column(String(255)) + started_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) diff --git a/app/generation/providers/__init__.py b/app/generation/providers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a9db5ad576edfd20dad8b04091b732f3ccaafe08 --- /dev/null +++ b/app/generation/providers/__init__.py @@ -0,0 +1,15 @@ +"""Provider contracts and audited optional generation-worker adapters.""" + +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.providers.flux import FluxProviderAdapter +from app.generation.providers.registry import GenerationProviderRegistry +from app.generation.providers.worker_client import RemoteWorkerClient +from app.generation.providers.wan import WanProviderAdapter + +__all__ = [ + "GenerationProviderAdapter", + "FluxProviderAdapter", + "GenerationProviderRegistry", + "RemoteWorkerClient", + "WanProviderAdapter", +] diff --git a/app/generation/providers/base.py b/app/generation/providers/base.py new file mode 100644 index 0000000000000000000000000000000000000000..c59f102b374aac04dd960f85b41fa703a7187d71 --- /dev/null +++ b/app/generation/providers/base.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +from abc import ABC +from collections.abc import AsyncIterator +from contextlib import AbstractAsyncContextManager +from pathlib import Path + +from app.generation.domain.capabilities import GenerationProviderCapabilities +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationProviderUnavailableError, + GenerationWorkerError, +) +from app.generation.domain.runtime import ( + WorkerCancellationResult, + WorkerHealth, + WorkerInfo, + WorkerJob, + WorkerOutput, + WorkerReadiness, +) +from app.generation.schemas.requests import GenerationRequestCreate +from app.security.models import CanonicalMediaAsset + + +class _UnavailableOutputStream(AbstractAsyncContextManager[AsyncIterator[bytes]]): + """Context-manager shape for the default unavailable output stream.""" + + def __init__(self, provider: str) -> None: + self.provider = provider + + async def __aenter__(self) -> AsyncIterator[bytes]: + raise GenerationProviderUnavailableError( + f"{self.provider} output streaming is unavailable." + ) + + async def __aexit__(self, *_: object) -> None: + return None + + +class GenerationProviderAdapter(ABC): + """Provider-neutral worker contract. + + No credential, endpoint, or model-specific logic belongs in REST, MCP, + SDK, n8n, or the generation service. A concrete adapter will be added + only alongside a verified WAN or FLUX worker integration. + """ + + capabilities: GenerationProviderCapabilities + + @property + def provider(self) -> str: + return self.capabilities.provider + + @property + def available(self) -> bool: + """Whether this process can safely accept new work for this adapter.""" + + return False + + def model_capability(self, model_id: str): + for model in self.capabilities.models: + if model.id == model_id: + return model + raise GenerationCapabilityUnsupportedError( + f"{self.provider} does not support model '{model_id}'." + ) + + async def validate_request( + self, payload: GenerationRequestCreate + ) -> dict[str, object]: + """Return a normalized, non-secret worker payload. + + Concrete adapters must validate their own strict Pydantic model and + must not pass through unknown fields. + """ + + raise GenerationCapabilityUnsupportedError( + f"{self.provider} generation is not implemented." + ) + + async def validate_input_asset( + self, payload: GenerationRequestCreate, asset: CanonicalMediaAsset + ) -> None: + """Validate a resolved canonical input descriptor before a job is queued. + + This hook intentionally receives a record owned by the service, never + a client path, URL, or arbitrary upload object. Adapters may impose + MIME/type constraints but byte-level readability is checked again by + the trusted dispatcher immediately before submission. + """ + + del payload, asset + return None + + async def info(self) -> WorkerInfo: + raise GenerationProviderUnavailableError( + f"{self.provider} worker metadata is unavailable." + ) + + async def health(self) -> WorkerHealth: + raise GenerationProviderUnavailableError( + f"{self.provider} worker health is unavailable." + ) + + async def ready(self) -> WorkerReadiness: + raise GenerationProviderUnavailableError( + f"{self.provider} worker readiness is unavailable." + ) + + async def submit( + self, + *, + payload: dict[str, object], + idempotency_key: str, + input_path: Path | None = None, + input_mime_type: str | None = None, + ) -> WorkerJob: + del input_path, input_mime_type + raise GenerationProviderUnavailableError( + f"{self.provider} generation worker is unavailable." + ) + + async def get_job(self, *, external_job_id: str) -> WorkerJob: + raise GenerationProviderUnavailableError( + f"{self.provider} status reconciliation is unavailable." + ) + + async def get_status(self, *, external_job_id: str) -> WorkerJob: + """Compatibility alias; future code should call ``get_job``.""" + + return await self.get_job(external_job_id=external_job_id) + + async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult: + if not self.capabilities.supports_cancellation: + from app.generation.domain.enums import WorkerCancellationStatus + + return WorkerCancellationResult(status=WorkerCancellationStatus.UNSUPPORTED) + raise GenerationProviderUnavailableError( + f"{self.provider} cancellation is unavailable." + ) + + async def retrieve_output(self, *, external_job_id: str) -> WorkerOutput: + raise GenerationProviderUnavailableError( + f"{self.provider} output retrieval is unavailable." + ) + + def stream_output( + self, output: WorkerOutput + ) -> AbstractAsyncContextManager[AsyncIterator[bytes]]: + """Return a scoped byte stream for a validated worker output descriptor.""" + + del output + return _UnavailableOutputStream(self.provider) + + def normalize_error(self, error: Exception) -> GenerationWorkerError | Exception: + """Keep provider error mapping in the adapter, never in transports.""" + + if isinstance(error, GenerationWorkerError): + return error + return error + + async def close(self) -> None: + """Close provider-owned worker clients when the container stops.""" + + return None diff --git a/app/generation/providers/flux.py b/app/generation/providers/flux.py new file mode 100644 index 0000000000000000000000000000000000000000..6e7ec36b7c21d6c70a4a9c015b590fce159e24b8 --- /dev/null +++ b/app/generation/providers/flux.py @@ -0,0 +1,406 @@ +"""FLUX.2 Klein adapter for the trusted MediaRouter worker contract.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator +from contextlib import AbstractAsyncContextManager +from pathlib import Path + +from app.core.config import Settings +from app.generation.domain.capabilities import ( + GenerationModelCapability, + GenerationProviderCapabilities, +) +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerJobStatus, + WorkerReadinessStatus, +) +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationOutputError, + GenerationProviderUnavailableError, + GenerationValidationError, + GenerationWorkerError, +) +from app.generation.domain.runtime import ( + WorkerCancellationResult, + WorkerHealth, + WorkerInfo, + WorkerJob, + WorkerOutput, + WorkerReadiness, + safe_worker_metadata, +) +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.providers.worker_client import RemoteWorkerClient +from app.generation.schemas.requests import FluxGenerationOptions, GenerationRequestCreate +from app.security.models import CanonicalMediaAsset + +FLUX_PROVIDER_ID = "flux" +FLUX_MODEL_ID = "flux.2-klein-4b" +FLUX_MODEL_NAME = "FLUX.2 Klein 4B" +FLUX_LICENSE = "Apache-2.0" +FLUX_TASK = "text-to-image" +FLUX_DISTILLED_MODEL_ID = "black-forest-labs/FLUX.2-klein-4B" +FLUX_BASE_MODEL_ID = "black-forest-labs/FLUX.2-klein-base-4B" +FLUX_INPUT_MAX_BYTES = 20 * 1024 * 1024 +FLUX_INPUT_MIME_PREFIX = "image/" + + +FLUX_MODEL_CAPABILITY = GenerationModelCapability( + id=FLUX_MODEL_ID, + name=FLUX_MODEL_NAME, + modality=GenerationModality.IMAGE, + # The worker also accepts up to four edit inputs. The established public + # generation envelope owns one canonical input asset, so this adapter + # deliberately supports zero or one input without extending that contract. + input_asset_supported=True, + input_schema={ + "input_asset": { + "required": False, + "media_type": "image/*", + "max_bytes": FLUX_INPUT_MAX_BYTES, + "maximum_count": 1, + }, + "prompt": {"required": True, "min_length": 1, "max_length": 4000}, + "mode_choice": { + "required": False, + "enum": ["Distilled (4 steps)", "Base (50 steps)"], + }, + "seed": {"required": False, "minimum": 0, "maximum": 2_147_483_647}, + "randomize_seed": {"required": False, "type": "boolean"}, + "width": {"required": False, "minimum": 256, "maximum": 1024, "multiple_of": 8}, + "height": {"required": False, "minimum": 256, "maximum": 1024, "multiple_of": 8}, + "num_inference_steps": {"required": False, "minimum": 1, "maximum": 100}, + "guidance_scale": {"required": False, "minimum": 0.0, "maximum": 10.0}, + "prompt_upsampling": {"required": False, "type": "boolean"}, + }, +) + + +class FluxProviderAdapter(GenerationProviderAdapter): + """Strict adapter for the audited FLUX.2 Klein asynchronous worker.""" + + capabilities = GenerationProviderCapabilities( + provider=FLUX_PROVIDER_ID, + name="FLUX.2 Klein", + models=[FLUX_MODEL_CAPABILITY], + implementation_status="implemented", + supports_cancellation=True, + supports_status_reconciliation=True, + ) + + def __init__( + self, + *, + client: RemoteWorkerClient | None, + configuration_error: str | None = None, + ) -> None: + self.client = client + self.configuration_error = configuration_error + + @classmethod + def from_settings(cls, settings: Settings) -> "FluxProviderAdapter": + """Build an optional adapter without affecting unrelated startup.""" + + url = settings.flux_space_url.strip() + token = settings.flux_space_token + if not url and token is None: + return cls(client=None) + if not url or token is None: + return cls( + client=None, + configuration_error="FLUX worker URL and token must be configured together.", + ) + secret = token.get_secret_value() + if not 32 <= len(secret) <= 512: + return cls(client=None, configuration_error="FLUX worker token has an invalid length.") + try: + from app.generation.domain.retry import GenerationRetryPolicy + + return cls( + client=RemoteWorkerClient( + base_url=url, + bearer_token=token, + connect_timeout_seconds=settings.ai_worker_connect_timeout_seconds, + request_timeout_seconds=settings.ai_worker_request_timeout_seconds, + read_timeout_seconds=settings.ai_worker_read_timeout_seconds, + retry_policy=GenerationRetryPolicy( + max_retries=settings.ai_worker_max_retries, + backoff_seconds=settings.ai_worker_retry_backoff_seconds, + ), + ) + ) + except (TypeError, ValueError): + return cls(client=None, configuration_error="FLUX worker configuration is invalid.") + + @property + def available(self) -> bool: + return self.client is not None + + async def validate_request( + self, payload: GenerationRequestCreate + ) -> dict[str, object]: + if payload.provider != self.provider or payload.model_id != FLUX_MODEL_ID: + raise GenerationCapabilityUnsupportedError("FLUX request targets an unsupported model.") + if payload.modality is not GenerationModality.IMAGE: + raise GenerationCapabilityUnsupportedError( + "FLUX.2 Klein supports image generation only." + ) + if payload.wan is not None: + raise GenerationCapabilityUnsupportedError( + "WAN controls cannot be supplied to a FLUX generation request." + ) + if not payload.prompt or not payload.prompt.strip() or len(payload.prompt) > 4_000: + raise GenerationValidationError("FLUX prompt must contain 1 through 4000 characters.") + if payload.flux is None: + return {"prompt": payload.prompt} + options = FluxGenerationOptions.model_validate(payload.flux) + return {"prompt": payload.prompt, "flux": options.model_dump(exclude_unset=True)} + + async def validate_input_asset( + self, payload: GenerationRequestCreate, asset: CanonicalMediaAsset + ) -> None: + del payload + mime_type = (asset.mime_type or "").split(";", 1)[0].strip().lower() + if not mime_type.startswith(FLUX_INPUT_MIME_PREFIX): + raise GenerationValidationError("FLUX requires a canonical image input asset.") + if asset.file_size <= 0 or asset.file_size > FLUX_INPUT_MAX_BYTES: + raise GenerationValidationError( + "FLUX input image exceeds the worker's 20 MiB per-image limit." + ) + + async def info(self) -> WorkerInfo: + info = await self._client().info() + variants = self._variants(info) + if ( + info.id != FLUX_MODEL_ID + or GenerationModality.IMAGE not in info.media_types + or info.metadata.get("license") != FLUX_LICENSE + or variants.get("distilled") != FLUX_DISTILLED_MODEL_ID + or variants.get("base") != FLUX_BASE_MODEL_ID + ): + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="Configured FLUX worker identity does not match the registered model.", + retryable=False, + ) + return info + + async def health(self) -> WorkerHealth: + return await self._client().health() + + async def ready(self) -> WorkerReadiness: + readiness = await self._client().ready() + if readiness.metadata.get("accepting_jobs") is not True: + return readiness.model_copy(update={"status": WorkerReadinessStatus.UNAVAILABLE}) + return readiness + + async def submit( + self, + *, + payload: dict[str, object], + idempotency_key: str, + input_path: Path | None = None, + input_mime_type: str | None = None, + ) -> WorkerJob: + fields = self._form_fields(payload) + if input_path is None: + response = await self._client().submit_form( + fields=fields, + idempotency_key=idempotency_key, + # FLUX accepts no idempotency key. A lost submission response + # is handled by the durable dispatcher as an ambiguity. + idempotent=False, + ) + return self._job_from_payload(response, expected_job_id=None) + + if input_mime_type is None: + raise GenerationValidationError("FLUX image input metadata is unavailable.") + mime_type = input_mime_type.split(";", 1)[0].strip().lower() + if not mime_type.startswith(FLUX_INPUT_MIME_PREFIX): + raise GenerationValidationError("FLUX requires an image input asset.") + try: + size = input_path.stat().st_size + except OSError as exc: + raise GenerationValidationError("FLUX input asset is no longer readable.") from exc + if size <= 0 or size > FLUX_INPUT_MAX_BYTES: + raise GenerationValidationError( + "FLUX input image exceeds the worker's 20 MiB per-image limit." + ) + response = await self._client().submit_multipart( + fields=fields, + file_field="input_images", + file_path=input_path, + filename=input_path.name, + mime_type=mime_type, + idempotency_key=idempotency_key, + idempotent=False, + ) + return self._job_from_payload(response, expected_job_id=None) + + async def get_job(self, *, external_job_id: str) -> WorkerJob: + return self._job_from_payload( + await self._client().get_job_payload(external_job_id), + expected_job_id=external_job_id, + ) + + async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult: + payload = await self._client().cancel_job_payload( + external_job_id, expected_statuses={200, 409} + ) + if payload.get("status") == "cancelled": + return WorkerCancellationResult(status=WorkerCancellationStatus.CANCELLED) + detail = payload.get("detail") + if isinstance(detail, dict) and detail.get("code") == "FLUX_JOB_NOT_CANCELLABLE": + return WorkerCancellationResult( + status=WorkerCancellationStatus.FAILED, + metadata={"reason": "job_not_cancellable"}, + ) + return WorkerCancellationResult( + status=WorkerCancellationStatus.FAILED, + metadata={"reason": "unexpected_cancellation_response"}, + ) + + async def retrieve_output(self, *, external_job_id: str) -> WorkerOutput: + job = await self.get_job(external_job_id=external_job_id) + if job.status is not WorkerJobStatus.COMPLETED or job.output is None: + raise GenerationOutputError("FLUX output is not ready.") + return job.output + + def stream_output( + self, output: WorkerOutput + ) -> AbstractAsyncContextManager[AsyncIterator[bytes]]: + if output.output_type is not GenerationModality.IMAGE or output.mime_type != "image/png": + return super().stream_output(output) + return self._client().stream_output(output) + + def normalize_error(self, error: Exception) -> GenerationWorkerError | Exception: + if isinstance(error, GenerationWorkerError): + return error + return GenerationWorkerError( + category=WorkerErrorCategory.UNKNOWN_ERROR, + message="FLUX worker operation failed unexpectedly.", + retryable=False, + ) + + async def close(self) -> None: + if self.client is not None: + await self.client.aclose() + + def _client(self) -> RemoteWorkerClient: + if self.client is None: + raise GenerationProviderUnavailableError("FLUX worker is not configured.") + return self.client + + @staticmethod + def _form_fields(payload: dict[str, object]) -> dict[str, str]: + prompt = payload.get("prompt") + raw_options = payload.get("flux", {}) + if not isinstance(prompt, str) or not prompt.strip() or not isinstance(raw_options, dict): + raise GenerationValidationError("FLUX generation request is invalid.") + try: + options = FluxGenerationOptions.model_validate(raw_options) + except ValueError as exc: + raise GenerationValidationError("FLUX generation options are invalid.") from exc + fields: dict[str, str] = {"prompt": prompt} + for name, value in options.model_dump(exclude_unset=True).items(): + if isinstance(value, bool): + fields[name] = "true" if value else "false" + else: + fields[name] = str(value) + return fields + + @staticmethod + def _variants(info: WorkerInfo) -> dict[str, object]: + matching = next((model for model in info.models if model.id == FLUX_MODEL_ID), None) + if matching is None: + return {} + variants = matching.metadata.get("variants") + return variants if isinstance(variants, dict) else {} + + @staticmethod + def _job_from_payload( + payload: dict[str, object], *, expected_job_id: str | None + ) -> WorkerJob: + job_id = payload.get("job_id") + raw_status = payload.get("status") + if not isinstance(job_id, str) or not isinstance(raw_status, str): + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="FLUX worker returned an invalid job response.", + retryable=False, + ) + if expected_job_id is not None and job_id != expected_job_id: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="FLUX worker returned an unexpected job identity.", + retryable=False, + ) + statuses = { + "queued": WorkerJobStatus.QUEUED, + "running": WorkerJobStatus.RUNNING, + "completed": WorkerJobStatus.COMPLETED, + "failed": WorkerJobStatus.FAILED, + "cancelled": WorkerJobStatus.CANCELLED, + } + state = statuses.get(raw_status.lower()) + if state is None: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="FLUX worker returned an unsupported job status.", + retryable=False, + ) + + output: WorkerOutput | None = None + if state is WorkerJobStatus.COMPLETED: + raw_output = payload.get("output") + if not isinstance(raw_output, dict) or raw_output.get("type") != "image": + raise GenerationOutputError("FLUX worker returned an invalid image output.") + filename = raw_output.get("filename") + if not isinstance(filename, str): + raise GenerationOutputError("FLUX worker returned an invalid image filename.") + output = WorkerOutput( + output_type=GenerationModality.IMAGE, + mime_type="image/png", + provider_output_id=job_id, + download_path=f"/v1/jobs/{job_id}/output", + filename=filename, + ) + + raw_error = payload.get("error") + error_code: str | None = None + error_message: str | None = None + if isinstance(raw_error, dict): + code = raw_error.get("code") + message = raw_error.get("message") + error_code = code if isinstance(code, str) else None + error_message = message if isinstance(message, str) else None + metadata = safe_worker_metadata( + { + key: value + for key, value in payload.items() + if key not in {"job_id", "status", "output", "error"} + } + ) + try: + return WorkerJob( + external_job_id=job_id, + status=state, + output=output, + error_category=( + WorkerErrorCategory.INFERENCE_ERROR if state is WorkerJobStatus.FAILED else None + ), + error_code=error_code, + error_message=error_message, + metadata=metadata if isinstance(metadata, dict) else {}, + ) + except ValueError as exc: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="FLUX worker returned invalid job metadata.", + retryable=False, + ) from exc diff --git a/app/generation/providers/registry.py b/app/generation/providers/registry.py new file mode 100644 index 0000000000000000000000000000000000000000..7ba9671f899c3c2170009cbe29ccc71330b37c13 --- /dev/null +++ b/app/generation/providers/registry.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from app.generation.domain.errors import GenerationProviderUnavailableError +from app.generation.providers.base import GenerationProviderAdapter + + +class GenerationProviderRegistry: + """Registry that exposes only concrete, verified adapters. + + Only fully implemented adapters are supplied by the container. WAN and + FLUX are optional and can be present but unavailable when their + configuration or readiness verification fails. + """ + + def __init__(self, providers: list[GenerationProviderAdapter] | None = None) -> None: + self._providers: dict[str, GenerationProviderAdapter] = {} + for provider in providers or []: + if provider.provider in self._providers: + raise ValueError(f"Duplicate generation provider '{provider.provider}'.") + self._providers[provider.provider] = provider + + def list(self) -> list[GenerationProviderAdapter]: + return [self._providers[key] for key in sorted(self._providers)] + + def get(self, provider: str) -> GenerationProviderAdapter: + key = provider.strip().lower() + try: + return self._providers[key] + except KeyError as exc: + raise GenerationProviderUnavailableError( + f"Generation provider '{key}' is not configured." + ) from exc + + async def close(self) -> None: + for provider in self._providers.values(): + await provider.close() diff --git a/app/generation/providers/wan.py b/app/generation/providers/wan.py new file mode 100644 index 0000000000000000000000000000000000000000..23325a7f6d76c0f2242154fca38ecec878af7bcc --- /dev/null +++ b/app/generation/providers/wan.py @@ -0,0 +1,416 @@ +"""WAN 2.2 image-to-video adapter for the trusted MediaRouter worker API. + +The worker protocol in this module was audited against the companion WAN +Space. This adapter deliberately knows only that protocol; orchestration, +tenancy, durable jobs, output storage, and public transport remain owned by +the provider-neutral generation services. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator +from contextlib import AbstractAsyncContextManager +from pathlib import Path + +from app.core.config import Settings +from app.generation.domain.capabilities import ( + GenerationModelCapability, + GenerationProviderCapabilities, +) +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerJobStatus, + WorkerReadinessStatus, +) +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationOutputError, + GenerationProviderUnavailableError, + GenerationValidationError, + GenerationWorkerError, +) +from app.generation.domain.runtime import ( + WorkerCancellationResult, + WorkerHealth, + WorkerInfo, + WorkerJob, + WorkerOutput, + WorkerReadiness, + safe_worker_metadata, +) +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.providers.worker_client import RemoteWorkerClient +from app.generation.schemas.requests import GenerationRequestCreate, WanGenerationOptions +from app.security.models import CanonicalMediaAsset + +WAN_PROVIDER_ID = "wan" +WAN_MODEL_ID = "wan2.2" +WAN_MODEL_NAME = "WAN 2.2 FP8 AOTI Faster" +WAN_UNDERLYING_MODEL_ID = "Wan-AI/Wan2.2-I2V-A14B-Diffusers" +WAN_INPUT_MAX_BYTES = 20 * 1024 * 1024 +WAN_INPUT_MIME_PREFIX = "image/" + + +WAN_MODEL_CAPABILITY = GenerationModelCapability( + id=WAN_MODEL_ID, + name=WAN_MODEL_NAME, + modality=GenerationModality.VIDEO, + input_asset_supported=True, + input_schema={ + "input_asset": { + "required": True, + "media_type": "image/*", + "max_bytes": WAN_INPUT_MAX_BYTES, + }, + "prompt": {"required": True, "min_length": 1, "max_length": 4000}, + "negative_prompt": {"required": False, "max_length": 4000}, + "duration_seconds": {"required": False, "minimum": 0.5, "maximum": 5.0}, + "steps": {"required": False, "minimum": 1, "maximum": 30}, + "guidance_scale": {"required": False, "minimum": 0.0, "maximum": 10.0}, + "guidance_scale_2": {"required": False, "minimum": 0.0, "maximum": 10.0}, + "seed": {"required": False, "minimum": 0, "maximum": 2_147_483_647}, + "randomize_seed": {"required": False, "type": "boolean"}, + "derived": { + "dimensions": "source image, resized by worker to its supported bounds", + "frames": "duration_seconds at worker-reported 16 fps", + }, + }, +) + + +class WanProviderAdapter(GenerationProviderAdapter): + """Strict adapter for the audited, authenticated WAN worker contract.""" + + capabilities = GenerationProviderCapabilities( + provider=WAN_PROVIDER_ID, + name="WAN 2.2", + models=[WAN_MODEL_CAPABILITY], + implementation_status="implemented", + supports_cancellation=True, + supports_status_reconciliation=True, + ) + + def __init__( + self, + *, + client: RemoteWorkerClient | None, + configuration_error: str | None = None, + ) -> None: + self.client = client + self.configuration_error = configuration_error + + @classmethod + def from_settings(cls, settings: Settings) -> "WanProviderAdapter": + """Build an optional adapter without allowing a bad config to abort startup.""" + + url = settings.wan_space_url.strip() + token = settings.wan_space_token + if not url and token is None: + return cls(client=None) + if not url or token is None: + return cls( + client=None, + configuration_error="WAN worker URL and token must be configured together.", + ) + secret = token.get_secret_value() + if not 32 <= len(secret) <= 4096: + return cls( + client=None, + configuration_error="WAN worker token has an invalid length.", + ) + try: + from app.generation.domain.retry import GenerationRetryPolicy + + return cls( + client=RemoteWorkerClient( + base_url=url, + bearer_token=token, + connect_timeout_seconds=settings.ai_worker_connect_timeout_seconds, + request_timeout_seconds=settings.ai_worker_request_timeout_seconds, + read_timeout_seconds=settings.ai_worker_read_timeout_seconds, + retry_policy=GenerationRetryPolicy( + max_retries=settings.ai_worker_max_retries, + backoff_seconds=settings.ai_worker_retry_backoff_seconds, + ), + ) + ) + except (TypeError, ValueError): + # Never include the configured URL or token in a startup error or + # log. An operator can correct configuration without affecting + # the non-generation application surface. + return cls(client=None, configuration_error="WAN worker configuration is invalid.") + + @property + def available(self) -> bool: + # Availability at the model level still requires a successful health, + # readiness and exact-info verification in GenerationModelRegistry. + return self.client is not None + + async def validate_request( + self, payload: GenerationRequestCreate + ) -> dict[str, object]: + if payload.provider != self.provider or payload.model_id != WAN_MODEL_ID: + raise GenerationCapabilityUnsupportedError("WAN request targets an unsupported model.") + if payload.modality is not GenerationModality.VIDEO: + raise GenerationCapabilityUnsupportedError("WAN 2.2 supports video generation only.") + if payload.input_asset_id is None: + raise GenerationValidationError("WAN 2.2 requires an image input asset.") + if not payload.prompt or not payload.prompt.strip() or len(payload.prompt) > 4_000: + raise GenerationValidationError("WAN prompt must contain 1 through 4000 characters.") + if payload.wan is None: + # Preserve the worker's audited defaults when no optional control + # was requested; do not persist an invented empty provider blob. + return {"prompt": payload.prompt} + # Validate again at the adapter boundary so internal callers cannot + # hand the worker an arbitrary model-dumped object. + options = WanGenerationOptions.model_validate(payload.wan) + return {"prompt": payload.prompt, "wan": options.model_dump(exclude_unset=True)} + + async def validate_input_asset( + self, payload: GenerationRequestCreate, asset: CanonicalMediaAsset + ) -> None: + del payload + mime_type = (asset.mime_type or "").split(";", 1)[0].strip().lower() + if not mime_type.startswith(WAN_INPUT_MIME_PREFIX): + raise GenerationValidationError("WAN 2.2 requires a canonical image input asset.") + if asset.file_size <= 0 or asset.file_size > WAN_INPUT_MAX_BYTES: + raise GenerationValidationError( + "WAN 2.2 input image exceeds the worker's 20 MiB limit." + ) + + async def info(self) -> WorkerInfo: + info = await self._client().info() + # The audited WAN worker is a single-model worker whose top-level + # identity is the public model identifier. Do not accept a different + # worker that merely happens to list ``wan2.2`` in a secondary model + # collection: it could expose different preprocessing, output, or + # cancellation semantics than this adapter has been reviewed for. + if ( + info.id != WAN_MODEL_ID + or GenerationModality.VIDEO not in info.media_types + or info.metadata.get("task") != "image-to-video" + or info.metadata.get("model_id") != WAN_UNDERLYING_MODEL_ID + or info.metadata.get("fps") != 16 + ): + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="Configured WAN worker identity does not match the registered model.", + retryable=False, + ) + return info + + async def health(self) -> WorkerHealth: + return await self._client().health() + + async def ready(self) -> WorkerReadiness: + readiness = await self._client().ready() + accepting_jobs = readiness.metadata.get("accepting_jobs") + # The audited worker returns 503/"not_ready" while it is not ready, + # but retain this extra guard for malformed or future responses. + if accepting_jobs is not True: + return readiness.model_copy( + update={"status": WorkerReadinessStatus.UNAVAILABLE} + ) + return readiness + + async def submit( + self, + *, + payload: dict[str, object], + idempotency_key: str, + input_path: Path | None = None, + input_mime_type: str | None = None, + ) -> WorkerJob: + if input_path is None or input_mime_type is None: + raise GenerationValidationError("WAN generation requires a verified image input asset.") + mime_type = input_mime_type.split(";", 1)[0].strip().lower() + if not mime_type.startswith(WAN_INPUT_MIME_PREFIX): + raise GenerationValidationError("WAN generation requires an image input asset.") + try: + size = input_path.stat().st_size + except OSError as exc: + raise GenerationValidationError("WAN input asset is no longer readable.") from exc + if size <= 0 or size > WAN_INPUT_MAX_BYTES: + raise GenerationValidationError("WAN input image exceeds the worker's 20 MiB limit.") + + form = self._multipart_fields(payload) + # The audited worker has no idempotency-key protocol. Include the + # canonical request ID for correlation only and forbid transport + # retries: a lost response is intentionally handled as ambiguous by + # the durable dispatcher rather than creating a duplicate video. + response = await self._client().submit_multipart( + fields=form, + file_field="image", + file_path=input_path, + filename=input_path.name, + mime_type=mime_type, + idempotency_key=idempotency_key, + idempotent=False, + ) + return self._job_from_payload(response, expected_job_id=None) + + async def get_job(self, *, external_job_id: str) -> WorkerJob: + return self._job_from_payload( + await self._client().get_job_payload(external_job_id), + expected_job_id=external_job_id, + ) + + async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult: + payload = await self._client().cancel_job_payload( + external_job_id, expected_statuses={200, 409} + ) + if payload.get("status") == "cancelled": + return WorkerCancellationResult(status=WorkerCancellationStatus.CANCELLED) + detail = payload.get("detail") + if isinstance(detail, dict) and detail.get("code") == "WAN_JOB_NOT_CANCELLABLE": + # The worker explicitly says a running GPU operation was not + # stopped. This is a cancellation failure, not a success or a + # provider-wide lack of cancellation capability. + return WorkerCancellationResult( + status=WorkerCancellationStatus.FAILED, + metadata={"reason": "job_not_cancellable"}, + ) + return WorkerCancellationResult( + status=WorkerCancellationStatus.FAILED, + metadata={"reason": "unexpected_cancellation_response"}, + ) + + async def retrieve_output(self, *, external_job_id: str) -> WorkerOutput: + job = await self.get_job(external_job_id=external_job_id) + if job.status is not WorkerJobStatus.COMPLETED or job.output is None: + raise GenerationOutputError("WAN output is not ready.") + return job.output + + def stream_output( + self, output: WorkerOutput + ) -> AbstractAsyncContextManager[AsyncIterator[bytes]]: + if output.output_type is not GenerationModality.VIDEO or output.mime_type != "video/mp4": + return super().stream_output(output) + return self._client().stream_output(output) + + def normalize_error(self, error: Exception) -> GenerationWorkerError | Exception: + if isinstance(error, GenerationWorkerError): + return error + return GenerationWorkerError( + category=WorkerErrorCategory.UNKNOWN_ERROR, + message="WAN worker operation failed unexpectedly.", + retryable=False, + ) + + async def close(self) -> None: + if self.client is not None: + await self.client.aclose() + + def _client(self) -> RemoteWorkerClient: + if self.client is None: + raise GenerationProviderUnavailableError("WAN worker is not configured.") + return self.client + + @staticmethod + def _multipart_fields(payload: dict[str, object]) -> dict[str, str]: + prompt = payload.get("prompt") + raw_options = payload.get("wan", {}) + if not isinstance(prompt, str) or not prompt.strip() or not isinstance(raw_options, dict): + raise GenerationValidationError("WAN generation request is invalid.") + try: + options = WanGenerationOptions.model_validate(raw_options) + except ValueError as exc: + raise GenerationValidationError("WAN generation options are invalid.") from exc + # Persisted specs use ``exclude_unset``. Sending only those fields + # preserves WAN's own documented defaults for omitted controls. + fields: dict[str, str] = {"prompt": prompt} + for name, value in options.model_dump(exclude_unset=True).items(): + if value is None: + continue + if isinstance(value, bool): + fields[name] = "true" if value else "false" + else: + fields[name] = str(value) + return fields + + @staticmethod + def _job_from_payload( + payload: dict[str, object], *, expected_job_id: str | None + ) -> WorkerJob: + job_id = payload.get("job_id") + raw_status = payload.get("status") + if not isinstance(job_id, str) or not isinstance(raw_status, str): + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="WAN worker returned an invalid job response.", + retryable=False, + ) + if expected_job_id is not None and job_id != expected_job_id: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="WAN worker returned an unexpected job identity.", + retryable=False, + ) + statuses = { + "queued": WorkerJobStatus.QUEUED, + "running": WorkerJobStatus.RUNNING, + "completed": WorkerJobStatus.COMPLETED, + "failed": WorkerJobStatus.FAILED, + "cancelled": WorkerJobStatus.CANCELLED, + } + state = statuses.get(raw_status.lower()) + if state is None: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="WAN worker returned an unsupported job status.", + retryable=False, + ) + + output: WorkerOutput | None = None + if state is WorkerJobStatus.COMPLETED: + raw_output = payload.get("output") + if not isinstance(raw_output, dict) or raw_output.get("type") != "video": + raise GenerationOutputError("WAN worker returned an invalid video output.") + filename = raw_output.get("filename") + if not isinstance(filename, str): + raise GenerationOutputError("WAN worker returned an invalid video filename.") + output = WorkerOutput( + output_type=GenerationModality.VIDEO, + mime_type="video/mp4", + provider_output_id=job_id, + download_path=f"/v1/jobs/{job_id}/output", + filename=filename, + ) + + error_code: str | None = None + error_message: str | None = None + raw_error = payload.get("error") + if isinstance(raw_error, dict): + code = raw_error.get("code") + message = raw_error.get("message") + error_code = code if isinstance(code, str) else None + error_message = message if isinstance(message, str) else None + metadata = safe_worker_metadata( + { + key: value + for key, value in payload.items() + if key not in {"job_id", "status", "output", "error"} + } + ) + try: + return WorkerJob( + external_job_id=job_id, + status=state, + output=output, + error_category=( + WorkerErrorCategory.INFERENCE_ERROR + if state is WorkerJobStatus.FAILED + else None + ), + error_code=error_code, + error_message=error_message, + metadata=metadata if isinstance(metadata, dict) else {}, + ) + except ValueError as exc: + raise GenerationWorkerError( + category=WorkerErrorCategory.PROVIDER_ERROR, + message="WAN worker returned invalid job metadata.", + retryable=False, + ) from exc diff --git a/app/generation/providers/worker_client.py b/app/generation/providers/worker_client.py new file mode 100644 index 0000000000000000000000000000000000000000..ace752b3af511ae145d8457e03978abecf246c6a --- /dev/null +++ b/app/generation/providers/worker_client.py @@ -0,0 +1,843 @@ +from __future__ import annotations + +import asyncio +import ipaddress +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Any +from urllib.parse import urlparse + +import httpx +from pydantic import SecretStr + +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerHealthStatus, + WorkerJobStatus, + WorkerReadinessStatus, +) +from app.generation.domain.errors import GenerationOutputError, GenerationWorkerError +from app.generation.domain.retry import GenerationRetryPolicy +from app.generation.domain.runtime import ( + WorkerCancellationResult, + WorkerHealth, + WorkerInfo, + WorkerJob, + WorkerModelInfo, + WorkerOutput, + WorkerReadiness, + safe_worker_metadata, +) + + +class RemoteWorkerClient: + """Strict HTTP client for a trusted, configured MediaRouter worker. + + The constructor is deliberately internal-facing: no REST, MCP, SDK, n8n, + or browser payload may supply its base URL or token. Redirects and proxy + environment variables are disabled, endpoint paths are fixed/validated, + and no request/response body is logged. + """ + + def __init__( + self, + *, + base_url: str, + bearer_token: SecretStr | str | None, + connect_timeout_seconds: float, + request_timeout_seconds: float, + read_timeout_seconds: float, + retry_policy: GenerationRetryPolicy, + http_client: httpx.AsyncClient | None = None, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + ) -> None: + self.base_url = self._validate_base_url(base_url) + self._bearer_token = ( + bearer_token.get_secret_value() + if isinstance(bearer_token, SecretStr) + else bearer_token + ) + self._request_timeout_seconds = request_timeout_seconds + self.retry_policy = retry_policy + self._sleep = sleep + self._client = http_client or httpx.AsyncClient( + timeout=httpx.Timeout( + connect=connect_timeout_seconds, + read=read_timeout_seconds, + write=request_timeout_seconds, + pool=connect_timeout_seconds, + ), + follow_redirects=False, + trust_env=False, + headers={"User-Agent": "mediarouter-generation-runtime/1"}, + ) + self._owns_client = http_client is None + + async def aclose(self) -> None: + if self._owns_client: + await self._client.aclose() + + async def health(self) -> WorkerHealth: + payload = await self._request_json("GET", "/health", idempotent=True) + return WorkerHealth( + status=self._health_status(payload.get("status")), + metadata=self._metadata(payload, known={"status"}), + ) + + async def ready(self) -> WorkerReadiness: + payload = await self._request_json( + "GET", "/ready", idempotent=True, readiness_endpoint=True + ) + model_ids = payload.get("model_ids") + if model_ids is None: + model = payload.get("model") + if model is None: + model_ids = [] + elif isinstance(model, str): + model_ids = [model] + else: + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned invalid readiness model IDs.", + ) + elif not isinstance(model_ids, list) or not all( + isinstance(value, str) for value in model_ids + ): + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned invalid readiness model IDs.", + ) + model_loaded = payload.get("model_loaded", False) + if not isinstance(model_loaded, bool): + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned invalid readiness metadata.", + ) + return WorkerReadiness( + status=self._readiness_status(payload.get("status")), + model_loaded=model_loaded, + model_ids=model_ids, + metadata=self._metadata( + payload, known={"status", "model_loaded", "model_ids", "model"} + ), + ) + + async def info(self) -> WorkerInfo: + payload = await self._request_json("GET", "/v1/info", idempotent=True) + identifier = payload.get("id") + name = payload.get("name") + if ( + not isinstance(identifier, str) + or not identifier + or not isinstance(name, str) + or not name + ): + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned invalid model metadata.", + ) + try: + media_types = self._media_types(payload) + models = self._worker_models(payload, fallback_id=identifier, fallback_name=name) + except ValueError as exc: + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned an unsupported media type.", + ) from exc + return WorkerInfo( + id=identifier, + name=name, + media_types=list(dict.fromkeys(media_types)), + models=models, + status=self._health_status(payload.get("status")), + metadata=self._metadata( + payload, + known={"id", "name", "type", "media_types", "models", "status"}, + ), + ) + + async def submit( + self, *, payload: dict[str, object], idempotency_key: str + ) -> WorkerJob: + if not idempotency_key.strip(): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation submission requires an idempotency key.", + ) + response = await self._request_json( + "POST", + "/v1/generate", + json_payload=payload, + headers={"Idempotency-Key": idempotency_key}, + idempotent=True, + expected_statuses={200, 202}, + ) + return self._worker_job(response) + + async def submit_form( + self, + *, + fields: dict[str, str], + idempotency_key: str, + idempotent: bool, + ) -> dict[str, object]: + """Submit a strict scalar form to a worker that does not need a file. + + Some workers use ``multipart/form-data`` only when an optional input + asset is supplied, but still require form fields for text-only work. + This provider-neutral primitive keeps that transport detail out of + adapters without falling back to an incompatible JSON request body. + ``idempotent`` remains explicit because workers can accept a job + without exposing any request-idempotency protocol. + """ + + if not idempotency_key.strip(): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation submission requires an idempotency key.", + ) + if not fields or any( + not isinstance(key, str) or not isinstance(value, str) + for key, value in fields.items() + ): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation submission fields are invalid.", + ) + return await self._request_json( + "POST", + "/v1/generate", + data=fields, + headers={"Idempotency-Key": idempotency_key}, + idempotent=idempotent, + expected_statuses={200, 202}, + ) + + async def submit_multipart( + self, + *, + fields: dict[str, str], + file_field: str, + file_path: Path, + filename: str, + mime_type: str, + idempotency_key: str, + idempotent: bool, + ) -> dict[str, object]: + """Submit one canonical local input file as multipart data. + + This is a transport primitive rather than a model-specific API. The + source file is opened by the server from a verified canonical asset; + it is never a client filesystem path. ``httpx`` streams the file + object while encoding multipart data, so large media is not loaded + into memory. A caller may opt out of automatic retries when its + worker does not offer submission idempotency. + """ + + if not idempotency_key.strip(): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation submission requires an idempotency key.", + ) + if not file_field or not filename or not mime_type: + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation submission file metadata is invalid.", + ) + source_input = file_path.expanduser() + if source_input.is_symlink(): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation input asset is unavailable.", + ) + source = source_input.resolve() + if not source.is_file(): + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation input asset is unavailable.", + ) + try: + # The WAN worker intentionally has no idempotency key support. + # Its adapter sets ``idempotent=False``, preventing an uncertain + # network failure from causing a second expensive GPU submission. + with source.open("rb") as stream: + return await self._request_json( + "POST", + "/v1/generate", + data=fields, + files={file_field: (filename, stream, mime_type)}, + headers={"Idempotency-Key": idempotency_key}, + idempotent=idempotent, + expected_statuses={200, 202}, + ) + except GenerationWorkerError: + raise + except OSError as exc: + raise self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation input asset could not be read.", + ) from exc + + async def get_job_payload(self, external_job_id: str) -> dict[str, object]: + """Return a fixed worker job response for adapter-specific parsing.""" + + path = f"/v1/jobs/{self._safe_external_id(external_job_id)}" + return await self._request_json("GET", path, idempotent=True, job_endpoint=True) + + async def cancel_job_payload( + self, + external_job_id: str, + *, + expected_statuses: set[int] | None = None, + ) -> dict[str, object]: + """Call the fixed cancellation endpoint without provider parsing.""" + + path = f"/v1/jobs/{self._safe_external_id(external_job_id)}/cancel" + return await self._request_json( + "POST", + path, + idempotent=True, + expected_statuses=expected_statuses or {200, 202, 204}, + job_endpoint=True, + ) + + async def get_job(self, external_job_id: str) -> WorkerJob: + return self._worker_job(await self.get_job_payload(external_job_id)) + + async def cancel(self, external_job_id: str) -> WorkerCancellationResult: + payload = await self.cancel_job_payload(external_job_id) + raw_status = str(payload.get("status", "")).strip().lower() + if raw_status in {"cancelled", "canceled"}: + status = WorkerCancellationStatus.CANCELLED + elif raw_status in {"requested", "cancel_requested", "cancellation_requested"}: + status = WorkerCancellationStatus.REQUESTED + elif raw_status in {"unsupported", "not_supported"}: + status = WorkerCancellationStatus.UNSUPPORTED + elif not raw_status: + # A successful 204 has no representation of whether a running + # GPU operation actually stopped. It can only mean the worker + # accepted the cancellation request, never that it completed it. + status = WorkerCancellationStatus.REQUESTED + else: + status = WorkerCancellationStatus.FAILED + return WorkerCancellationResult( + status=status, + metadata=self._metadata(payload, known={"status"}), + ) + + async def retrieve_output(self, external_job_id: str) -> WorkerOutput: + job = await self.get_job(external_job_id) + if job.status is not WorkerJobStatus.COMPLETED or job.output is None: + raise GenerationOutputError("Generation output is not ready.") + return job.output + + @asynccontextmanager + async def stream_output(self, output: WorkerOutput) -> AsyncIterator[AsyncIterator[bytes]]: + """Stream a worker-owned relative output path without buffering it. + + The caller must write into a controlled MediaRouter staging location, + verify the optional checksum, then register it through + ``CanonicalAssetService``. No worker filesystem path is ever trusted. + """ + + path = self._safe_worker_path(output.download_path) + context, response = await self._open_stream(path) + try: + yield response.aiter_bytes() + finally: + await context.__aexit__(None, None, None) + + async def _request_json( + self, + method: str, + path: str, + *, + json_payload: dict[str, object] | None = None, + data: dict[str, str] | None = None, + files: Any | None = None, + headers: dict[str, str] | None = None, + idempotent: bool, + expected_statuses: set[int] | None = None, + readiness_endpoint: bool = False, + job_endpoint: bool = False, + ) -> dict[str, object]: + expected = expected_statuses or {200} + safe_path = self._safe_worker_path(path) + retry_number = 0 + while True: + try: + response = await asyncio.wait_for( + self._client.request( + method, + self._url_for(safe_path), + headers=self._headers(headers), + json=json_payload, + data=data, + files=files, + ), + timeout=self._request_timeout_seconds, + ) + if response.status_code not in expected: + raise self._response_error( + response.status_code, + readiness_endpoint=readiness_endpoint, + job_endpoint=job_endpoint, + ) + try: + payload = response.json() if response.content else {} + except (ValueError, UnicodeDecodeError) as exc: + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned an invalid JSON response.", + ) from exc + if not isinstance(payload, dict): + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned an invalid response shape.", + ) + return payload + except asyncio.CancelledError: + raise + except Exception as exc: + error = self._normalise_exception( + exc, + readiness_endpoint=readiness_endpoint, + job_endpoint=job_endpoint, + ) + decision = self.retry_policy.decide( + category=error.category, + http_status=error.http_status, + retry_number=retry_number, + idempotent=idempotent, + ) + if not decision.retryable: + raise error from None + retry_number += 1 + await self._sleep(decision.delay_seconds) + + async def _open_stream( + self, path: str + ) -> tuple[Any, httpx.Response]: + # httpx exposes stream() as an async context manager. It is kept + # private to this method so all callers close it in a finally block. + retry_number = 0 + while True: + context = self._client.stream( + "GET", self._url_for(path), headers=self._headers(None) + ) + try: + response = await asyncio.wait_for( + context.__aenter__(), timeout=self._request_timeout_seconds + ) + if response.status_code != 200: + raise self._response_error(response.status_code) + return context, response + except asyncio.CancelledError: + await context.__aexit__(None, None, None) + raise + except Exception as exc: + await context.__aexit__(None, None, None) + error = self._normalise_exception(exc) + decision = self.retry_policy.decide( + category=error.category, + http_status=error.http_status, + retry_number=retry_number, + idempotent=True, + ) + if not decision.retryable: + raise error from None + retry_number += 1 + await self._sleep(decision.delay_seconds) + + def _worker_job(self, payload: dict[str, object]) -> WorkerJob: + raw_status = str(payload.get("status", "")).strip().lower() + statuses = { + "queued": WorkerJobStatus.QUEUED, + "running": WorkerJobStatus.RUNNING, + "processing": WorkerJobStatus.RUNNING, + "completed": WorkerJobStatus.COMPLETED, + "succeeded": WorkerJobStatus.COMPLETED, + "failed": WorkerJobStatus.FAILED, + "cancelled": WorkerJobStatus.CANCELLED, + "canceled": WorkerJobStatus.CANCELLED, + } + job_id = payload.get("job_id", payload.get("external_job_id")) + if not isinstance(job_id, str) or raw_status not in statuses: + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned an invalid job response.", + ) + raw_output = payload.get("output") + try: + output = self._worker_output(raw_output) if isinstance(raw_output, dict) else None + return WorkerJob( + external_job_id=job_id, + status=statuses[raw_status], + output=output, + error_category=self._worker_error_category(payload, statuses[raw_status]), + error_code=self._worker_error_code(payload), + error_message=None, + metadata=self._metadata( + payload, + known={"job_id", "external_job_id", "status", "output", "error"}, + ), + ) + except (TypeError, ValueError) as exc: + raise self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned invalid job metadata.", + ) from exc + + def _worker_output(self, payload: dict[str, object]) -> WorkerOutput: + type_value = payload.get("output_type", payload.get("type")) + mime_type = payload.get("mime_type") + provider_output_id = payload.get("provider_output_id", payload.get("id")) + download_path = payload.get("download_path") + if not all( + isinstance(value, str) + for value in (type_value, mime_type, provider_output_id, download_path) + ): + raise GenerationOutputError() + filename = payload.get("filename") + sha256 = payload.get("sha256") + byte_size = payload.get("byte_size") + if filename is not None and not isinstance(filename, str): + raise GenerationOutputError() + if sha256 is not None and not isinstance(sha256, str): + raise GenerationOutputError() + if byte_size is not None and ( + not isinstance(byte_size, int) or isinstance(byte_size, bool) + ): + raise GenerationOutputError() + metadata = self._metadata( + payload, + known={ + "output_type", + "type", + "mime_type", + "provider_output_id", + "id", + "download_path", + "filename", + "sha256", + "byte_size", + }, + ) + return WorkerOutput( + output_type=type_value, + mime_type=mime_type, + provider_output_id=provider_output_id, + download_path=self._safe_worker_path(download_path), + filename=filename, + sha256=sha256, + byte_size=byte_size, + metadata=metadata, + ) + + @staticmethod + def _media_types(payload: dict[str, object]) -> list[GenerationModality]: + values = payload.get("media_types") + if values is None: + legacy_type = payload.get("type") + if legacy_type is None: + values = [] + elif isinstance(legacy_type, str): + values = [legacy_type] + else: + raise ValueError("worker media type is invalid") + elif not isinstance(values, list) or not all(isinstance(value, str) for value in values): + raise ValueError("worker media types are invalid") + return [GenerationModality(value) for value in values] + + def _worker_models( + self, payload: dict[str, object], *, fallback_id: str, fallback_name: str + ) -> list[WorkerModelInfo]: + raw_models = payload.get("models") + if raw_models is None: + return [ + WorkerModelInfo( + id=fallback_id, + name=fallback_name, + media_types=list(dict.fromkeys(self._media_types(payload))), + ) + ] + if isinstance(raw_models, dict): + # A compact single-model worker may expose named model variants as + # a JSON object instead of a list of independently selectable + # models. It is still one discovered top-level model; retain the + # safe variant map as metadata for the concrete adapter to verify. + variants = safe_worker_metadata(raw_models) + if not isinstance(variants, dict): + raise ValueError("worker model variants must be an object") + return [ + WorkerModelInfo( + id=fallback_id, + name=fallback_name, + media_types=list(dict.fromkeys(self._media_types(payload))), + metadata={"variants": variants}, + ) + ] + if not isinstance(raw_models, list) or not raw_models: + raise ValueError("worker models must be a non-empty list") + models: list[WorkerModelInfo] = [] + for raw_model in raw_models: + if not isinstance(raw_model, dict): + raise ValueError("worker model must be an object") + model_id = raw_model.get("id") + model_name = raw_model.get("name") + if not isinstance(model_id, str) or not model_id: + raise ValueError("worker model ID is invalid") + if not isinstance(model_name, str) or not model_name: + raise ValueError("worker model name is invalid") + models.append( + WorkerModelInfo( + id=model_id, + name=model_name, + media_types=list(dict.fromkeys(self._media_types(raw_model))), + metadata=self._metadata( + raw_model, known={"id", "name", "type", "media_types"} + ), + ) + ) + return models + + @staticmethod + def _worker_error_code(payload: dict[str, object]) -> str | None: + raw_error = payload.get("error") + if not isinstance(raw_error, dict): + return None + code = raw_error.get("code") + return code if isinstance(code, str) and len(code) <= 100 else None + + @staticmethod + def _worker_error_category( + payload: dict[str, object], status: WorkerJobStatus + ) -> WorkerErrorCategory | None: + raw_error = payload.get("error") + if isinstance(raw_error, dict): + category = raw_error.get("category") + if isinstance(category, str): + try: + return WorkerErrorCategory(category) + except ValueError: + pass + # A terminal worker failure is an inference failure unless the worker + # explicitly supplied a supported, non-secret category. + return WorkerErrorCategory.INFERENCE_ERROR if status is WorkerJobStatus.FAILED else None + + def _normalise_exception( + self, + exc: Exception, + *, + readiness_endpoint: bool = False, + job_endpoint: bool = False, + ) -> GenerationWorkerError: + if isinstance(exc, GenerationWorkerError): + return exc + if isinstance(exc, asyncio.TimeoutError) or isinstance( + exc, (httpx.ReadTimeout, httpx.WriteTimeout, httpx.PoolTimeout) + ): + return self._error(WorkerErrorCategory.TIMEOUT, "Generation worker request timed out.") + if isinstance(exc, (httpx.ConnectTimeout, httpx.NetworkError)): + return self._error( + WorkerErrorCategory.WORKER_UNAVAILABLE, + "Generation worker is unavailable.", + ) + if isinstance(exc, httpx.RequestError): + return self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker request failed.", + ) + del readiness_endpoint, job_endpoint + # Never propagate implementation exception text: httpx exceptions can + # include request URLs and caller implementations can include secrets. + return self._error( + WorkerErrorCategory.UNKNOWN_ERROR, + "Generation worker operation failed unexpectedly.", + ) + + def _response_error( + self, + status_code: int, + *, + readiness_endpoint: bool = False, + job_endpoint: bool = False, + ) -> GenerationWorkerError: + if status_code in {400, 422}: + return self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation worker rejected the request.", + http_status=status_code, + ) + if status_code == 401: + return self._error( + WorkerErrorCategory.AUTHENTICATION_ERROR, + "Generation worker authentication failed.", + http_status=status_code, + ) + if status_code == 403: + return self._error( + WorkerErrorCategory.AUTHORIZATION_ERROR, + "Generation worker authorization failed.", + http_status=status_code, + ) + if status_code == 404 and job_endpoint: + return self._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation worker job was not found.", + http_status=status_code, + ) + if status_code == 429: + return self._error( + WorkerErrorCategory.RATE_LIMITED, + "Generation worker is rate limited.", + http_status=status_code, + ) + if status_code == 503 and readiness_endpoint: + return self._error( + WorkerErrorCategory.WORKER_NOT_READY, + "Generation worker is not ready.", + http_status=status_code, + ) + if status_code in {502, 503, 504}: + return self._error( + WorkerErrorCategory.WORKER_UNAVAILABLE, + "Generation worker is temporarily unavailable.", + http_status=status_code, + ) + if status_code >= 500: + return self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker failed to process the request.", + http_status=status_code, + ) + return self._error( + WorkerErrorCategory.PROVIDER_ERROR, + "Generation worker returned an unsupported response.", + http_status=status_code, + ) + + @staticmethod + def _error( + category: WorkerErrorCategory, + message: str, + *, + http_status: int | None = None, + ) -> GenerationWorkerError: + return GenerationWorkerError( + category=category, + message=message, + retryable=False, + http_status=http_status, + ) + + def _headers(self, extra: dict[str, str] | None) -> dict[str, str]: + headers = dict(extra or {}) + if self._bearer_token: + headers["Authorization"] = f"Bearer {self._bearer_token}" + return headers + + def _url_for(self, path: str) -> str: + return f"{self.base_url}{path}" + + @staticmethod + def _safe_external_id(value: str) -> str: + if not value or len(value) > 255 or any( + character not in "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789._:-" + for character in value + ): + raise RemoteWorkerClient._error( + WorkerErrorCategory.INVALID_REQUEST, + "Generation worker job identifier is invalid.", + ) + return value + + @staticmethod + def _safe_worker_path(value: str) -> str: + parsed = urlparse(value) + if ( + not value.startswith("/") + or parsed.scheme + or parsed.netloc + or parsed.query + or parsed.fragment + or "\\" in value + or "%" in value + or "//" in value + or any(part in {"", ".", ".."} for part in value.split("/")[1:]) + ): + raise GenerationOutputError("Generation worker returned an unsafe endpoint path.") + return value + + @staticmethod + def _validate_base_url(value: str) -> str: + parsed = urlparse(value.strip()) + if ( + parsed.scheme not in {"https", "http"} + or not parsed.hostname + or parsed.username + or parsed.password + or parsed.query + or parsed.fragment + ): + raise ValueError("Generation worker base URL must be an absolute HTTP(S) origin.") + host = parsed.hostname.lower() + try: + address = ipaddress.ip_address(host) + except ValueError: + address = None + is_loopback = host == "localhost" or (address is not None and address.is_loopback) + if address is not None and not address.is_global and not is_loopback: + raise ValueError("Generation worker base URL uses a prohibited address.") + if parsed.scheme == "http" and not is_loopback: + raise ValueError("Generation workers require HTTPS outside local development.") + path = parsed.path.rstrip("/") + if path and ( + "\\" in path + or "%" in path + or "//" in path + or any(part in {"", ".", ".."} for part in path.split("/")[1:]) + ): + raise ValueError("Generation worker base URL contains an unsafe path.") + return f"{parsed.scheme}://{parsed.netloc}{path}" + + @staticmethod + def _health_status(value: object) -> WorkerHealthStatus: + normalized = str(value or "").strip().lower() + if normalized in {"ok", "healthy", "ready"}: + return WorkerHealthStatus.HEALTHY + if normalized in {"starting", "loading", "initializing"}: + return WorkerHealthStatus.STARTING + if normalized in {"unavailable", "offline"}: + return WorkerHealthStatus.UNAVAILABLE + if normalized in {"unhealthy", "failed", "error"}: + return WorkerHealthStatus.UNHEALTHY + return WorkerHealthStatus.UNKNOWN + + @staticmethod + def _readiness_status(value: object) -> WorkerReadinessStatus: + normalized = str(value or "").strip().lower() + if normalized == "ready": + return WorkerReadinessStatus.READY + if normalized in {"starting", "loading", "initializing"}: + return WorkerReadinessStatus.STARTING + if normalized in { + "not_ready", + "unavailable", + "offline", + "unhealthy", + "failed", + "error", + }: + return WorkerReadinessStatus.UNAVAILABLE + return WorkerReadinessStatus.UNKNOWN + + @staticmethod + def _metadata(payload: dict[str, object], *, known: set[str]) -> dict[str, object]: + data = safe_worker_metadata( + {key: value for key, value in payload.items() if key not in known} + ) + return data if isinstance(data, dict) else {} diff --git a/app/generation/repositories/__init__.py b/app/generation/repositories/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..10385aa9294205a5dc82d5ce78b99c25f85d8afb --- /dev/null +++ b/app/generation/repositories/__init__.py @@ -0,0 +1,3 @@ +from app.generation.repositories.generation import GenerationRepository + +__all__ = ["GenerationRepository"] diff --git a/app/generation/repositories/generation.py b/app/generation/repositories/generation.py new file mode 100644 index 0000000000000000000000000000000000000000..6427baddc33b3ae57519c1c06f97938fa9ee9293 --- /dev/null +++ b/app/generation/repositories/generation.py @@ -0,0 +1,871 @@ +from __future__ import annotations + +import re +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone + +from sqlalchemy import or_, select, update +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from app.generation.domain.enums import GenerationJobStatus, GenerationRequestStatus +from app.generation.domain.errors import ( + GenerationJobNotFoundError, + GenerationOutputConflictError, + GenerationProviderJobConflictError, + GenerationRequestNotFoundError, +) +from app.generation.domain.state_machine import validate_transition +from app.generation.domain.runtime import safe_worker_metadata +from app.generation.models import GenerationJob, GenerationJobAttempt, GenerationRequest +from app.projects.models import ProjectGenerationJob +from app.security.database import SecurityDatabase +from app.security.models import CanonicalMediaAsset + + +_EXTERNAL_JOB_ID = re.compile(r"^[A-Za-z0-9._:-]{1,255}$") + + +@dataclass(frozen=True, slots=True) +class GenerationDispatchRecord: + """Trusted worker view of one tenant-owned generation execution. + + This record is produced only after a durable row lock/claim. It is never + serialised through REST and therefore may contain the opaque worker job ID + needed for reconciliation (but never a worker URL or credential). + """ + + workspace_id: str + user_id: str + generation_request_id: str + generation_job_id: str + provider: str + model_id: str + modality: str + input_asset_id: str | None + project_id: str | None + product_surface: str + spec: dict[str, object] + idempotency_key: str + attempt_number: int + external_job_id: str | None + status: GenerationJobStatus + + +class GenerationRepository: + """Tenant-scoped persistence for immutable requests and durable jobs.""" + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def create( + self, *, request: GenerationRequest, job: GenerationJob, user_id: str + ) -> tuple[GenerationRequest, GenerationJob]: + """Atomically persist one request and its one logical execution.""" + + try: + async with self.database.tenant_session( + workspace_id=request.workspace_id, user_id=user_id + ) as session: + session.add(request) + await session.flush() + job.generation_request_id = request.id + session.add(job) + await session.flush() + if request.project_id is not None: + session.add( + ProjectGenerationJob( + workspace_id=request.workspace_id, + project_id=request.project_id, + generation_job_id=job.id, + attached_by=user_id, + ) + ) + await session.commit() + await session.refresh(request) + await session.refresh(job) + return request, job + except IntegrityError: + # The idempotency unique constraint arbitrates concurrent API + # submissions. The service reads and fingerprints the winner. + existing = await self.get_by_idempotency( + request.workspace_id, request.idempotency_key, user_id=user_id + ) + if existing is None: + raise + return existing + + async def get_by_idempotency( + self, workspace_id: str, idempotency_key: str, *, user_id: str + ) -> tuple[GenerationRequest, GenerationJob] | None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + request = await session.scalar( + select(GenerationRequest).where( + GenerationRequest.workspace_id == workspace_id, + GenerationRequest.idempotency_key == idempotency_key, + ) + ) + if request is None: + return None + job = await session.scalar( + select(GenerationJob).where( + GenerationJob.generation_request_id == request.id, + GenerationJob.workspace_id == workspace_id, + ) + ) + if job is None: + # A one-to-one job is transactionally created with each + # request. Treat any absent row as database corruption rather + # than returning a partial idempotency result. + raise RuntimeError("Generation request is missing its logical job.") + return request, job + + async def list_requests( + self, + workspace_id: str, + *, + user_id: str, + offset: int = 0, + limit: int = 100, + product_surface: str | None = None, + ) -> list[tuple[GenerationRequest, GenerationJob]]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + statement = ( + select(GenerationRequest, GenerationJob) + .join( + GenerationJob, + GenerationJob.generation_request_id == GenerationRequest.id, + ) + .where(GenerationRequest.workspace_id == workspace_id) + .order_by(GenerationRequest.created_at.desc()) + .offset(offset) + .limit(limit) + ) + if product_surface is not None: + statement = statement.where(GenerationRequest.product_surface == product_surface) + rows = await session.execute(statement) + return list(rows.tuples().all()) + + async def get_request( + self, workspace_id: str, request_id: str, *, user_id: str + ) -> tuple[GenerationRequest, GenerationJob]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + row = ( + ( + await session.execute( + select(GenerationRequest, GenerationJob) + .join( + GenerationJob, + GenerationJob.generation_request_id == GenerationRequest.id, + ) + .where( + GenerationRequest.id == request_id, + GenerationRequest.workspace_id == workspace_id, + ) + ) + ) + .tuples() + .one_or_none() + ) + if row is None: + raise GenerationRequestNotFoundError( + "Generation request was not found in this workspace." + ) + return row + + async def get_job(self, workspace_id: str, job_id: str, *, user_id: str) -> GenerationJob: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob).where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + ) + if job is None: + raise GenerationJobNotFoundError("Generation job was not found in this workspace.") + return job + + async def cancel(self, workspace_id: str, job_id: str, *, user_id: str) -> GenerationJob: + """Cancel queued work immediately; request cancellation otherwise. + + The status distinction avoids falsely reporting that a running remote + generation was stopped before a future adapter confirms cancellation. + """ + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise GenerationJobNotFoundError("Generation job was not found in this workspace.") + current = GenerationJobStatus(job.status) + if current in { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + }: + return job + destination = ( + GenerationJobStatus.CANCELLED + if current in {GenerationJobStatus.QUEUED, GenerationJobStatus.RETRYING} + else GenerationJobStatus.CANCEL_REQUESTED + ) + validate_transition(current, destination) + now = datetime.now(timezone.utc) + job.status = destination.value + job.updated_at = now + if destination is GenerationJobStatus.CANCELLED: + job.completed_at = now + await self._set_request_status( + session, + workspace_id=workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=job.completed_at, + ) + if destination is GenerationJobStatus.CANCELLED: + await self._finish_open_attempt( + session, + generation_job_id=job.id, + status=GenerationJobStatus.CANCELLED.value, + error_code=None, + error_message=None, + completed_at=now, + ) + await session.commit() + await session.refresh(job) + return job + + async def transition_job( + self, + workspace_id: str, + job_id: str, + status: GenerationJobStatus, + *, + error_code: str | None = None, + error_message: str | None = None, + next_attempt_at: datetime | None = None, + user_id: str, + ) -> GenerationJob: + """Trusted worker transition hook for a future provider adapter.""" + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise GenerationJobNotFoundError("Generation job was not found in this workspace.") + destination = validate_transition(job.status, status) + now = datetime.now(timezone.utc) + job.status = destination.value + job.error_code, job.error_message = self._safe_error(error_code, error_message) + job.next_attempt_at = next_attempt_at + if destination is GenerationJobStatus.SUBMITTING and job.started_at is None: + job.started_at = now + if destination in { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + }: + job.completed_at = now + await self._set_request_status( + session, + workspace_id=workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=job.completed_at, + ) + if destination in { + GenerationJobStatus.RETRYING, + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + }: + await self._finish_open_attempt( + session, + generation_job_id=job.id, + status=destination.value, + error_code=job.error_code, + error_message=job.error_message, + completed_at=now, + ) + await session.commit() + await session.refresh(job) + return job + + async def bind_external_job( + self, + workspace_id: str, + job_id: str, + *, + user_id: str, + provider: str, + external_job_id: str, + provider_metadata: dict[str, object] | None = None, + ) -> GenerationJob: + """Bind one trusted worker job identity to one logical tenant job. + + The external identity never enters through a public API. This + repository method still checks the workspace and provider and uses a + uniqueness constraint as a final guard against cross-workspace job + attachment during concurrent worker recovery. + """ + + if _EXTERNAL_JOB_ID.fullmatch(external_job_id) is None: + raise GenerationProviderJobConflictError("Generation worker job ID is invalid.") + + try: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise GenerationJobNotFoundError( + "Generation job was not found in this workspace." + ) + if job.provider != provider: + raise GenerationProviderJobConflictError( + "Generation worker job belongs to a different provider." + ) + if GenerationJobStatus(job.status) not in { + GenerationJobStatus.SUBMITTING, + GenerationJobStatus.RUNNING, + GenerationJobStatus.CANCEL_REQUESTED, + }: + raise GenerationProviderJobConflictError( + "Generation job is not in a state that can bind a worker job." + ) + if job.external_job_id is not None: + if job.external_job_id != external_job_id: + raise GenerationProviderJobConflictError( + "Generation job is already bound to a different worker job." + ) + return job + + conflicting = await session.scalar( + select(GenerationJob.id).where( + GenerationJob.provider == provider, + GenerationJob.external_job_id == external_job_id, + GenerationJob.id != job.id, + ) + ) + if conflicting is not None: + raise GenerationProviderJobConflictError( + "Generation worker job is already bound to another workspace." + ) + job.external_job_id = external_job_id + job.provider_metadata_json = self._merge_safe_metadata( + job.provider_metadata_json, provider_metadata + ) + job.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(job) + return job + except IntegrityError as exc: + # The unique database constraint arbitrates a binding race across + # independently running worker processes without leaking either + # workspace or provider details. + raise GenerationProviderJobConflictError( + "Generation worker job is already bound to another logical job." + ) from exc + + async def complete_with_output_asset( + self, + workspace_id: str, + job_id: str, + *, + user_id: str, + external_job_id: str, + output_asset_id: str, + provider_metadata: dict[str, object] | None = None, + ) -> GenerationJob: + """Atomically attach an owned canonical output and complete the job.""" + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise GenerationJobNotFoundError("Generation job was not found in this workspace.") + if job.external_job_id != external_job_id: + raise GenerationProviderJobConflictError( + "Generation worker job is not bound to this logical job." + ) + asset = await session.scalar( + select(CanonicalMediaAsset).where( + CanonicalMediaAsset.id == output_asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + ) + ) + if asset is None: + raise GenerationOutputConflictError( + "Generation output asset is not owned by this workspace." + ) + if job.output_asset_id is not None: + if job.output_asset_id == output_asset_id: + return job + raise GenerationOutputConflictError( + "Generation job already has a different canonical output asset." + ) + + destination = validate_transition(job.status, GenerationJobStatus.SUCCEEDED) + now = datetime.now(timezone.utc) + job.status = destination.value + job.output_asset_id = output_asset_id + job.provider_metadata_json = self._merge_safe_metadata( + job.provider_metadata_json, provider_metadata + ) + job.error_code = None + job.error_message = None + job.next_attempt_at = None + job.completed_at = now + job.updated_at = now + await self._set_request_status( + session, + workspace_id=workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=now, + ) + await self._finish_open_attempt( + session, + generation_job_id=job.id, + status=GenerationJobStatus.SUCCEEDED.value, + error_code=None, + error_message=None, + completed_at=now, + ) + await session.commit() + await session.refresh(job) + return job + + async def start_attempt( + self, workspace_id: str, job_id: str, *, user_id: str + ) -> GenerationJobAttempt: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise GenerationJobNotFoundError("Generation job was not found in this workspace.") + job.attempt_count += 1 + attempt = GenerationJobAttempt( + generation_job_id=job.id, + attempt_number=job.attempt_count, + status="started", + ) + session.add(attempt) + await session.commit() + await session.refresh(attempt) + return attempt + + async def claim_dispatchable(self, *, limit: int) -> list[GenerationDispatchRecord]: + """Atomically claim unbound queued/retry jobs for one submission. + + The worker/service role is the only caller. ``FOR UPDATE SKIP + LOCKED`` prevents two worker processes from submitting the same + logical job; SQLite safely serialises writers and ignores the lock + clause. A claim persists ``submitting`` and an attempt before network + I/O, so crash recovery can distinguish a never-started job from an + ambiguous lost worker response. + """ + + now = datetime.now(timezone.utc) + claimed: list[GenerationDispatchRecord] = [] + async with self.database.session() as session: + rows = await session.execute( + select(GenerationJob, GenerationRequest) + .join( + GenerationRequest, + GenerationRequest.id == GenerationJob.generation_request_id, + ) + .where( + GenerationJob.external_job_id.is_(None), + GenerationJob.status.in_( + [GenerationJobStatus.QUEUED.value, GenerationJobStatus.RETRYING.value] + ), + or_( + GenerationJob.next_attempt_at.is_(None), + GenerationJob.next_attempt_at <= now, + ), + ) + .order_by(GenerationJob.created_at) + .limit(limit) + .with_for_update(skip_locked=True) + ) + for job, request in rows.tuples().all(): + if job.attempt_count >= job.max_attempts: + result = await session.execute( + update(GenerationJob) + .where( + GenerationJob.id == job.id, + GenerationJob.status == job.status, + GenerationJob.external_job_id.is_(None), + GenerationJob.attempt_count == job.attempt_count, + ) + .values( + status=GenerationJobStatus.FAILED.value, + error_code="GENERATION_RETRY_LIMIT_EXCEEDED", + error_message="Generation retry limit was reached before submission.", + completed_at=now, + updated_at=now, + ) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + continue + destination = validate_transition(job.status, GenerationJobStatus.FAILED) + job.status = destination.value + job.error_code = "GENERATION_RETRY_LIMIT_EXCEEDED" + job.error_message = "Generation retry limit was reached before submission." + job.completed_at = now + job.updated_at = now + await self._set_request_status( + session, + workspace_id=job.workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=now, + ) + continue + original_status = job.status + original_attempt_count = job.attempt_count + destination = validate_transition(original_status, GenerationJobStatus.SUBMITTING) + result = await session.execute( + update(GenerationJob) + .where( + GenerationJob.id == job.id, + GenerationJob.status == original_status, + GenerationJob.external_job_id.is_(None), + GenerationJob.attempt_count == original_attempt_count, + ) + .values( + status=destination.value, + error_code=None, + error_message=None, + next_attempt_at=None, + updated_at=now, + started_at=job.started_at or now, + attempt_count=original_attempt_count + 1, + ) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + continue + job.status = destination.value + job.error_code = None + job.error_message = None + job.next_attempt_at = None + job.updated_at = now + if job.started_at is None: + job.started_at = now + job.attempt_count = original_attempt_count + 1 + attempt = GenerationJobAttempt( + generation_job_id=job.id, + attempt_number=job.attempt_count, + status="submitting", + ) + session.add(attempt) + await self._set_request_status( + session, + workspace_id=job.workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=None, + ) + claimed.append( + GenerationDispatchRecord( + workspace_id=job.workspace_id, + user_id=request.created_by_user_id, + generation_request_id=request.id, + generation_job_id=job.id, + provider=job.provider, + model_id=request.model_id, + modality=request.modality, + input_asset_id=request.input_asset_id, + project_id=request.project_id, + product_surface=request.product_surface, + spec=dict(request.spec_json or {}), + idempotency_key=request.idempotency_key, + attempt_number=job.attempt_count, + external_job_id=None, + status=destination, + ) + ) + await session.commit() + return claimed + + async def list_reconcilable( + self, *, limit: int, lease_seconds: int + ) -> list[GenerationDispatchRecord]: + """Return bound active jobs for idempotent status reconciliation.""" + + now = datetime.now(timezone.utc) + lease_until = now + timedelta(seconds=max(1, lease_seconds)) + async with self.database.session() as session: + rows = await session.execute( + select(GenerationJob, GenerationRequest) + .join( + GenerationRequest, + GenerationRequest.id == GenerationJob.generation_request_id, + ) + .where( + GenerationJob.external_job_id.is_not(None), + GenerationJob.status.in_( + [ + GenerationJobStatus.QUEUED.value, + GenerationJobStatus.SUBMITTING.value, + GenerationJobStatus.RUNNING.value, + GenerationJobStatus.CANCEL_REQUESTED.value, + ] + ), + or_( + GenerationJob.next_attempt_at.is_(None), + GenerationJob.next_attempt_at <= now, + ), + ) + .order_by(GenerationJob.updated_at) + .limit(limit) + ) + records: list[GenerationDispatchRecord] = [] + for job, request in rows.tuples().all(): + result = await session.execute( + update(GenerationJob) + .where( + GenerationJob.id == job.id, + GenerationJob.status == job.status, + GenerationJob.external_job_id == job.external_job_id, + or_( + GenerationJob.next_attempt_at.is_(None), + GenerationJob.next_attempt_at <= now, + ), + ) + .values(next_attempt_at=lease_until, updated_at=now) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + continue + job.next_attempt_at = lease_until + job.updated_at = now + records.append( + GenerationDispatchRecord( + workspace_id=job.workspace_id, + user_id=request.created_by_user_id, + generation_request_id=request.id, + generation_job_id=job.id, + provider=job.provider, + model_id=request.model_id, + modality=request.modality, + input_asset_id=request.input_asset_id, + project_id=request.project_id, + spec=dict(request.spec_json or {}), + idempotency_key=request.idempotency_key, + attempt_number=job.attempt_count, + external_job_id=job.external_job_id, + status=GenerationJobStatus(job.status), + ) + ) + await session.commit() + return records + + async def set_reconciliation_due( + self, + workspace_id: str, + job_id: str, + *, + user_id: str, + due_at: datetime | None, + ) -> None: + """Release or defer one bound active job's durable reconciliation lease.""" + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(GenerationJob) + .where( + GenerationJob.id == job_id, + GenerationJob.workspace_id == workspace_id, + GenerationJob.external_job_id.is_not(None), + GenerationJob.status.in_( + [ + GenerationJobStatus.QUEUED.value, + GenerationJobStatus.SUBMITTING.value, + GenerationJobStatus.RUNNING.value, + GenerationJobStatus.CANCEL_REQUESTED.value, + ] + ), + ) + .with_for_update() + ) + if job is None: + return + job.next_attempt_at = due_at + job.updated_at = datetime.now(timezone.utc) + await session.commit() + + async def fail_stale_unbound_submissions(self, *, stale_before: datetime) -> int: + """Fail unrecoverably ambiguous submissions that lost their response. + + A remote worker without an idempotency/reconciliation lookup cannot be + safely resubmitted after MediaRouter crashes or loses the job ID. A + terminal, auditable failure is preferable to silently making a second + video. This is intentionally provider-neutral; a future idempotent + provider may choose a different recovery route before calling it. + """ + + count = 0 + async with self.database.session() as session: + rows = await session.scalars( + select(GenerationJob) + .where( + GenerationJob.status == GenerationJobStatus.SUBMITTING.value, + GenerationJob.external_job_id.is_(None), + GenerationJob.updated_at < stale_before, + ) + .with_for_update(skip_locked=True) + ) + now = datetime.now(timezone.utc) + for job in rows.all(): + destination = validate_transition(job.status, GenerationJobStatus.FAILED) + job.status = destination.value + job.error_code = "GENERATION_SUBMISSION_AMBIGUOUS" + job.error_message = ( + "Worker submission could not be reconciled safely; it was not retried." + ) + job.completed_at = now + job.updated_at = now + await self._set_request_status( + session, + workspace_id=job.workspace_id, + request_id=job.generation_request_id, + status=destination, + completed_at=now, + ) + await self._finish_open_attempt( + session, + generation_job_id=job.id, + status="failed_ambiguous", + error_code=job.error_code, + error_message=job.error_message, + completed_at=now, + ) + count += 1 + await session.commit() + return count + + @staticmethod + def _merge_safe_metadata( + existing: dict[str, object] | None, incoming: dict[str, object] | None + ) -> dict[str, object]: + """Prevent a worker response from becoming a token/URL secret store.""" + + safe_existing = safe_worker_metadata(existing or {}) + safe_incoming = safe_worker_metadata(incoming or {}) + assert isinstance(safe_existing, dict) + assert isinstance(safe_incoming, dict) + return {**safe_existing, **safe_incoming} + + @staticmethod + async def _finish_open_attempt( + session: AsyncSession, + *, + generation_job_id: str, + status: str, + error_code: str | None, + error_message: str | None, + completed_at: datetime, + ) -> None: + """Close the newest open attempt without creating a second job.""" + + attempt = await session.scalar( + select(GenerationJobAttempt) + .where( + GenerationJobAttempt.generation_job_id == generation_job_id, + GenerationJobAttempt.completed_at.is_(None), + ) + .order_by(GenerationJobAttempt.attempt_number.desc()) + .limit(1) + .with_for_update() + ) + if attempt is None: + return + attempt.status = status[:32] + attempt.error_code = error_code + attempt.error_message = error_message + attempt.completed_at = completed_at + + @staticmethod + def _safe_error(code: str | None, message: str | None) -> tuple[str | None, str | None]: + """Persist only safe, bounded worker error information.""" + + safe_code = code if code and _EXTERNAL_JOB_ID.fullmatch(code) else None + safe_message = safe_worker_metadata(message) if message is not None else None + if not isinstance(safe_message, str): + safe_message = None + return safe_code, safe_message[:500] if safe_message is not None else None + + @staticmethod + async def _set_request_status( + session: AsyncSession, + *, + workspace_id: str, + request_id: str, + status: GenerationJobStatus, + completed_at: datetime | None, + ) -> None: + request = await session.scalar( + select(GenerationRequest).where( + GenerationRequest.id == request_id, + GenerationRequest.workspace_id == workspace_id, + ) + ) + if request is None: + raise GenerationRequestNotFoundError( + "Generation request was not found in this workspace." + ) + request.status = GenerationRequestStatus(status.value).value + request.completed_at = completed_at diff --git a/app/generation/schemas/__init__.py b/app/generation/schemas/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..f873a4e635690064cecacfadcbdb0a0f77b25377 --- /dev/null +++ b/app/generation/schemas/__init__.py @@ -0,0 +1,17 @@ +"""Strict transport schemas for the generation foundation.""" + +from app.generation.schemas.requests import ( + GenerationJobView, + GenerationProviderView, + GenerationRequestCreate, + GenerationRequestView, + WanGenerationOptions, +) + +__all__ = [ + "GenerationJobView", + "GenerationProviderView", + "GenerationRequestCreate", + "GenerationRequestView", + "WanGenerationOptions", +] diff --git a/app/generation/schemas/requests.py b/app/generation/schemas/requests.py new file mode 100644 index 0000000000000000000000000000000000000000..92e686076b1712e8461e044f3516f22578995684 --- /dev/null +++ b/app/generation/schemas/requests.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +from app.generation.domain.capabilities import GenerationProviderCapabilities +from app.generation.domain.enums import ( + GenerationJobStatus, + GenerationModality, + GenerationRequestStatus, +) + + +class GenerationRequestCreate(BaseModel): + """Small provider-neutral submission envelope. + + Provider-specific controls are deliberately absent until a concrete + adapter supplies a typed schema. Callers can reference only a canonical + MediaRouter asset, never a filesystem path or arbitrary remote URL. + """ + + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + provider: str = Field(min_length=1, max_length=64, pattern=r"^[a-z][a-z0-9_-]{0,63}$") + model_id: str = Field(min_length=1, max_length=255) + modality: GenerationModality + prompt: str = Field(min_length=1, max_length=10_000) + input_asset_id: str | None = Field( + default=None, + pattern=(r"^[0-9a-f]{8}-[0-9a-f]{4}-[1-5][0-9a-f]{3}-" r"[89ab][0-9a-f]{3}-[0-9a-f]{12}$"), + ) + brand_kit_version_id: str | None = Field(default=None, min_length=1, max_length=36) + # Concrete-provider controls remain closed, typed schemas rather than an + # escape hatch for arbitrary worker payload. Provider-specific fields must + # stay independently reviewed and cannot become a generic provider-payload + # escape hatch. + wan: "WanGenerationOptions | None" = None + flux: "FluxGenerationOptions | None" = None + + @field_validator("model_id") + @classmethod + def normalized_model_id(cls, value: str) -> str: + # Keep model IDs stable and URL-safe; adapters can impose narrower + # constraints but cannot accept hidden whitespace or control chars. + if any(character.isspace() for character in value): + raise ValueError("model_id must not contain whitespace") + return value + + +class WanGenerationOptions(BaseModel): + """Exact optional controls accepted by the audited WAN 2.2 worker. + + Dimensions and frame counts are intentionally absent: WAN derives both + from the canonical source image and ``duration_seconds``. Unknown fields + (including ``width``, ``height`` and ``num_frames``) are rejected before a + GPU request can be created. + """ + + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + negative_prompt: str | None = Field(default=None, max_length=4_000) + duration_seconds: float = Field(default=5.0, ge=0.5, le=5.0, strict=True) + steps: int = Field(default=4, ge=1, le=30, strict=True) + guidance_scale: float = Field(default=1.0, ge=0.0, le=10.0, strict=True) + guidance_scale_2: float = Field(default=1.0, ge=0.0, le=10.0, strict=True) + seed: int = Field(default=42, ge=0, le=2_147_483_647, strict=True) + randomize_seed: bool = Field(default=False, strict=True) + + @field_validator("negative_prompt") + @classmethod + def negative_prompt_is_bounded(cls, value: str | None) -> str | None: + # The worker permits an empty negative prompt (which explicitly + # overrides its built-in default); preserve that real behavior. + return value + + +class FluxGenerationOptions(BaseModel): + """Exact optional controls accepted by the audited FLUX worker. + + FLUX has no negative-prompt or scheduler field in its worker API. The + mode selector chooses between its two loaded pipelines; it is closed to + the values actually accepted by the worker. + """ + + model_config = ConfigDict(extra="forbid", str_strip_whitespace=True) + + mode_choice: Literal["Distilled (4 steps)", "Base (50 steps)"] = "Distilled (4 steps)" + seed: int = Field(default=42, ge=0, le=2_147_483_647, strict=True) + randomize_seed: bool = Field(default=False, strict=True) + width: int = Field(default=1024, ge=256, le=1024, strict=True) + height: int = Field(default=1024, ge=256, le=1024, strict=True) + num_inference_steps: int = Field(default=4, ge=1, le=100, strict=True) + guidance_scale: float = Field(default=1.0, ge=0.0, le=10.0, strict=True) + prompt_upsampling: bool = Field(default=False, strict=True) + + @model_validator(mode="after") + def dimensions_are_worker_compatible(self) -> "FluxGenerationOptions": + if self.width % 8 or self.height % 8: + raise ValueError("width and height must be multiples of 8") + return self + + +# ``GenerationRequestCreate`` intentionally appears before the provider schemas +# in this module so its provider-neutral fields remain easy to audit. Rebuild +# its forward references after the closed provider schemas exist. +GenerationRequestCreate.model_rebuild() + + +class GenerationProviderView(BaseModel): + """Public discovery state with no worker endpoints or credentials.""" + + model_config = ConfigDict(extra="forbid") + + capabilities: GenerationProviderCapabilities + available: bool + + +class GenerationJobView(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + generation_request_id: str + status: GenerationJobStatus + attempt_count: int + max_attempts: int + brand_kit_version_id: str | None = None + output_asset_id: str | None = None + error_code: str | None = None + error_message: str | None = None + next_attempt_at: datetime | None = None + created_at: datetime + started_at: datetime | None = None + completed_at: datetime | None = None + updated_at: datetime + + @classmethod + def from_record(cls, record: object) -> "GenerationJobView": + return cls.model_validate( + { + "id": getattr(record, "id"), + "generation_request_id": getattr(record, "generation_request_id"), + "status": getattr(record, "status"), + "attempt_count": getattr(record, "attempt_number"), + "max_attempts": getattr(record, "max_attempts"), + "brand_kit_version_id": getattr(record, "brand_kit_version_id", None), + "output_asset_id": getattr(record, "output_asset_id"), + "error_code": getattr(record, "error_code"), + "error_message": getattr(record, "error_message"), + "next_attempt_at": getattr(record, "next_attempt_at"), + "created_at": getattr(record, "created_at"), + "started_at": getattr(record, "started_at"), + "completed_at": getattr(record, "completed_at"), + "updated_at": getattr(record, "updated_at"), + } + ) + + +class GenerationRequestView(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + provider: str + model_id: str + modality: GenerationModality + prompt: str + input_asset_id: str | None = None + project_id: str | None = None + brand_kit_version_id: str | None = None + product_surface: Literal["generation", "ai_studio"] = Field( + default="generation", + exclude=True, + ) + status: GenerationRequestStatus + created_at: datetime + updated_at: datetime + completed_at: datetime | None = None + job: GenerationJobView + + @classmethod + def from_records(cls, request: object, job: object) -> "GenerationRequestView": + spec = getattr(request, "spec_json", {}) + if not isinstance(spec, dict): + spec = {} + return cls( + id=getattr(request, "id"), + provider=getattr(request, "provider"), + model_id=getattr(request, "model_id"), + modality=getattr(request, "modality"), + prompt=str(spec.get("prompt", "")), + input_asset_id=getattr(request, "input_asset_id"), + project_id=getattr(request, "project_id", None), + brand_kit_version_id=getattr(job, "brand_kit_version_id", None), + product_surface=getattr(request, "product_surface", "generation"), + status=getattr(request, "status"), + created_at=getattr(request, "created_at"), + updated_at=getattr(request, "updated_at"), + completed_at=getattr(request, "completed_at"), + job=GenerationJobView.from_record(job), + ) diff --git a/app/generation/services/__init__.py b/app/generation/services/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..330be25b27529ec8dd37d0b96ec39d7963d35194 --- /dev/null +++ b/app/generation/services/__init__.py @@ -0,0 +1,4 @@ +from app.generation.services.generation_service import GenerationService +from app.generation.services.output_ingestion import GenerationOutputIngestor + +__all__ = ["GenerationOutputIngestor", "GenerationService"] diff --git a/app/generation/services/generation_service.py b/app/generation/services/generation_service.py new file mode 100644 index 0000000000000000000000000000000000000000..e9b8d0bce65bc9a35789f25d84e586cdf61cc02a --- /dev/null +++ b/app/generation/services/generation_service.py @@ -0,0 +1,660 @@ +from __future__ import annotations + +import hashlib +import json +from datetime import datetime, timedelta, timezone +from typing import Literal + +from app.core.config import Settings +from app.core.logger import get_logger +from app.generation.domain.enums import ( + GenerationJobStatus, + WorkerCancellationStatus, + WorkerJobStatus, +) +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationCancellationError, + GenerationIdempotencyConflictError, + GenerationInputAssetNotFoundError, + GenerationOutputError, + GenerationProviderUnavailableError, + GenerationValidationError, +) +from app.generation.model_registry import GenerationModelRegistry, GenerationModelView +from app.generation.models import GenerationJob, GenerationRequest +from app.generation.domain.runtime import WorkerJob +from app.generation.providers.registry import GenerationProviderRegistry +from app.generation.repositories.generation import GenerationRepository +from app.generation.repositories.generation import GenerationDispatchRecord +from app.generation.schemas.requests import ( + GenerationJobView, + GenerationProviderView, + GenerationRequestCreate, + GenerationRequestView, +) +from app.generation.services.output_ingestion import GenerationOutputIngestor +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.database import SecurityDatabase + +logger = get_logger(__name__) + + +class GenerationService: + """Transport-neutral generation facade. + + This service owns request validation, tenant isolation, durable job + creation, and idempotency. It intentionally never opens a WAN or FLUX + connection; adding a worker adapter later is the only route to accepting + generation work. + """ + + def __init__( + self, + *, + settings: Settings, + database: SecurityDatabase, + assets: CanonicalAssetService, + repository: GenerationRepository, + providers: GenerationProviderRegistry, + models: GenerationModelRegistry, + output_ingestor: GenerationOutputIngestor, + ) -> None: + self.settings = settings + self.database = database + self.assets = assets + self.repository = repository + self.providers = providers + self.models = models + self.output_ingestor = output_ingestor + self.ready = False + + async def initialize(self) -> None: + # SecurityDatabase initialization is owned by the application + # lifespan. It creates the local SQLite metadata or verifies the + # production PostgreSQL migration before this service becomes ready. + self.ready = self.settings.generation_enabled and await self.database.schema_ready() + if self.ready: + logger.info("generation foundation initialized") + elif self.settings.generation_enabled: + logger.warning( + "generation foundation schema unavailable; apply the generation migration" + ) + else: + logger.info("generation foundation disabled") + + def ensure_ready(self) -> None: + if not self.settings.generation_enabled: + raise GenerationProviderUnavailableError("Generation is disabled.") + if not self.ready: + raise GenerationProviderUnavailableError( + "Generation storage is unavailable. Apply the generation migration." + ) + + def list_providers(self) -> list[GenerationProviderView]: + self.ensure_ready() + return [self._provider_view(adapter.provider) for adapter in self.providers.list()] + + def get_provider(self, provider: str) -> GenerationProviderView: + self.ensure_ready() + return self._provider_view(provider) + + def _provider_view(self, provider: str) -> GenerationProviderView: + """Return public availability only after model verification. + + An adapter's ``available`` flag intentionally means its server-owned + configuration can be contacted; it is used by the dispatcher to run + health/readiness discovery. It must not by itself be advertised to + clients as a usable provider. At least one registered model must have + passed the exact health, identity, modality, and readiness checks. + """ + + adapter = self.providers.get(provider) + verified_model_available = any( + model.available for model in self.models.list(provider_id=adapter.provider) + ) + return GenerationProviderView( + capabilities=adapter.capabilities, + available=adapter.available and verified_model_available, + ) + + def list_models(self, *, provider: str | None = None) -> list[GenerationModelView]: + self.ensure_ready() + return self.models.list(provider_id=provider.strip().lower() if provider else None) + + def get_model(self, provider: str, model_id: str) -> GenerationModelView: + self.ensure_ready() + return self.models.get(provider.strip().lower(), model_id) + + async def refresh_provider_runtime(self, provider: str) -> list[GenerationModelView]: + """Verify configured worker metadata and readiness before advertising models. + + No startup caller invokes this while the registry is empty. Future + provider integrations may call it after their trusted configuration is + loaded; a health success alone never makes a model available. + """ + + self.ensure_ready() + adapter = self.providers.get(provider) + try: + health = await adapter.health() + if health.status.value != "healthy": + self.models.mark_unavailable(adapter.provider) + return self.models.list(provider_id=adapter.provider) + info = await adapter.info() + readiness = await adapter.ready() + except Exception: + # Do not expose a worker exception or leave a stale model marked + # available. The future provider worker owns detailed diagnostics. + self.models.mark_unavailable(adapter.provider) + raise + return self.models.verify_readiness( + provider_id=adapter.provider, + worker_info=info, + readiness=readiness, + provider_configured=adapter.available, + ) + + async def create( + self, + *, + workspace_id: str, + user_id: str, + payload: GenerationRequestCreate, + idempotency_key: str, + project_id: str | None = None, + brand_kit_version_id: str | None = None, + product_surface: Literal["generation", "ai_studio"] = "generation", + ) -> GenerationRequestView: + self.ensure_ready() + key = idempotency_key.strip() + if not key: + raise GenerationValidationError("Idempotency-Key must not be blank.") + + base_spec = self._request_spec(payload) + requested_fingerprint = self._fingerprint( + provider=payload.provider, + model_id=payload.model_id, + modality=payload.modality.value, + input_asset_id=payload.input_asset_id, + project_id=project_id, + brand_kit_version_id=brand_kit_version_id, + product_surface=product_surface, + spec=base_spec, + ) + existing = await self.repository.get_by_idempotency(workspace_id, key, user_id=user_id) + if existing is not None: + request, job = existing + if request.request_fingerprint != requested_fingerprint: + raise GenerationIdempotencyConflictError( + "Idempotency-Key is already associated with a different generation request." + ) + return GenerationRequestView.from_records(request, job) + + adapter = self.providers.get(payload.provider) + model = adapter.model_capability(payload.model_id) + registered_model = self.models.get(adapter.provider, payload.model_id) + if not registered_model.available: + raise GenerationProviderUnavailableError( + f"Generation model '{payload.model_id}' is not ready to accept work." + ) + if registered_model.model.modality != model.modality: + raise GenerationCapabilityUnsupportedError( + f"Generation model '{payload.model_id}' has inconsistent capability metadata." + ) + if model.modality != payload.modality: + raise GenerationCapabilityUnsupportedError( + f"Model '{payload.model_id}' does not support {payload.modality.value} generation." + ) + if payload.input_asset_id and not model.input_asset_supported: + raise GenerationCapabilityUnsupportedError( + f"Model '{payload.model_id}' does not support canonical input assets." + ) + if not adapter.available: + raise GenerationProviderUnavailableError( + f"Generation provider '{adapter.provider}' is not ready to accept work." + ) + + input_asset = None + if payload.input_asset_id: + try: + input_asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=payload.input_asset_id, + ) + except CanonicalAssetNotFoundError as exc: + raise GenerationInputAssetNotFoundError( + "Input asset was not found in this workspace." + ) from exc + + normalized = await adapter.validate_request(payload) + if input_asset is not None: + await adapter.validate_input_asset(payload, input_asset) + spec = self._validated_spec(payload, normalized) + fingerprint = self._fingerprint( + provider=adapter.provider, + model_id=payload.model_id, + modality=payload.modality.value, + input_asset_id=payload.input_asset_id, + project_id=project_id, + brand_kit_version_id=brand_kit_version_id, + product_surface=product_surface, + spec=spec, + ) + if fingerprint != requested_fingerprint: + raise GenerationCapabilityUnsupportedError( + "Generation adapter returned a request outside the foundation contract." + ) + + request = GenerationRequest( + workspace_id=workspace_id, + created_by_user_id=user_id, + provider=adapter.provider, + model_id=payload.model_id, + modality=payload.modality.value, + input_asset_id=payload.input_asset_id, + project_id=project_id, + product_surface=product_surface, + spec_json=spec, + request_fingerprint=fingerprint, + idempotency_key=key, + status=GenerationJobStatus.QUEUED.value, + ) + job = GenerationJob( + generation_request_id="", + workspace_id=workspace_id, + provider=adapter.provider, + status=GenerationJobStatus.QUEUED.value, + brand_kit_version_id=brand_kit_version_id, + max_attempts=self.settings.generation_job_retry_limit + 1, + ) + created_request, created_job = await self.repository.create( + request=request, job=job, user_id=user_id + ) + if created_request.request_fingerprint != fingerprint: + raise GenerationIdempotencyConflictError( + "Idempotency-Key is already associated with a different generation request." + ) + logger.info( + "generation request queued", + extra={ + "generation_request_id": created_request.id, + "generation_job_id": created_job.id, + "workspace_id": workspace_id, + "provider": adapter.provider, + "model_id": payload.model_id, + }, + ) + return GenerationRequestView.from_records(created_request, created_job) + + async def list_requests( + self, + workspace_id: str, + user_id: str, + *, + offset: int = 0, + limit: int = 100, + product_surface: str | None = None, + ) -> list[GenerationRequestView]: + self.ensure_ready() + records = await self.repository.list_requests( + workspace_id, + user_id=user_id, + offset=offset, + limit=limit, + product_surface=product_surface, + ) + return [GenerationRequestView.from_records(request, job) for request, job in records] + + async def get_request( + self, workspace_id: str, user_id: str, request_id: str + ) -> GenerationRequestView: + self.ensure_ready() + request, job = await self.repository.get_request(workspace_id, request_id, user_id=user_id) + return GenerationRequestView.from_records(request, job) + + async def get_job(self, workspace_id: str, user_id: str, job_id: str) -> GenerationJobView: + self.ensure_ready() + return GenerationJobView.from_record( + await self.repository.get_job(workspace_id, job_id, user_id=user_id) + ) + + async def cancel(self, workspace_id: str, user_id: str, job_id: str) -> GenerationJobView: + self.ensure_ready() + current = await self.repository.get_job(workspace_id, job_id, user_id=user_id) + status = GenerationJobStatus(current.status) + active = { + GenerationJobStatus.QUEUED, + GenerationJobStatus.SUBMITTING, + GenerationJobStatus.RUNNING, + GenerationJobStatus.CANCEL_REQUESTED, + } + if status in active and current.external_job_id: + adapter = self.providers.get(current.provider) + try: + result = await adapter.cancel(external_job_id=current.external_job_id) + except Exception as exc: + try: + normalized = adapter.normalize_error(exc) + except Exception: + # An adapter's error normalizer is diagnostic-only. A + # programming error there must not expose worker context + # or turn cancellation into a false success. + normalized = None + if isinstance(normalized, GenerationCancellationError): + raise normalized + # Never surface provider exception text: worker clients can + # include request context and third-party adapters may not. + raise GenerationCancellationError() from None + if result.status is WorkerCancellationStatus.UNSUPPORTED: + raise GenerationCapabilityUnsupportedError( + f"{adapter.provider} does not support remote generation cancellation." + ) + if result.status is WorkerCancellationStatus.FAILED: + raise GenerationCancellationError() + if result.status is WorkerCancellationStatus.CANCELLED: + job = await self.repository.transition_job( + workspace_id, + job_id, + GenerationJobStatus.CANCELLED, + user_id=user_id, + ) + else: + # The worker accepted a request but did not claim that its GPU + # inference stopped. Preserve this distinction locally. + if status is GenerationJobStatus.CANCEL_REQUESTED: + job = current + else: + job = await self.repository.transition_job( + workspace_id, + job_id, + GenerationJobStatus.CANCEL_REQUESTED, + user_id=user_id, + ) + else: + # Queued jobs are safe to cancel locally. Active jobs without an + # external identity remain cancellation_requested for the worker + # to observe before/after it establishes a remote job. + job = await self.repository.cancel(workspace_id, job_id, user_id=user_id) + logger.info( + "generation cancellation requested", + extra={ + "generation_job_id": job.id, + "workspace_id": workspace_id, + "status": job.status, + }, + ) + return GenerationJobView.from_record(job) + + async def bind_provider_job( + self, + *, + workspace_id: str, + user_id: str, + job_id: str, + worker_job_id: str, + provider_metadata: dict[str, object] | None = None, + ) -> GenerationJobView: + """Bind a worker-issued opaque job ID through the tenant boundary. + + This is an internal worker integration hook, not an HTTP endpoint. + Future dispatchers must use it after their state-machine transition, + rather than updating ``GenerationJob.external_job_id`` directly. + """ + + self.ensure_ready() + current = await self.repository.get_job(workspace_id, job_id, user_id=user_id) + adapter = self.providers.get(current.provider) + job = await self.repository.bind_external_job( + workspace_id, + job_id, + user_id=user_id, + provider=adapter.provider, + external_job_id=worker_job_id, + provider_metadata=provider_metadata, + ) + return GenerationJobView.from_record(job) + + async def ingest_completed_provider_output( + self, + *, + workspace_id: str, + user_id: str, + job_id: str, + worker_job: WorkerJob | None = None, + ) -> GenerationJobView: + """Ingest a completed worker output into the canonical asset pipeline. + + The public API cannot call this method or provide an output URL/path. + A future trusted dispatcher calls it only after status reconciliation. + The job's bound provider and external ID are the sole source of worker + selection, preserving workspace ownership throughout the handoff. + """ + + self.ensure_ready() + current = await self.repository.get_job(workspace_id, job_id, user_id=user_id) + if current.output_asset_id: + return GenerationJobView.from_record(current) + if not current.external_job_id: + raise GenerationOutputError("Generation job has no bound worker output.") + if GenerationJobStatus(current.status) not in { + GenerationJobStatus.QUEUED, + GenerationJobStatus.SUBMITTING, + GenerationJobStatus.RUNNING, + GenerationJobStatus.CANCEL_REQUESTED, + }: + raise GenerationOutputError("Generation job is not ready to ingest a worker output.") + + adapter = self.providers.get(current.provider) + worker_job = worker_job or await adapter.get_job(external_job_id=current.external_job_id) + if worker_job.external_job_id != current.external_job_id: + raise GenerationOutputError("Generation worker returned an unexpected job identity.") + if worker_job.status is not WorkerJobStatus.COMPLETED or worker_job.output is None: + raise GenerationOutputError("Generation worker output is not ready.") + + async with adapter.stream_output(worker_job.output) as chunks: + asset = await self.output_ingestor.ingest( + workspace_id=workspace_id, + user_id=user_id, + generation_request_id=current.generation_request_id, + generation_job_id=current.id, + provider_id=adapter.provider, + project_id=( + await self.repository.get_request( + workspace_id, + current.generation_request_id, + user_id=user_id, + ) + )[0].project_id, + output=worker_job.output, + chunks=chunks, + ) + job = await self.repository.complete_with_output_asset( + workspace_id, + job_id, + user_id=user_id, + external_job_id=current.external_job_id, + output_asset_id=asset.id, + provider_metadata={ + "provider_output_id": worker_job.output.provider_output_id, + "output_type": worker_job.output.output_type.value, + "mime_type": worker_job.output.mime_type, + "worker_metadata": worker_job.metadata, + }, + ) + logger.info( + "generation worker output ingested", + extra={ + "generation_job_id": job.id, + "workspace_id": workspace_id, + "provider": adapter.provider, + "output_asset_id": asset.id, + }, + ) + return GenerationJobView.from_record(job) + + # The following methods form the trusted worker boundary. They are not + # mounted as REST/MCP/SDK endpoints, so external callers cannot claim + # arbitrary jobs, supply an external job identity, or alter state. + + async def claim_dispatchable_jobs(self, *, limit: int) -> list[GenerationDispatchRecord]: + self.ensure_ready() + return await self.repository.claim_dispatchable(limit=limit) + + async def list_reconcilable_jobs(self, *, limit: int) -> list[GenerationDispatchRecord]: + self.ensure_ready() + return await self.repository.list_reconcilable( + limit=limit, + lease_seconds=self.settings.generation_job_stale_after_seconds, + ) + + async def set_worker_reconciliation_due( + self, + *, + workspace_id: str, + user_id: str, + job_id: str, + due_at: datetime | None, + ) -> None: + self.ensure_ready() + await self.repository.set_reconciliation_due( + workspace_id, + job_id, + user_id=user_id, + due_at=due_at, + ) + + async def fail_stale_unbound_submissions(self) -> int: + self.ensure_ready() + stale_before = datetime.now(timezone.utc) - timedelta( + seconds=self.settings.generation_job_stale_after_seconds + ) + return await self.repository.fail_stale_unbound_submissions(stale_before=stale_before) + + async def transition_worker_job( + self, + *, + workspace_id: str, + user_id: str, + job_id: str, + status: GenerationJobStatus, + error_code: str | None = None, + error_message: str | None = None, + next_attempt_at: datetime | None = None, + ) -> GenerationJobView: + self.ensure_ready() + job = await self.repository.transition_job( + workspace_id, + job_id, + status, + error_code=error_code, + error_message=error_message, + next_attempt_at=next_attempt_at, + user_id=user_id, + ) + return GenerationJobView.from_record(job) + + def submission_retry_at( + self, + *, + attempt_number: int, + category, + http_status: int | None, + ) -> datetime | None: + """Classify a known-not-accepted submission with the shared policy.""" + + from app.generation.domain.retry import GenerationRetryPolicy + + decision = GenerationRetryPolicy( + max_retries=self.settings.ai_worker_max_retries, + backoff_seconds=self.settings.ai_worker_retry_backoff_seconds, + ).decide( + category=category, + http_status=http_status, + retry_number=max(0, attempt_number - 1), + idempotent=True, + ) + if not decision.retryable: + return None + return datetime.now(timezone.utc) + timedelta(seconds=decision.delay_seconds) + + async def close(self) -> None: + await self.providers.close() + + @staticmethod + def _validated_spec( + payload: GenerationRequestCreate, normalized: dict[str, object] + ) -> dict[str, object]: + """Guard adapter output before it becomes durable API-visible input.""" + + if not isinstance(normalized, dict): + raise GenerationCapabilityUnsupportedError( + "Generation adapter returned an invalid validated request." + ) + # A provider cannot rewrite the base intent or smuggle an opaque + # payload into persistence. Provider controls are accepted only when + # a closed typed schema has been reviewed at this boundary. + allowed = {"prompt", "wan", "flux"} + extra = set(normalized) - allowed + if extra: + raise GenerationCapabilityUnsupportedError( + "Generation adapter returned unsupported request fields." + ) + prompt = normalized.get("prompt") + if not isinstance(prompt, str) or prompt != payload.prompt: + raise GenerationCapabilityUnsupportedError( + "Generation adapter returned an invalid prompt." + ) + expected = GenerationService._request_spec(payload) + if normalized != expected: + raise GenerationCapabilityUnsupportedError( + "Generation adapter returned unsupported request controls." + ) + return expected + + @staticmethod + def _request_spec(payload: GenerationRequestCreate) -> dict[str, object]: + """Return the complete typed public generation intent for persistence.""" + + spec: dict[str, object] = {"prompt": payload.prompt} + if payload.wan is not None: + # ``exclude_unset`` retains the worker's audited defaults when a + # caller does not send a parameter, instead of silently inventing + # a MediaRouter-side override. + spec["wan"] = payload.wan.model_dump(exclude_unset=True) + if payload.flux is not None: + # As with WAN, preserve the audited worker defaults for omitted + # controls. The closed Pydantic schema rejects every unsupported + # FLUX field before the request can be persisted or dispatched. + spec["flux"] = payload.flux.model_dump(exclude_unset=True) + return spec + + @staticmethod + def _fingerprint( + *, + provider: str, + model_id: str, + modality: str, + input_asset_id: str | None, + project_id: str | None, + brand_kit_version_id: str | None, + product_surface: str, + spec: dict[str, object], + ) -> str: + canonical = json.dumps( + { + "provider": provider, + "model_id": model_id, + "modality": modality, + "input_asset_id": input_asset_id, + "project_id": project_id, + "brand_kit_version_id": brand_kit_version_id, + "product_surface": product_surface, + "spec": spec, + }, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ) + return hashlib.sha256(canonical.encode("utf-8")).hexdigest() diff --git a/app/generation/services/output_ingestion.py b/app/generation/services/output_ingestion.py new file mode 100644 index 0000000000000000000000000000000000000000..f46695aa31ce0d870c9be08d94bab204102fc54f --- /dev/null +++ b/app/generation/services/output_ingestion.py @@ -0,0 +1,298 @@ +from __future__ import annotations + +import asyncio +import hashlib +import hmac +import struct +from collections.abc import AsyncIterator +from pathlib import Path +from uuid import uuid4 + +import aiofiles + +from app.core.config import Settings +from app.generation.domain.errors import GenerationOutputError +from app.generation.domain.runtime import WorkerOutput, safe_worker_metadata +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.models import CanonicalMediaAsset +from app.services.cleanup import CleanupService +from app.services.ffprobe_service import FFprobeService +from app.services.validator import MediaValidator + + +_OUTPUT_EXTENSIONS = { + "image/apng": ".apng", + "image/avif": ".avif", + "image/bmp": ".bmp", + "image/gif": ".gif", + "image/heic": ".heic", + "image/jpeg": ".jpg", + "image/png": ".png", + "image/tiff": ".tiff", + "image/webp": ".webp", + "video/mp4": ".mp4", + "video/mpeg": ".mpeg", + "video/quicktime": ".mov", + "video/webm": ".webm", +} +_PNG_SIGNATURE = b"\x89PNG\r\n\x1a\n" + + +class GenerationOutputIngestor: + """Streams a validated worker output into MediaRouter's canonical asset store. + + The remote worker supplies only a relative endpoint through ``WorkerOutput``. + Its bytes are staged in a controlled request directory, size/checksum + verified, published through ``CleanupService``, then registered through + ``CanonicalAssetService``. No worker filesystem path, absolute output URL, + or client-provided destination is accepted. + """ + + def __init__( + self, + *, + settings: Settings, + cleanup: CleanupService, + assets: CanonicalAssetService, + ffprobe: FFprobeService | None = None, + validator: MediaValidator | None = None, + ) -> None: + self.settings = settings + self.cleanup = cleanup + self.assets = assets + self.ffprobe = ffprobe + self.validator = validator + + async def ingest( + self, + *, + workspace_id: str, + user_id: str, + generation_request_id: str, + generation_job_id: str, + provider_id: str, + project_id: str | None, + output: WorkerOutput, + chunks: AsyncIterator[bytes], + ) -> CanonicalMediaAsset: + """Create exactly one canonical output asset from a worker byte stream.""" + + filename = self._filename(generation_job_id, output) + workspace = await self.cleanup.create_workspace(generation_request_id) + staging = workspace.outputs / f".{uuid4().hex}.part" + published: Path | None = None + try: + existing = await self._existing_asset( + workspace_id=workspace_id, + generation_request_id=generation_request_id, + filename=filename, + ) + if existing is not None: + return existing + written, digest = await self._write_stream(staging, chunks) + self._validate_received_output(output, written, digest) + media_metadata = await self._validate_media_output( + staging, output=output, written=written + ) + published = await self.cleanup.publish_new(generation_request_id, staging, filename) + if published is None: + # Another process won a recovery race. It owns the file and + # must also have registered it; never replace it from this + # attempt merely because a worker returned different bytes. + existing = await self._existing_asset( + workspace_id=workspace_id, + generation_request_id=generation_request_id, + filename=filename, + ) + if existing is None: + raise GenerationOutputError( + "A concurrent generation output could not be reconciled safely." + ) + return existing + metadata = safe_worker_metadata( + { + "generation": { + "provider": provider_id, + "provider_output_id": output.provider_output_id, + "output_type": output.output_type.value, + "worker_metadata": output.metadata, + "media": media_metadata, + } + } + ) + assert isinstance(metadata, dict) + return await self.assets.register_output( + workspace_id=workspace_id, + user_id=user_id, + request_id=generation_request_id, + path=published, + mime_type=output.mime_type, + metadata=metadata, + project_id=project_id, + ) + except Exception: + # A final file without a canonical record must not remain + # downloadable. ``publish`` atomically moves the staging file, + # so clean either location based on the point of failure. + if published is not None: + await self._remove_file(published) + raise + finally: + await self._remove_file(staging) + await self.cleanup.complete(generation_request_id) + + async def _existing_asset( + self, *, workspace_id: str, generation_request_id: str, filename: str + ) -> CanonicalMediaAsset | None: + try: + asset = await self.assets.get_owned( + workspace_id=workspace_id, + request_id=generation_request_id, + filename=filename, + ) + except CanonicalAssetNotFoundError: + return None + try: + path = self.cleanup.resolve_download(generation_request_id, filename) + await self.assets.verify_file(asset, path) + except CanonicalAssetNotFoundError as exc: + raise GenerationOutputError( + "A previously registered generation output is no longer valid." + ) from exc + return asset + + async def _write_stream( + self, destination: Path, chunks: AsyncIterator[bytes] + ) -> tuple[int, str]: + digest = hashlib.sha256() + written = 0 + try: + async with aiofiles.open(destination, "xb") as stream: + async for chunk in chunks: + if not isinstance(chunk, bytes): + raise GenerationOutputError( + "Generation worker returned a non-binary output stream." + ) + if not chunk: + continue + written += len(chunk) + if written > self.settings.max_upload_size: + raise GenerationOutputError( + "Generation worker output exceeds the configured size limit." + ) + digest.update(chunk) + await stream.write(chunk) + except GenerationOutputError: + raise + except OSError as exc: + raise GenerationOutputError("Generation output could not be stored safely.") from exc + return written, digest.hexdigest() + + @staticmethod + def _filename(generation_job_id: str, output: WorkerOutput) -> str: + extension = _OUTPUT_EXTENSIONS.get(output.mime_type) + if extension is None: + raise GenerationOutputError( + "Generation worker returned an unsupported output MIME type." + ) + # The canonical filename is service-generated and stable per logical + # job; provider-supplied names are descriptive only and never paths. + return f"generation-{generation_job_id}{extension}" + + @staticmethod + def _validate_received_output(output: WorkerOutput, written: int, digest: str) -> None: + if written <= 0: + raise GenerationOutputError("Generation worker returned an empty output.") + if output.byte_size is not None and output.byte_size != written: + raise GenerationOutputError( + "Generation worker output size did not match its descriptor." + ) + if output.sha256 is not None and not hmac.compare_digest(output.sha256, digest): + raise GenerationOutputError( + "Generation worker output checksum did not match its descriptor." + ) + + async def _validate_media_output( + self, path: Path, *, output: WorkerOutput, written: int + ) -> dict[str, object]: + """Run existing media validation on generated video before publication.""" + + if self.validator is not None: + try: + extension = _OUTPUT_EXTENSIONS.get(output.mime_type) + if extension is None: + raise ValueError("unsupported output extension") + self.validator.validate_declared( + f"generation-output{extension}", output.mime_type, written + ) + except Exception as exc: + raise GenerationOutputError( + "Generation worker output has an unsupported media declaration." + ) from exc + if output.output_type.value == "image": + return await self._validate_image_output(path, output=output) + if output.output_type.value != "video": + return {} + if self.ffprobe is None or self.validator is None: + # Unit-level foundation construction can omit process services. + # The production container always supplies both for video models. + return {} + try: + metadata = await self.ffprobe.probe(path) + self.validator.validate_probe(metadata) + if not metadata.get("video_streams"): + raise ValueError("missing video stream") + except Exception as exc: + raise GenerationOutputError( + "Generation worker output is not a readable video." + ) from exc + safe = safe_worker_metadata(metadata) + return safe if isinstance(safe, dict) else {} + + async def _validate_image_output( + self, path: Path, *, output: WorkerOutput + ) -> dict[str, object]: + """Validate the PNG output contract and retain its actual dimensions. + + FLUX's audited worker emits PNG only. Parsing its fixed-size PNG + signature/IHDR header avoids adding a second image-processing runtime + solely for metadata, rejects a mismatched declared MIME type, and + keeps canonical-asset metadata useful to later MediaRouter services. + Other image types remain unavailable until a reviewed provider adds + an equally strict verifier. + """ + + if output.mime_type != "image/png": + raise GenerationOutputError( + "Generation worker returned an unsupported image MIME type." + ) + try: + header = await asyncio.to_thread(self._read_png_header, path) + if len(header) != 24 or header[:8] != _PNG_SIGNATURE or header[12:16] != b"IHDR": + raise ValueError("invalid PNG header") + width, height = struct.unpack(">II", header[16:24]) + if width < 1 or height < 1 or width * height > self.settings.max_resolution_pixels: + raise ValueError("invalid PNG dimensions") + except (OSError, ValueError, struct.error) as exc: + raise GenerationOutputError( + "Generation worker output is not a readable PNG image." + ) from exc + return { + "container": "png", + "codec": "png", + "resolution": {"width": width, "height": height}, + } + + @staticmethod + def _read_png_header(path: Path) -> bytes: + with path.open("rb") as stream: + return stream.read(24) + + @staticmethod + async def _remove_file(path: Path) -> None: + try: + await asyncio.to_thread(path.unlink, missing_ok=True) + except OSError: + # The output is in a controlled per-request directory. Cleanup + # will retry directory removal; do not mask the original failure. + return diff --git a/app/generation/workers/__init__.py b/app/generation/workers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a9ac6a6e4d6f9cf6a048417e148136bb2b444171 --- /dev/null +++ b/app/generation/workers/__init__.py @@ -0,0 +1,6 @@ +"""Durable provider-neutral generation dispatch workers.""" + +from app.generation.workers.generation_worker import GenerationWorker + +__all__ = ["GenerationWorker"] + diff --git a/app/generation/workers/generation_worker.py b/app/generation/workers/generation_worker.py new file mode 100644 index 0000000000000000000000000000000000000000..f00901fa0a63cd81f3f2fbf3bc93c1dda36ee9fe --- /dev/null +++ b/app/generation/workers/generation_worker.py @@ -0,0 +1,601 @@ +"""Provider-neutral durable dispatcher for remote generation workers.""" + +from __future__ import annotations + +import asyncio +from datetime import datetime, timedelta, timezone +from pathlib import Path + +from app.core.config import Settings +from app.core.logger import get_logger +from app.generation.domain.enums import ( + GenerationJobStatus, + WorkerErrorCategory, + WorkerJobStatus, +) +from app.generation.domain.errors import ( + GenerationError, + GenerationOutputError, + GenerationValidationError, + GenerationWorkerError, +) +from app.generation.domain.runtime import WorkerJob +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.repositories.generation import GenerationDispatchRecord +from app.generation.services.generation_service import GenerationService +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService +from app.security.models import CanonicalMediaAsset +from app.services.cleanup import CleanupService + +logger = get_logger(__name__) + + +class GenerationWorker: + """Claim, dispatch, reconcile, and ingest provider-neutral generation jobs. + + This worker owns no provider URLs, credentials, schema, or storage. It + calls the existing GenerationService internal boundary, which preserves + authoritative workspace ownership and the durable state machine. + """ + + def __init__( + self, + *, + settings: Settings, + generation: GenerationService, + assets: CanonicalAssetService, + cleanup: CleanupService, + audit: AuditService, + ) -> None: + self.settings = settings + self.generation = generation + self.assets = assets + self.cleanup = cleanup + self.audit = audit + self._task: asyncio.Task[None] | None = None + self._stop = asyncio.Event() + self._poll_not_before: dict[str, datetime] = {} + self._poll_failures: dict[str, int] = {} + + async def start(self) -> None: + if not self.settings.generation_worker_enabled: + logger.info("generation worker disabled") + return + self._stop.clear() + self._task = asyncio.create_task(self._run(), name="generation-dispatch-worker") + + async def stop(self) -> None: + self._stop.set() + if self._task is not None: + self._task.cancel() + try: + await self._task + except asyncio.CancelledError: + logger.debug("generation worker task cancelled") + self._task = None + + async def run_once(self) -> None: + """Execute one bounded dispatch/reconciliation cycle for tests and runtime.""" + + if not self.settings.generation_enabled or not self.generation.ready: + return + await self._refresh_configured_providers() + stale = await self.generation.fail_stale_unbound_submissions() + if stale: + logger.warning("generation submissions failed as ambiguous", extra={"count": stale}) + await self._reconcile_bound_jobs() + for item in await self.generation.claim_dispatchable_jobs( + limit=self.settings.generation_worker_batch_size + ): + await self._dispatch(item) + + async def _run(self) -> None: + while not self._stop.is_set(): + try: + await self.run_once() + except asyncio.CancelledError: + raise + except Exception: + # A single broken worker/provider cannot take down the ASGI + # process or prevent the next durable jobs from being handled. + logger.exception("generation worker iteration failed") + try: + await asyncio.wait_for( + self._stop.wait(), timeout=self.settings.generation_worker_interval_seconds + ) + except asyncio.TimeoutError: + continue + + async def _refresh_configured_providers(self) -> None: + for adapter in self.generation.providers.list(): + if not adapter.available: + continue + try: + await self.generation.refresh_provider_runtime(adapter.provider) + except Exception: + # refresh_provider_runtime marks models unavailable. The + # worker client already redacts transport errors. + logger.warning( + "generation worker provider unavailable", + extra={"provider": adapter.provider}, + ) + + async def _dispatch(self, item: GenerationDispatchRecord) -> None: + adapter = self.generation.providers.get(item.provider) + bound = False + try: + registered_model = self.generation.get_model(item.provider, item.model_id) + if not registered_model.available: + # No WAN request has been sent, so this is safe to retry as a + # local availability delay rather than a duplicate-risk + # submission. The durable attempt limit still bounds it. + if item.attempt_number <= self.settings.generation_job_retry_limit: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.RETRYING, + error_code="GENERATION_MODEL_UNAVAILABLE", + error_message="Generation model is not ready to accept work.", + next_attempt_at=datetime.now(timezone.utc) + + timedelta(seconds=self.settings.generation_worker_poll_backoff_seconds), + ) + else: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.FAILED, + error_code="GENERATION_MODEL_UNAVAILABLE", + error_message=( + "Generation model did not become ready within the retry limit." + ), + ) + await self._audit_terminal( + item, + "ai.generation_failed", + error_code="GENERATION_MODEL_UNAVAILABLE", + ) + return + current = await self.generation.get_job( + item.workspace_id, item.user_id, item.generation_job_id + ) + if current.status is GenerationJobStatus.CANCEL_REQUESTED: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.CANCELLED, + ) + await self._audit_terminal(item, "ai.generation_cancelled") + return + if current.status is not GenerationJobStatus.SUBMITTING: + return + + input_path: Path | None = None + input_mime_type: str | None = None + if item.input_asset_id is not None: + asset, input_path = await self._resolve_input(item) + input_mime_type = asset.mime_type + remote_job = await adapter.submit( + payload=item.spec, + # This ID is only an internal correlation value for workers + # that support idempotency. An adapter for a worker without + # that protocol disables transport resubmission. + idempotency_key=item.generation_request_id, + input_path=input_path, + input_mime_type=input_mime_type, + ) + await self.generation.bind_provider_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + worker_job_id=remote_job.external_job_id, + provider_metadata={ + "model_id": item.model_id, + "worker_metadata": remote_job.metadata, + }, + ) + bound = True + logger.info( + "generation provider submission accepted", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "model_id": item.model_id, + "attempt_number": item.attempt_number, + }, + ) + await self._apply_remote_job(item, adapter, remote_job) + except asyncio.CancelledError: + raise + except Exception as exc: + if bound: + await self._handle_bound_execution_failure(item, adapter, exc) + else: + await self._handle_submission_failure(item, adapter, exc) + + async def _resolve_input( + self, item: GenerationDispatchRecord + ) -> tuple[CanonicalMediaAsset, Path]: + if item.input_asset_id is None: + raise GenerationValidationError("Generation model requires a canonical input asset.") + try: + asset = await self.assets.get_owned_by_id( + workspace_id=item.workspace_id, + user_id=item.user_id, + asset_id=item.input_asset_id, + ) + path = self.cleanup.resolve_download(asset.request_id, asset.filename) + await self.assets.verify_file(asset, path) + if path.is_symlink() or not path.is_file(): + raise CanonicalAssetNotFoundError("Input asset is unsafe.") + return asset, path + except (CanonicalAssetNotFoundError, OSError) as exc: + raise GenerationValidationError( + "Generation input asset is no longer readable." + ) from exc + + async def _handle_submission_failure( + self, + item: GenerationDispatchRecord, + adapter: GenerationProviderAdapter, + exc: Exception, + ) -> None: + error = self._normalise(adapter, exc) + # A client-side timeout/connection failure, and a gateway 502/504, + # may have reached a non-idempotent remote worker. A provider without + # request lookup/idempotency support must never resubmit that ambiguity. + if self._submission_is_ambiguous(error): + code = "GENERATION_SUBMISSION_AMBIGUOUS" + message = "Worker submission could not be reconciled safely; it was not retried." + destination = GenerationJobStatus.FAILED + else: + retry_at = self._safe_submission_retry(item, error) + if retry_at is not None: + code = "GENERATION_SUBMISSION_RETRYING" + message = "Generation worker did not accept the request; retry scheduled." + destination = GenerationJobStatus.RETRYING + else: + code = "GENERATION_SUBMISSION_FAILED" + message = "Generation worker rejected or failed the submission." + destination = GenerationJobStatus.FAILED + try: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=destination, + error_code=code, + error_message=message, + next_attempt_at=retry_at if destination is GenerationJobStatus.RETRYING else None, + ) + except Exception: + logger.exception( + "generation submission failure could not be persisted", + extra={"generation_job_id": item.generation_job_id, "provider": item.provider}, + ) + return + logger.warning( + "generation provider submission failed", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "model_id": item.model_id, + "attempt_number": item.attempt_number, + "category": error.category.value, + "http_status": error.http_status, + "retrying": destination is GenerationJobStatus.RETRYING, + "ambiguous": destination is GenerationJobStatus.FAILED + and self._submission_is_ambiguous(error), + }, + ) + if destination is GenerationJobStatus.FAILED: + await self._audit_terminal(item, "ai.generation_failed", error_code=code) + + async def _handle_bound_execution_failure( + self, item: GenerationDispatchRecord, adapter: GenerationProviderAdapter, exc: Exception + ) -> None: + """Contain post-bind errors without ever attempting a second submit.""" + + error = self._normalise(adapter, exc) + current = await self.generation.get_job( + item.workspace_id, item.user_id, item.generation_job_id + ) + if current.status in { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + }: + return + if self._is_transient_poll_error(error): + self._defer_poll(item.generation_job_id) + await self.generation.set_worker_reconciliation_due( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + due_at=self._poll_not_before[item.generation_job_id], + ) + return + try: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.FAILED, + error_code="GENERATION_PROVIDER_EXECUTION_FAILED", + error_message="Generation worker execution could not be reconciled.", + ) + await self._audit_terminal( + item, + "ai.generation_failed", + error_code="GENERATION_PROVIDER_EXECUTION_FAILED", + ) + except Exception: + logger.exception( + "generation bound execution failure could not be persisted", + extra={"generation_job_id": item.generation_job_id, "provider": item.provider}, + ) + + async def _reconcile_bound_jobs(self) -> None: + for item in await self.generation.list_reconcilable_jobs( + limit=self.settings.generation_worker_batch_size + ): + if item.external_job_id is None: + continue + adapter = self.generation.providers.get(item.provider) + try: + remote_job = await adapter.get_job(external_job_id=item.external_job_id) + if remote_job.external_job_id != item.external_job_id: + raise GenerationOutputError("Worker returned an unexpected job identity.") + self._clear_poll_backoff(item.generation_job_id) + await self._apply_remote_job(item, adapter, remote_job) + await self.generation.set_worker_reconciliation_due( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + due_at=None, + ) + except asyncio.CancelledError: + raise + except Exception as exc: + error = self._normalise(adapter, exc) + current = await self.generation.get_job( + item.workspace_id, item.user_id, item.generation_job_id + ) + if current.status in { + GenerationJobStatus.SUCCEEDED, + GenerationJobStatus.FAILED, + GenerationJobStatus.CANCELLED, + }: + self._clear_poll_backoff(item.generation_job_id) + continue + if self._is_transient_poll_error(error): + self._defer_poll(item.generation_job_id) + await self.generation.set_worker_reconciliation_due( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + due_at=self._poll_not_before[item.generation_job_id], + ) + logger.warning( + "generation provider poll deferred", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "category": error.category.value, + "http_status": error.http_status, + }, + ) + continue + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.FAILED, + error_code="GENERATION_PROVIDER_STATUS_FAILED", + error_message="Generation worker status reconciliation failed.", + ) + await self._audit_terminal( + item, + "ai.generation_failed", + error_code="GENERATION_PROVIDER_STATUS_FAILED", + ) + logger.warning( + "generation provider reconciliation failed", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "category": error.category.value, + }, + ) + + async def _apply_remote_job( + self, + item: GenerationDispatchRecord, + adapter: GenerationProviderAdapter, + remote_job: WorkerJob, + ) -> None: + current = await self.generation.get_job( + item.workspace_id, item.user_id, item.generation_job_id + ) + current_status = current.status + + if current_status is GenerationJobStatus.CANCEL_REQUESTED and remote_job.status in { + WorkerJobStatus.QUEUED, + WorkerJobStatus.RUNNING, + }: + # A cancellation may have raced the original worker submission. + # Reuse the service cancellation boundary; it will never claim a + # running WAN inference stopped unless the worker confirms it. + try: + await self.generation.cancel( + item.workspace_id, item.user_id, item.generation_job_id + ) + except GenerationError: + pass + return + + if remote_job.status is WorkerJobStatus.QUEUED: + if current_status is GenerationJobStatus.SUBMITTING: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.QUEUED, + ) + return + if remote_job.status is WorkerJobStatus.RUNNING: + if current_status in {GenerationJobStatus.QUEUED, GenerationJobStatus.SUBMITTING}: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.RUNNING, + ) + logger.info( + "generation provider job running", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "model_id": item.model_id, + "attempt_number": item.attempt_number, + }, + ) + return + if remote_job.status is WorkerJobStatus.COMPLETED: + await self.generation.ingest_completed_provider_output( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + worker_job=remote_job, + ) + await self._audit_terminal(item, "ai.generation_completed") + logger.info( + "generation provider job completed", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "model_id": item.model_id, + "attempt_number": item.attempt_number, + }, + ) + return + if remote_job.status is WorkerJobStatus.CANCELLED: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.CANCELLED, + ) + await self._audit_terminal(item, "ai.generation_cancelled") + return + if remote_job.status is WorkerJobStatus.FAILED: + await self.generation.transition_worker_job( + workspace_id=item.workspace_id, + user_id=item.user_id, + job_id=item.generation_job_id, + status=GenerationJobStatus.FAILED, + error_code=remote_job.error_code or "GENERATION_INFERENCE_FAILED", + error_message="Generation worker reported an inference failure.", + ) + await self._audit_terminal( + item, + "ai.generation_failed", + error_code=remote_job.error_code or "GENERATION_INFERENCE_FAILED", + ) + logger.warning( + "generation provider job failed", + extra={ + "generation_job_id": item.generation_job_id, + "provider": item.provider, + "model_id": item.model_id, + "attempt_number": item.attempt_number, + }, + ) + return + raise GenerationOutputError("Generation worker returned an unsupported job status.") + + async def _audit_terminal( + self, + item: GenerationDispatchRecord, + event_type: str, + *, + error_code: str | None = None, + ) -> None: + if item.product_surface != "ai_studio": + return + await self.audit.record_event( + workspace_id=item.workspace_id, + user_id=item.user_id, + request_id=item.generation_request_id, + event_type=event_type, + entity_type="generation_job", + entity_id=item.generation_job_id, + metadata={ + "operation": ("generate_image" if item.modality == "image" else "generate_video"), + "model": item.model_id, + "project_id": item.project_id, + "error_code": error_code, + }, + ) + + @staticmethod + def _normalise(adapter: GenerationProviderAdapter, exc: Exception) -> GenerationWorkerError: + try: + normalized = adapter.normalize_error(exc) + except Exception: + normalized = None + if isinstance(normalized, GenerationWorkerError): + return normalized + return GenerationWorkerError( + category=WorkerErrorCategory.UNKNOWN_ERROR, + message="Generation worker operation failed unexpectedly.", + retryable=False, + ) + + @staticmethod + def _submission_is_ambiguous(error: GenerationWorkerError) -> bool: + return error.category is WorkerErrorCategory.TIMEOUT or ( + error.category is WorkerErrorCategory.WORKER_UNAVAILABLE + and error.http_status in {None, 502, 504} + ) + + def _safe_submission_retry( + self, item: GenerationDispatchRecord, error: GenerationWorkerError + ) -> datetime | None: + # A received 429/503 means the worker did not accept a job. Other + # potentially transient submit failures are treated as ambiguous above + # unless a future worker offers explicit idempotency/reconciliation. + if error.http_status not in {429, 503}: + return None + if item.attempt_number > self.settings.generation_job_retry_limit: + return None + return self.generation.submission_retry_at( + attempt_number=item.attempt_number, + category=error.category, + http_status=error.http_status, + ) + + @staticmethod + def _is_transient_poll_error(error: GenerationWorkerError) -> bool: + return error.category in { + WorkerErrorCategory.WORKER_UNAVAILABLE, + WorkerErrorCategory.WORKER_NOT_READY, + WorkerErrorCategory.TIMEOUT, + WorkerErrorCategory.RATE_LIMITED, + } or error.http_status in {429, 502, 503, 504} + + def _defer_poll(self, job_id: str) -> None: + failures = self._poll_failures.get(job_id, 0) + 1 + self._poll_failures[job_id] = failures + delay = min( + self.settings.generation_worker_poll_backoff_seconds * (2 ** (failures - 1)), + 60.0, + ) + self._poll_not_before[job_id] = datetime.now(timezone.utc) + timedelta(seconds=delay) + + def _clear_poll_backoff(self, job_id: str) -> None: + self._poll_not_before.pop(job_id, None) + self._poll_failures.pop(job_id, None) diff --git a/app/mcp/registry.py b/app/mcp/registry.py index 5483a477bace657dd2a33c39c576f20397cc1ec9..bcdf62ac7b7bb1b87c89251c10029ff428a92fc4 100644 --- a/app/mcp/registry.py +++ b/app/mcp/registry.py @@ -89,6 +89,9 @@ TEMPLATE_TOOLS = [ "template_details", "run_template", "template_categories", + "list_marketplace_templates", + "get_marketplace_template", + "apply_marketplace_template", ] @@ -489,9 +492,7 @@ class MCPRegistry: if lease is not None: await lease.release() if context is not None and not via_http: - await self._audit_mcp( - context, request_id, tool_name, response_code, started - ) + await self._audit_mcp(context, request_id, tool_name, response_code, started) async def _audit_mcp( self, diff --git a/app/mcp/server.py b/app/mcp/server.py index d5a598ec410044c45aac3b96c840ce243cf5f08a..e41b0ce4e5f8076ce67b8e62c132d0c13352c75c 100644 --- a/app/mcp/server.py +++ b/app/mcp/server.py @@ -16,7 +16,10 @@ from app.core.logger import configure_logging from app.mcp.prompts import register_prompts from app.mcp.registry import MCPRegistry from app.mcp.resources import register_resources +from app.mcp.tools.ai import register_ai_tools +from app.mcp.tools.analytics import register_analytics_tools from app.mcp.tools.audio import register_audio_tools +from app.mcp.tools.brand import register_brand_tools from app.mcp.tools.image import register_image_tools from app.mcp.tools.probe import register_probe_tools from app.mcp.tools.social import register_social_tools @@ -58,6 +61,9 @@ def create_mcp_server(container: Container) -> FastMCP[Any]: register_system_tools(server, registry) register_template_tools(server, registry) register_social_tools(server, registry) + register_brand_tools(server, registry) + register_ai_tools(server, registry) + register_analytics_tools(server, registry) register_resources(server, registry) register_prompts(server) return server @@ -71,8 +77,20 @@ async def run_server(transport: Literal["stdio", "streamable-http"] = "stdio") - worker = CleanupWorker(container.cleanup, settings.cleanup_interval_seconds) server = create_mcp_server(container) await container.security_database.initialize() + if not await container.security_database.schema_ready(): + missing = ", ".join(await container.security_database.missing_schema_objects()) + raise RuntimeError( + "Security schema is unavailable; apply app/security/migrations/. " f"Missing: {missing}" + ) + await container.security_database.verify_execution_boundary( + expected_role=settings.security_database_role, + enforce_rls=settings.security_enforce_rls, + ) await container.api_keys.ensure_bootstrap_admin() + await container.tenants.ensure_all_api_key_principals() await container.social.initialize() + await container.analytics.initialize(container.social.ready) + await container.social.adopt_legacy_workspaces(await container.tenants.list_principals()) await worker.start() try: if transport == "stdio": @@ -84,9 +102,7 @@ async def run_server(transport: Literal["stdio", "streamable-http"] = "stdio") - else "" ) if not configured_key: - raise RuntimeError( - "MCP_STDIO_API_KEY is required when AUTH_ENABLED=true" - ) + raise RuntimeError("MCP_STDIO_API_KEY is required when AUTH_ENABLED=true") context = await container.api_keys.authenticate(configured_key) if not context.allows("mcp:read"): raise RuntimeError("MCP_STDIO_API_KEY is missing the mcp:read scope") diff --git a/app/mcp/tools/ai.py b/app/mcp/tools/ai.py new file mode 100644 index 0000000000000000000000000000000000000000..0ca8d0823b7520ee0ce0312afc4aebb358631a23 --- /dev/null +++ b/app/mcp/tools/ai.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +from typing import Any +from uuid import uuid4 + +from mcp.server.fastmcp import FastMCP +from pydantic import TypeAdapter + +from app.ai.schemas import AiGenerationRequest +from app.mcp.registry import MCPRegistry +from app.security.context import auth_context + +_generation_request = TypeAdapter(AiGenerationRequest) + + +def register_ai_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: + """Expose AI Studio through the same authorized provider-neutral service.""" + + def identity() -> tuple[str, str, str | None]: + context = auth_context.get() + if context is None: + return "", "", None + return context.workspace_id or "", context.user_id or "", context.api_key_id + + @server.tool( + name="ai.capabilities", + description="List AI operations and models currently ready in MediaRouter.", + ) + async def ai_capabilities() -> dict[str, Any]: + async def action() -> dict[str, Any]: + return registry.container.ai.capabilities().model_dump(mode="json") + + return await registry.run_metadata_tool("ai.capabilities", action, required_scope="ai:read") + + @server.tool( + name="ai.generate", + description="Submit a real, idempotent AI generation using an advertised capability.", + ) + async def ai_generate(payload: dict[str, Any], idempotency_key: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, user_id, api_key_id = identity() + job = await registry.container.ai.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=str(uuid4()), + payload=_generation_request.validate_python(payload), + idempotency_key=idempotency_key, + ) + return {"job": job.model_dump(mode="json")} + + return await registry.run_metadata_tool("ai.generate", action, required_scope="ai:generate") + + @server.tool(name="ai.list_jobs", description="List workspace AI generation history.") + async def ai_list_jobs(offset: int = 0, limit: int = 25) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, user_id, _ = identity() + history = await registry.container.ai.history( + workspace_id=workspace_id, + user_id=user_id, + offset=offset, + limit=limit, + ) + return history.model_dump(mode="json") + + return await registry.run_metadata_tool("ai.list_jobs", action, required_scope="ai:read") + + @server.tool(name="ai.get_job", description="Get one workspace-owned AI job.") + async def ai_get_job(generation_id: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, user_id, _ = identity() + job = await registry.container.ai.get( + workspace_id=workspace_id, + user_id=user_id, + generation_id=generation_id, + ) + return {"job": job.model_dump(mode="json")} + + return await registry.run_metadata_tool("ai.get_job", action, required_scope="ai:read") + + @server.tool(name="ai.cancel_job", description="Cancel a cancellable AI job.") + async def ai_cancel_job(generation_id: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, user_id, api_key_id = identity() + job = await registry.container.ai.cancel( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=str(uuid4()), + generation_id=generation_id, + ) + return {"job": job.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "ai.cancel_job", action, required_scope="ai:generate" + ) diff --git a/app/mcp/tools/analytics.py b/app/mcp/tools/analytics.py new file mode 100644 index 0000000000000000000000000000000000000000..c58bfed1b3802b8a8c2bd1367b09dd79e0427960 --- /dev/null +++ b/app/mcp/tools/analytics.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from typing import Any + +from mcp.server.fastmcp import FastMCP + +from app.analytics.schemas import AnalyticsQuery, AnalyticsSyncRequest +from app.mcp.registry import MCPRegistry +from app.security.context import auth_context + + +def register_analytics_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: + """Register narrow analytics tools over the shared domain service.""" + + def identity() -> tuple[str, str]: + context = auth_context.get() + return ( + context.workspace_id or "" if context else "", + context.user_id or "" if context else "", + ) + + @server.tool( + name="analytics.overview", description="Get workspace or project analytics totals." + ) + async def analytics_overview(filters: dict[str, Any] | None = None) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, _ = identity() + result = await registry.container.analytics.overview( + workspace_id, AnalyticsQuery.model_validate(filters or {}) + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.overview", action, required_scope="analytics:read" + ) + + @server.tool( + name="analytics.timeseries", description="Get timezone-safe analytics time series." + ) + async def analytics_timeseries(filters: dict[str, Any] | None = None) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, _ = identity() + result = await registry.container.analytics.timeseries( + workspace_id, AnalyticsQuery.model_validate(filters or {}) + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.timeseries", action, required_scope="analytics:read" + ) + + @server.tool(name="analytics.platforms", description="Compare available provider metrics.") + async def analytics_platforms(filters: dict[str, Any] | None = None) -> dict[str, Any]: + return await analytics_overview(filters) + + @server.tool( + name="analytics.top_content", description="List top content by an available metric." + ) + async def analytics_top_content(filters: dict[str, Any] | None = None) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, _ = identity() + query = AnalyticsQuery.model_validate({**(filters or {}), "limit": 10}) + result = await registry.container.analytics.posts(workspace_id, query) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.top_content", action, required_scope="analytics:read" + ) + + @server.tool(name="analytics.post", description="Get analytics for one MediaRouter post.") + async def analytics_post(post_id: str, filters: dict[str, Any] | None = None) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, _ = identity() + result = await registry.container.analytics.post( + workspace_id, post_id, AnalyticsQuery.model_validate(filters or {}) + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.post", action, required_scope="analytics:read" + ) + + @server.tool( + name="analytics.sync", description="Queue an idempotent analytics synchronization." + ) + async def analytics_sync(payload: dict[str, Any]) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, user_id = identity() + result = await registry.container.analytics.create_sync( + workspace_id, user_id, AnalyticsSyncRequest.model_validate(payload) + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.sync", action, required_scope="analytics:sync" + ) + + @server.tool(name="analytics.sync_status", description="Get one analytics sync status.") + async def analytics_sync_status(sync_run_id: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + workspace_id, _ = identity() + run = await registry.container.analytics.repository.get_sync(workspace_id, sync_run_id) + return registry.container.analytics._sync_view(run).model_dump(mode="json") + + return await registry.run_metadata_tool( + "analytics.sync_status", action, required_scope="analytics:read" + ) diff --git a/app/mcp/tools/brand.py b/app/mcp/tools/brand.py new file mode 100644 index 0000000000000000000000000000000000000000..f185fc984b7b4347a27e086ac8b408e952f92383 --- /dev/null +++ b/app/mcp/tools/brand.py @@ -0,0 +1,72 @@ +from __future__ import annotations +from typing import Any +from mcp.server.fastmcp import FastMCP +from app.mcp.registry import MCPRegistry +from app.security.context import auth_context +from app.core.exceptions import InputError + +def register_brand_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: + """Register brand kit discovery, retrieval, and validation tools.""" + + @server.tool(description="List brand kits in the current workspace.") + async def list_brand_kits( + search: str | None = None, + status: str = "active", + offset: int = 0, + limit: int = 20, + ) -> dict[str, Any]: + context = auth_context.get() + if not context or not context.workspace_id or not context.user_id: + raise InputError("Authentication required.") + + async def action() -> dict[str, Any]: + kits = await registry.container.brand_kits.list_brand_kits( + workspace_id=context.workspace_id, + user_id=context.user_id + ) + # Search filtering + items = kits + if search: + items = [k for k in items if search.lower() in k.name.lower()] + + return { + "kits": [k.model_dump() for k in items[offset:offset+limit]], + "total": len(items) + } + + return await registry.run_metadata_tool("list_brand_kits", action) + + @server.tool(description="Get a brand kit by ID.") + async def get_brand_kit(brand_kit_id: str) -> dict[str, Any]: + context = auth_context.get() + if not context or not context.workspace_id or not context.user_id: + raise InputError("Authentication required.") + + async def action() -> dict[str, Any]: + kit = await registry.container.brand_kits.get_brand_kit( + workspace_id=context.workspace_id, + brand_kit_id=brand_kit_id, + user_id=context.user_id + ) + return kit.model_dump() + + return await registry.run_metadata_tool("get_brand_kit", action) + + @server.tool(description="Validate a brand kit version.") + async def validate_brand_kit(brand_kit_id: str) -> dict[str, Any]: + context = auth_context.get() + if not context or not context.workspace_id or not context.user_id: + raise InputError("Authentication required.") + + async def action() -> dict[str, Any]: + # Retrieve latest version + version = await registry.container.brand_kits.get_brand_kit_version( + workspace_id=context.workspace_id, + brand_kit_id=brand_kit_id, + user_id=context.user_id + ) + from app.brand.services.validation_service import BrandKitValidationService + validator = BrandKitValidationService() + return validator.validate(version) + + return await registry.run_metadata_tool("validate_brand_kit", action) diff --git a/app/mcp/tools/collaboration.py b/app/mcp/tools/collaboration.py new file mode 100644 index 0000000000000000000000000000000000000000..cfc72db1e7010e34e250d8cf1389f574a0982a66 --- /dev/null +++ b/app/mcp/tools/collaboration.py @@ -0,0 +1,25 @@ +from __future__ import annotations +from typing import Any +from app.mcp.registry import register_tool +from app.projects.services.collaboration_service import CollaborationService +from app.security.context import auth_context + +@register_tool("collaboration.list_teams") +async def list_teams(ctx: Any) -> list[dict[str, Any]]: + """List all teams in the workspace.""" + auth = auth_context.get() + if not auth or not auth.workspace_id: + raise Exception("Unauthorized") + + service: CollaborationService = ctx.container.collaboration + return await service.list_teams(auth.workspace_id) + +@register_tool("collaboration.create_team") +async def create_team(ctx: Any, name: str) -> dict[str, Any]: + """Create a new team in the workspace.""" + auth = auth_context.get() + if not auth or not auth.workspace_id: + raise Exception("Unauthorized") + + service: CollaborationService = ctx.container.collaboration + return (await service.create_team(auth.workspace_id, name)).model_dump() diff --git a/app/mcp/tools/social.py b/app/mcp/tools/social.py index f0c118f31c7c135d154ea2bbb96dd77ea2f60ecc..1fe0fc50c70192f64e63751356d5e66ca43ba0d9 100644 --- a/app/mcp/tools/social.py +++ b/app/mcp/tools/social.py @@ -6,19 +6,21 @@ from mcp.server.fastmcp import FastMCP from app.mcp.registry import MCPRegistry from app.security.context import auth_context -from app.social.schemas.posts import SocialPostCreate from app.social.schemas.assets import SocialMediaAssetRegister -from app.social.schemas.scheduling import SocialScheduleCreate +from app.social.schemas.operations import PublishingBulkRequest +from app.social.schemas.posts import SocialPostCreate, SocialPostDuplicateRequest +from app.social.schemas.scheduling import SocialRescheduleRequest, SocialScheduleCreate def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: """Register thin MCP transports over the shared SocialService facade.""" - def identity() -> str: + def identity() -> tuple[str, str, str]: context = auth_context.get() if context is None: # _execute normally rejects this first. - return "" - return context.api_key_id + return "", "", "" + # API-key credentials are deliberately not used as tenant IDs. + return context.workspace_id or "", context.user_id or "", context.api_key_id @server.tool( name="social.list_providers", @@ -54,7 +56,8 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - items = await social.media_assets.list(identity(), offset=offset, limit=limit) + workspace_id, _, _ = identity() + items = await social.media_assets.list(workspace_id, offset=offset, limit=limit) return {"assets": [item.model_dump(mode="json") for item in items]} return await registry.run_metadata_tool( @@ -69,8 +72,9 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() + workspace_id, _, _ = identity() item = await social.media_assets.register( - identity(), + workspace_id, SocialMediaAssetRegister(request_id=request_id, filename=filename), ) return {"asset": item.model_dump(mode="json")} @@ -79,12 +83,16 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: "social.register_media_asset", action, required_scope="assets:write" ) - @server.tool(name="social.list_accounts", description="List workspace social accounts and connection states.") + @server.tool( + name="social.list_accounts", + description="List workspace social accounts and connection states.", + ) async def social_list_accounts(offset: int = 0, limit: int = 100) -> dict[str, Any]: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - items = await social.accounts.list(identity(), offset=offset, limit=limit) + workspace_id, _, _ = identity() + items = await social.accounts.list(workspace_id, offset=offset, limit=limit) return {"accounts": [item.model_dump(mode="json") for item in items]} return await registry.run_metadata_tool( @@ -96,7 +104,8 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - item = await social.accounts.get(identity(), account_id) + workspace_id, _, _ = identity() + item = await social.accounts.get(workspace_id, account_id) return {"account": item.model_dump(mode="json")} return await registry.run_metadata_tool( @@ -110,16 +119,16 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - user_id = identity() + workspace_id, user_id, api_key_id = identity() item = await social.publishing.create( - workspace_id=user_id, + workspace_id=workspace_id, user_id=user_id, payload=SocialPostCreate.model_validate(payload), idempotency_key=idempotency_key, ) await social.audit.record( - workspace_id=user_id, - api_key_id=user_id, + workspace_id=workspace_id, + api_key_id=api_key_id, event_type="SOCIAL_POST_CREATED", social_post_id=item.id, ) @@ -129,13 +138,16 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: "social.create_post", action, required_scope="social:posts:write" ) - @server.tool(name="social.publish_post", description="Queue an idempotent social post publication.") + @server.tool( + name="social.publish_post", description="Queue an idempotent social post publication." + ) async def social_publish_post(post_id: str, idempotency_key: str) -> dict[str, Any]: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() + workspace_id, _, _ = identity() jobs = await social.publishing.queue( - identity(), post_id, idempotency_key=idempotency_key + workspace_id, post_id, idempotency_key=idempotency_key ) return {"jobs": [job.model_dump(mode="json") for job in jobs]} @@ -143,15 +155,51 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: "social.publish_post", action, required_scope="social:posts:publish" ) - @server.tool(name="social.schedule_post", description="Schedule a social post using an aware timestamp and IANA timezone.") + @server.tool( + name="social.validate_post", + description="Validate every target of a canonical social post without publishing it.", + ) + async def social_validate_post(post_id: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + result = await social.publishing.validate_post_targets(workspace_id, post_id) + return {"validation": result.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "social.validate_post", action, required_scope="social:posts:write" + ) + + @server.tool( + name="social.list_jobs", + description="List durable workspace publishing jobs and per-target states.", + ) + async def social_list_jobs(offset: int = 0, limit: int = 100) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + items = await social.jobs.list(workspace_id, offset=offset, limit=limit) + return {"jobs": [item.model_dump(mode="json") for item in items]} + + return await registry.run_metadata_tool( + "social.list_jobs", action, required_scope="social:posts:read" + ) + + @server.tool( + name="social.schedule_post", + description="Schedule a social post using an aware timestamp and IANA timezone.", + ) async def social_schedule_post( post_id: str, scheduled_at: str, timezone: str ) -> dict[str, Any]: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() + workspace_id, _, _ = identity() schedule = await social.scheduling.schedule( - identity(), + workspace_id, post_id, SocialScheduleCreate( scheduled_at=scheduled_at, # type: ignore[arg-type] @@ -169,19 +217,46 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - item = await social.publishing.cancel(identity(), post_id) + workspace_id, _, _ = identity() + item = await social.publishing.cancel(workspace_id, post_id) return {"post": item.model_dump(mode="json")} return await registry.run_metadata_tool( "social.cancel_post", action, required_scope="social:posts:write" ) - @server.tool(name="social.get_post", description="Get a multi-target social post and target states.") + @server.tool( + name="social.retry_target", + description="Safely retry one failed publishing target when no provider acceptance is uncertain.", + ) + async def social_retry_target( + post_id: str, target_id: str, idempotency_key: str + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + job = await social.publishing.retry_target( + workspace_id, + post_id, + target_id, + idempotency_key=idempotency_key, + ) + return {"job": job.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "social.retry_target", action, required_scope="social:posts:publish" + ) + + @server.tool( + name="social.get_post", description="Get a multi-target social post and target states." + ) async def social_get_post(post_id: str) -> dict[str, Any]: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - item = await social.publishing.get(identity(), post_id) + workspace_id, _, _ = identity() + item = await social.publishing.get(workspace_id, post_id) return {"post": item.model_dump(mode="json")} return await registry.run_metadata_tool( @@ -193,23 +268,138 @@ def register_social_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() - item = await social.jobs.get(identity(), job_id) + workspace_id, _, _ = identity() + item = await social.jobs.get(workspace_id, job_id) return {"job": item.model_dump(mode="json")} return await registry.run_metadata_tool( "social.get_job", action, required_scope="social:posts:read" ) - @server.tool(name="social.get_analytics", description="Get normalized metrics for an account or post.") + @server.tool( + name="publishing.list_queue", description="List authoritative publishing queue items." + ) + async def publishing_list_queue( + status: str | None = None, provider: str | None = None, limit: int = 50 + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + queue = await social.operations.queue( + workspace_id, offset=0, limit=min(limit, 200), status=status, provider=provider + ) + return {"queue": queue.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "publishing.list_queue", action, required_scope="social:posts:read" + ) + + @server.tool( + name="publishing.list_calendar", + description="List scheduled publishing operations in a UTC range.", + ) + async def publishing_list_calendar(starts_at: str, ends_at: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + from datetime import datetime + + calendar = await social.operations.calendar( + workspace_id, + starts_at=datetime.fromisoformat(starts_at), + ends_at=datetime.fromisoformat(ends_at), + offset=0, + limit=200, + ) + return {"calendar": calendar.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "publishing.list_calendar", action, required_scope="social:schedules:read" + ) + + @server.tool( + name="publishing.reschedule", + description="Reschedule an eligible post with revision protection.", + ) + async def publishing_reschedule( + post_id: str, scheduled_at: str, timezone: str, expected_revision: int + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, _, _ = identity() + result = await social.operations.reschedule( + workspace_id, + post_id, + SocialRescheduleRequest( + scheduled_at=scheduled_at, + timezone=timezone, + expected_revision=expected_revision, + ), + ) + return {"schedule": result.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "publishing.reschedule", action, required_scope="social:schedules:write" + ) + + @server.tool( + name="publishing.duplicate", + description="Duplicate a publishing post without provider execution state.", + ) + async def publishing_duplicate( + post_id: str, expected_revision: int, idempotency_key: str + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, user_id, _ = identity() + result = await social.operations.duplicate( + workspace_id, + user_id, + post_id, + SocialPostDuplicateRequest(expected_revision=expected_revision), + idempotency_key=idempotency_key, + ) + return {"post": result.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "publishing.duplicate", action, required_scope="social:posts:write" + ) + + @server.tool(name="publishing.bulk", description="Queue a bounded durable publishing batch.") + async def publishing_bulk(payload: dict[str, Any], idempotency_key: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + social = registry.container.social + social.ensure_ready() + workspace_id, user_id, _ = identity() + result = await social.operations.create_batch( + workspace_id, + user_id, + PublishingBulkRequest.model_validate(payload), + idempotency_key=idempotency_key, + ) + return {"batch": result.model_dump(mode="json")} + + return await registry.run_metadata_tool( + "publishing.bulk", action, required_scope="social:schedules:write" + ) + + @server.tool( + name="social.get_analytics", description="Get normalized metrics for an account or post." + ) async def social_get_analytics( resource: Literal["account", "post"], resource_id: str ) -> dict[str, Any]: async def action() -> dict[str, Any]: social = registry.container.social social.ensure_ready() + workspace_id, _, _ = identity() if resource == "account": - return await social.analytics.account(identity(), resource_id) - return await social.analytics.post(identity(), resource_id) + return await social.analytics.account(workspace_id, resource_id) + return await social.analytics.post(workspace_id, resource_id) return await registry.run_metadata_tool( "social.get_analytics", action, required_scope="social:analytics:read" diff --git a/app/mcp/tools/templates.py b/app/mcp/tools/templates.py index 33e639a24b158a3b7e75974e093d2bf298b40bd6..c95ec8f80295533a3f0281e7ebdece2ecf0fb082 100644 --- a/app/mcp/tools/templates.py +++ b/app/mcp/tools/templates.py @@ -5,6 +5,9 @@ from typing import Any from mcp.server.fastmcp import FastMCP from app.mcp.registry import MCPRegistry, MediaInput +from app.security.context import auth_context +from app.templates.marketplace_errors import MarketplaceTemplateForbidden +from app.templates.marketplace_schemas import SlotBinding, TemplateApply def register_template_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: @@ -47,3 +50,93 @@ def register_template_tools(server: FastMCP[Any], registry: MCPRegistry) -> None return {"categories": registry.container.template_registry.categories()} return await registry.run_metadata_tool("template_categories", action) + + @server.tool(description="List visible versioned Content Studio marketplace templates.") + async def list_marketplace_templates( + search: str | None = None, + category: str | None = None, + limit: int = 24, + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + context = auth_context.get() + if context is None or not context.workspace_id or not context.user_id: + return {"items": [], "total": 0} + result = await registry.container.template_marketplace.list( + workspace_id=context.workspace_id, + user_id=context.user_id, + search=search, + category=category, + aspect_ratio=None, + min_duration_ms=None, + max_duration_ms=None, + media_type=None, + visibility=None, + status=None, + capability=None, + available_only=False, + offset=0, + limit=min(max(limit, 1), 100), + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "list_marketplace_templates", action, required_scope="templates:read" + ) + + @server.tool(description="Get one visible versioned marketplace template.") + async def get_marketplace_template(template_id: str) -> dict[str, Any]: + async def action() -> dict[str, Any]: + context = auth_context.get() + if context is None or not context.workspace_id or not context.user_id: + return {} + result = await registry.container.template_marketplace.get( + workspace_id=context.workspace_id, + user_id=context.user_id, + template_id=template_id, + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "get_marketplace_template", action, required_scope="templates:read" + ) + + @server.tool(description="Apply a marketplace template to an existing project.") + async def apply_marketplace_template( + template_id: str, + project_id: str, + slot_bindings: dict[str, dict[str, str]], + idempotency_key: str, + template_version_id: str | None = None, + ) -> dict[str, Any]: + async def action() -> dict[str, Any]: + context = auth_context.get() + if context is None or not context.workspace_id or not context.user_id: + return {} + if not context.allows("projects:update"): + raise MarketplaceTemplateForbidden("projects:update permission is required") + if any( + value.get("asset_id") for value in slot_bindings.values() + ) and not context.allows("assets:read"): + raise MarketplaceTemplateForbidden("assets:read permission is required") + payload = TemplateApply( + project_id=project_id, + template_version_id=template_version_id, + slot_bindings={ + key: SlotBinding.model_validate(value) for key, value in slot_bindings.items() + }, + ) + result = await registry.container.template_marketplace.apply( + workspace_id=context.workspace_id, + user_id=context.user_id, + api_key_id=context.api_key_id, + request_id=idempotency_key[:36], + template_id=template_id, + payload=payload, + idempotency_key=idempotency_key, + instantiate=False, + ) + return result.model_dump(mode="json") + + return await registry.run_metadata_tool( + "apply_marketplace_template", action, required_scope="templates:apply" + ) diff --git a/app/projects/__init__.py b/app/projects/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..9c00ab39e750b71c7c226e243b45215eb57509f9 --- /dev/null +++ b/app/projects/__init__.py @@ -0,0 +1 @@ +"""Authoritative workspace-owned project domain.""" diff --git a/app/projects/api.py b/app/projects/api.py new file mode 100644 index 0000000000000000000000000000000000000000..582e55df285e7e7a151a84ac17692371f8ac6f15 --- /dev/null +++ b/app/projects/api.py @@ -0,0 +1,431 @@ +from __future__ import annotations + +from typing import Annotated, Any +from uuid import UUID + +from fastapi import APIRouter, Header, Path, Query, Request, Response, status + +from app.projects.editor_schemas import ( + EditorSaveRequest, + EditorStateResponse, + ProjectRenderCreate, + ProjectRenderListResponse, + ProjectRenderResponse, +) +from app.projects.schemas import ( + ProjectAssetAttach, + ProjectAssetListResponse, + ProjectAssetResponse, + ProjectCreate, + ProjectGenerationJobAttach, + ProjectGenerationJobListResponse, + ProjectGenerationJobResponse, + ProjectListResponse, + ProjectResponse, + ProjectStatus, + ProjectUpdate, +) +from app.projects.schemas.collaboration import ( + TeamResponse, + InvitationCreate, + InvitationResponse, + TeamBase, + MemberResponse +) +from app.projects.schemas.approval import ApprovalRequest, ReviewComment +from app.security.errors import ForbiddenError + +router = APIRouter(prefix="/v1/projects", tags=["projects"]) + +@router.post("/workspace/teams", response_model=TeamResponse) +async def create_team(request: Request, payload: TeamBase) -> TeamResponse: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.create_team(workspace_id, payload.name) + +@router.get("/workspace/teams", response_model=list[TeamResponse]) +async def list_teams(request: Request) -> list[TeamResponse]: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.list_teams(workspace_id) + +@router.get("/workspace/members", response_model=list[MemberResponse]) +async def list_members(request: Request) -> list[MemberResponse]: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.list_members(workspace_id) + +@router.get("/workspace/notifications/preferences", response_model=list[dict[str, Any]]) +async def list_notification_preferences(request: Request) -> list[dict[str, Any]]: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.notifications.get_preferences(workspace_id, user_id) + +@router.post("/workspace/notifications/preferences", response_model=dict[str, Any]) +async def update_notification_preference(request: Request, payload: dict[str, Any]) -> dict[str, Any]: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.notifications.update_preference( + workspace_id, user_id, payload["event_type"], payload["enabled"] + ) + +@router.post("/workspace/invitations", response_model=InvitationResponse) +async def invite_member(request: Request, payload: InvitationCreate) -> InvitationResponse: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.invite_member(workspace_id, payload.email, payload.role) + +@router.patch("/workspace/teams/{team_id}", response_model=TeamResponse) +async def update_team(request: Request, team_id: str, payload: TeamBase) -> TeamResponse: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.update_team(workspace_id, team_id, payload.name) + +@router.delete("/workspace/teams/{team_id}", status_code=status.HTTP_204_NO_CONTENT) +async def archive_team(request: Request, team_id: str) -> Response: + workspace_id, _, _, _ = _identity(request) + await request.app.state.container.collaboration.archive_team(workspace_id, team_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + +@router.get("/workspace/workflows/{workflow_id}/requests", response_model=list[ApprovalRequest]) +async def list_approval_requests(request: Request, workflow_id: str) -> list[ApprovalRequest]: + return await request.app.state.container.approval.list_requests(workflow_id) + + +@router.post("/workspace/requests/{request_id}/approve", response_model=ApprovalRequest) +async def approve_request( + request: Request, + request_id: str +) -> ApprovalRequest: + _, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.approve_request(request_id, user_id) + +@router.post("/workspace/requests/{request_id}/reject", response_model=ApprovalRequest) +async def reject_request( + request: Request, + request_id: str +) -> ApprovalRequest: + _, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.reject_request(request_id, user_id) + +@router.post("/workspace/requests/{request_id}/comments", response_model=ReviewComment) +async def add_review_comment( + request: Request, + request_id: str, + content: str +) -> ReviewComment: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.add_comment( + request_id, user_id, workspace_id, content + ) + +@router.delete("/workspace/members/{user_id}", status_code=status.HTTP_204_NO_CONTENT) +async def remove_member(request: Request, user_id: str) -> Response: + workspace_id, actor_user_id, _, _ = _identity(request) + await request.app.state.container.collaboration.remove_member(workspace_id, actor_user_id, user_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + +@router.patch("/workspace/members/{user_id}/role", response_model=MemberResponse) +async def update_member_role(request: Request, user_id: str, new_role: str) -> MemberResponse: + workspace_id, actor_user_id, _, _ = _identity(request) + await request.app.state.container.collaboration.update_member_role(workspace_id, actor_user_id, user_id, new_role) + # Return updated membership + membership = await request.app.state.container.collaboration.get_membership(workspace_id, user_id) + return MemberResponse(id=membership.id, workspace_id=membership.workspace_id, user_id=membership.user_id, role=membership.role, created_at=membership.created_at.isoformat()) + + +def _identity(request: Request) -> tuple[str, str, str, str]: + context = request.state.auth + if not context.workspace_id or not context.user_id: + # API-key middleware resolves this authoritative membership. Client + # request bodies and query strings are never tenant selectors. + raise ForbiddenError + return ( + context.workspace_id, + context.user_id, + context.api_key_id, + request.state.request_id, + ) + + +@router.get("", response_model=ProjectListResponse) +async def list_projects( + request: Request, + project_status: Annotated[ProjectStatus | None, Query(alias="status")] = ProjectStatus.ACTIVE, + search: Annotated[str | None, Query(min_length=1, max_length=200)] = None, + limit: Annotated[int, Query(ge=1, le=100)] = 50, + cursor: Annotated[str | None, Query(min_length=1, max_length=1024)] = None, +) -> ProjectListResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.projects.list( + workspace_id=workspace_id, + user_id=user_id, + status=project_status, + search=search, + limit=limit, + cursor=cursor, + ) + + +@router.post("", response_model=ProjectResponse, status_code=status.HTTP_201_CREATED) +async def create_project(request: Request, payload: ProjectCreate) -> ProjectResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.projects.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + payload=payload, + ) + + +@router.get("/{project_id}/editor", response_model=EditorStateResponse) +async def get_project_editor( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> EditorStateResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.editor.get( + workspace_id=workspace_id, user_id=user_id, project_id=str(project_id) + ) + + +@router.put("/{project_id}/editor", response_model=EditorStateResponse) +async def save_project_editor( + request: Request, + payload: EditorSaveRequest, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> EditorStateResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.editor.save( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + payload=payload, + ) + + +@router.get("/{project_id}/renders", response_model=ProjectRenderListResponse) +async def list_project_renders( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectRenderListResponse: + workspace_id, user_id, _, _ = _identity(request) + return ProjectRenderListResponse( + items=await request.app.state.container.renders.list( + workspace_id=workspace_id, user_id=user_id, project_id=str(project_id) + ) + ) + + +@router.post( + "/{project_id}/renders", + response_model=ProjectRenderResponse, + status_code=status.HTTP_202_ACCEPTED, +) +async def create_project_render( + request: Request, + payload: ProjectRenderCreate, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], + idempotency_key: Annotated[str | None, Header(alias="Idempotency-Key")] = None, +) -> ProjectRenderResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + request.app.state.container.api_keys.authorize(request.state.auth, "jobs:create") + return await request.app.state.container.renders.create( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + payload=payload, + idempotency_key=idempotency_key or "", + ) + + +@router.get("/{project_id}/renders/{render_id}", response_model=ProjectRenderResponse) +async def get_project_render( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], + render_id: Annotated[UUID, Path(description="Canonical render job UUID")], +) -> ProjectRenderResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.renders.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(project_id), + render_id=str(render_id), + ) + + +@router.post("/{project_id}/renders/{render_id}/cancel", response_model=ProjectRenderResponse) +async def cancel_project_render( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], + render_id: Annotated[UUID, Path(description="Canonical render job UUID")], +) -> ProjectRenderResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + request.app.state.container.api_keys.authorize(request.state.auth, "jobs:cancel") + return await request.app.state.container.renders.cancel( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + render_id=str(render_id), + ) + + +@router.get("/{project_id}/assets", response_model=ProjectAssetListResponse) +async def list_project_assets( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectAssetListResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.projects.list_assets( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(project_id), + ) + + +@router.post( + "/{project_id}/assets", + response_model=ProjectAssetResponse, + status_code=status.HTTP_201_CREATED, +) +async def attach_project_asset( + request: Request, + payload: ProjectAssetAttach, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectAssetResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.projects.attach_asset( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + asset_id=payload.asset_id, + ) + + +@router.get("/{project_id}/collaborators", response_model=list[MemberResponse]) +async def list_project_collaborators( + request: Request, + project_id: UUID +) -> list[MemberResponse]: + return await request.app.state.container.collaboration.list_project_collaborators(str(project_id)) + +@router.post("/{project_id}/collaborators", response_model=MemberResponse) +async def add_project_collaborator( + request: Request, + project_id: UUID, + user_id: str, + role: str +) -> MemberResponse: + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.collaboration.add_project_collaborator( + workspace_id, str(project_id), user_id, role + ) + +@router.delete("/{project_id}/collaborators/{user_id}", status_code=status.HTTP_204_NO_CONTENT) +async def remove_project_collaborator( + request: Request, + project_id: UUID, + user_id: str +) -> Response: + await request.app.state.container.collaboration.remove_project_collaborator(str(project_id), user_id) + return Response(status_code=status.HTTP_204_NO_CONTENT) + + +@router.get("/{project_id}/jobs", response_model=ProjectGenerationJobListResponse) +async def list_project_generation_jobs( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectGenerationJobListResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.projects.list_generation_jobs( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(project_id), + ) + + +@router.post( + "/{project_id}/jobs", + response_model=ProjectGenerationJobResponse, + status_code=status.HTTP_201_CREATED, +) +async def attach_project_generation_job( + request: Request, + payload: ProjectGenerationJobAttach, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectGenerationJobResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.projects.attach_generation_job( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + generation_job_id=payload.generation_job_id, + ) + + +@router.delete("/{project_id}/jobs/{job_id}", status_code=status.HTTP_204_NO_CONTENT) +async def detach_project_generation_job( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], + job_id: Annotated[UUID, Path(description="Durable generation job UUID")], +) -> Response: + workspace_id, user_id, api_key_id, request_id = _identity(request) + await request.app.state.container.projects.detach_generation_job( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + generation_job_id=str(job_id), + ) + return Response(status_code=status.HTTP_204_NO_CONTENT) + + +@router.get("/{project_id}", response_model=ProjectResponse) +async def get_project( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectResponse: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.projects.get( + workspace_id=workspace_id, + user_id=user_id, + project_id=str(project_id), + ) + + +@router.patch("/{project_id}", response_model=ProjectResponse) +async def update_project( + request: Request, + payload: ProjectUpdate, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> ProjectResponse: + workspace_id, user_id, api_key_id, request_id = _identity(request) + return await request.app.state.container.projects.update( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + payload=payload, + ) + + +@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT) +async def delete_project( + request: Request, + project_id: Annotated[UUID, Path(description="Canonical project UUID")], +) -> Response: + workspace_id, user_id, api_key_id, request_id = _identity(request) + await request.app.state.container.projects.delete( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=str(project_id), + ) + return Response(status_code=status.HTTP_204_NO_CONTENT) diff --git a/app/projects/editor_schemas.py b/app/projects/editor_schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..d3757e29af31d90c7c53f50c0fde75f0a11dadf5 --- /dev/null +++ b/app/projects/editor_schemas.py @@ -0,0 +1,295 @@ +from __future__ import annotations + +import json +from datetime import datetime +from typing import Annotated, Literal + +from pydantic import ( + BaseModel, + ConfigDict, + Field, + StrictBool, + StrictFloat, + StrictInt, + model_validator, +) + +EDITOR_SCHEMA_VERSION = 1 +EDITOR_STATE_MAX_BYTES = 1_048_576 +EDITOR_ID_MAX_LENGTH = 128 +EDITOR_LABEL_MAX_LENGTH = 500 + +MetadataValue = str | StrictInt | StrictFloat | StrictBool | None +Metadata = dict[str, MetadataValue] + + +class EditorModel(BaseModel): + model_config = ConfigDict(extra="forbid", populate_by_name=True) + + +class ClipTransform(EditorModel): + x: float = Field(allow_inf_nan=False, ge=-100_000, le=100_000) + y: float = Field(allow_inf_nan=False, ge=-100_000, le=100_000) + scale_x: float = Field(alias="scaleX", allow_inf_nan=False, gt=0, le=100) + scale_y: float = Field(alias="scaleY", allow_inf_nan=False, gt=0, le=100) + rotation: float = Field(allow_inf_nan=False, ge=-36_000, le=36_000) + + +class BaseClip(EditorModel): + id: str = Field(min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + track_id: str = Field(alias="trackId", min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + label: str = Field(min_length=1, max_length=EDITOR_LABEL_MAX_LENGTH) + start_ms: int = Field(alias="startMs", ge=0) + duration_ms: int = Field(alias="durationMs", gt=0) + visible: StrictBool + opacity: float = Field(allow_inf_nan=False, ge=0, le=1) + metadata: Metadata = Field(default_factory=dict, max_length=100) + + +class SourceClip(BaseClip): + asset_id: str = Field(alias="assetId", min_length=1, max_length=36) + source_start_ms: int = Field(alias="sourceStartMs", ge=0) + source_duration_ms: int = Field(alias="sourceDurationMs", gt=0) + + @model_validator(mode="after") + def validate_source_range(self) -> SourceClip: + if self.source_duration_ms != self.duration_ms: + raise ValueError("Source duration must match timeline duration") + return self + + +class MediaClip(SourceClip): + kind: Literal["media"] + media_type: Literal["video", "image"] = Field(alias="mediaType") + transform: ClipTransform + volume: float = Field(allow_inf_nan=False, ge=0, le=1) + + +class AudioClip(SourceClip): + kind: Literal["audio"] + volume: float = Field(allow_inf_nan=False, ge=0, le=1) + fade_in_ms: int = Field(alias="fadeInMs", ge=0) + fade_out_ms: int = Field(alias="fadeOutMs", ge=0) + + @model_validator(mode="after") + def validate_fades(self) -> AudioClip: + if self.fade_in_ms > self.duration_ms or self.fade_out_ms > self.duration_ms: + raise ValueError("Audio fades cannot exceed clip duration") + return self + + +class CaptionStyle(EditorModel): + align: Literal["left", "center", "right"] + position: Literal["top", "center", "bottom"] + + +class CaptionClip(BaseClip): + kind: Literal["caption"] + text: str = Field(max_length=10_000) + style: CaptionStyle + + +class EffectClip(BaseClip): + kind: Literal["effect"] + effect_type: str = Field(alias="effectType", min_length=1, max_length=100) + target_clip_id: str | None = Field( + default=None, alias="targetClipId", min_length=1, max_length=EDITOR_ID_MAX_LENGTH + ) + parameters: Metadata = Field(default_factory=dict, max_length=100) + + +TimelineClip = Annotated[ + MediaClip | AudioClip | CaptionClip | EffectClip, Field(discriminator="kind") +] + + +class Track(EditorModel): + id: str = Field(min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + type: Literal["video", "audio", "caption", "overlay"] + name: str = Field(min_length=1, max_length=200) + order: int = Field(ge=0) + muted: StrictBool + locked: StrictBool + visible: StrictBool + clips: list[TimelineClip] + + @model_validator(mode="after") + def validate_clip_ownership(self) -> Track: + for clip in self.clips: + if clip.track_id != self.id: + raise ValueError("Clip trackId must reference its containing track") + valid = { + "video": {"media"}, + "audio": {"audio"}, + "caption": {"caption"}, + "overlay": {"media", "caption", "effect"}, + }[self.type] + if clip.kind not in valid: + raise ValueError(f"Clip kind {clip.kind} is invalid for a {self.type} track") + return self + + +class Transition(EditorModel): + id: str = Field(min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + type: str = Field(min_length=1, max_length=100) + from_clip_id: str = Field(alias="fromClipId", min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + to_clip_id: str = Field(alias="toClipId", min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + duration_ms: int = Field(alias="durationMs", gt=0) + metadata: Metadata = Field(default_factory=dict, max_length=100) + + +class Marker(EditorModel): + id: str = Field(min_length=1, max_length=EDITOR_ID_MAX_LENGTH) + time_ms: int = Field(alias="timeMs", ge=0) + label: str = Field(min_length=1, max_length=500) + color_token: str | None = Field(default=None, alias="colorToken", max_length=100) + + +class Timeline(EditorModel): + time_unit: Literal["milliseconds"] = Field(alias="timeUnit") + tracks: list[Track] + transitions: list[Transition] + markers: list[Marker] + + @model_validator(mode="after") + def validate_structure(self) -> Timeline: + track_ids = [track.id for track in self.tracks] + track_orders = [track.order for track in self.tracks] + clip_ids = [clip.id for track in self.tracks for clip in track.clips] + marker_ids = [marker.id for marker in self.markers] + transition_ids = [item.id for item in self.transitions] + for label, values in ( + ("track", track_ids), + ("track order", track_orders), + ("clip", clip_ids), + ("marker", marker_ids), + ("transition", transition_ids), + ): + if len(values) != len(set(values)): + raise ValueError(f"Duplicate {label} identifiers are not allowed") + if track_orders and set(track_orders) != set(range(len(track_orders))): + raise ValueError("Track order must be contiguous and zero-based") + known_clips = set(clip_ids) + for transition in self.transitions: + if ( + transition.from_clip_id not in known_clips + or transition.to_clip_id not in known_clips + ): + raise ValueError("Transition clip references are invalid") + if transition.from_clip_id == transition.to_clip_id: + raise ValueError("A transition must reference two clips") + for track in self.tracks: + for clip in track.clips: + if isinstance(clip, EffectClip) and ( + clip.target_clip_id is not None and clip.target_clip_id not in known_clips + ): + raise ValueError("Effect target clip reference is invalid") + return self + + +class EditorRenderSettings(EditorModel): + format: Literal["mp4", "webm"] + width: int = Field(ge=2) + height: int = Field(ge=2) + frame_rate: float = Field(alias="frameRate", allow_inf_nan=False, gt=0, le=120) + + @model_validator(mode="after") + def validate_even_resolution(self) -> EditorRenderSettings: + if self.width % 2 or self.height % 2: + raise ValueError("Render width and height must be even") + return self + + +class EditorDocument(EditorModel): + schema_version: Literal[EDITOR_SCHEMA_VERSION] = Field(alias="schemaVersion") + project_id: str = Field(alias="projectId", min_length=1, max_length=36) + timeline: Timeline + render_settings: EditorRenderSettings = Field(alias="renderSettings") + + def json_bytes(self) -> bytes: + return json.dumps( + self.model_dump(by_alias=True), + ensure_ascii=False, + allow_nan=False, + separators=(",", ":"), + ).encode("utf-8") + + def asset_ids(self) -> set[str]: + return { + clip.asset_id + for track in self.timeline.tracks + for clip in track.clips + if isinstance(clip, (MediaClip, AudioClip)) + } + + def duration_ms(self) -> int: + return max( + ( + clip.start_ms + clip.duration_ms + for track in self.timeline.tracks + for clip in track.clips + ), + default=0, + ) + + +class EditorSaveRequest(EditorModel): + expected_revision: int = Field(ge=0) + schema_version: Literal[EDITOR_SCHEMA_VERSION] + state: EditorDocument + + @model_validator(mode="after") + def validate_versions(self) -> EditorSaveRequest: + if self.state.schema_version != self.schema_version: + raise ValueError("Editor schema versions do not match") + if len(self.state.json_bytes()) > EDITOR_STATE_MAX_BYTES: + raise ValueError("Editor state exceeds maximum allowed size") + return self + + +class EditorStateResponse(BaseModel): + project_id: str + revision: int = Field(ge=1) + schema_version: int + state: EditorDocument + created_at: datetime + updated_at: datetime + updated_by: str + + +class ProjectRenderCreate(EditorModel): + editor_revision: int = Field(ge=1) + output_format: Literal["mp4", "webm"] = "mp4" + width: int = Field(default=1920, ge=2) + height: int = Field(default=1080, ge=2) + frame_rate: float = Field(default=30, allow_inf_nan=False, gt=0, le=120) + quality: Literal["draft", "standard", "high"] = "standard" + preset: Literal["fast", "balanced", "quality"] = "balanced" + + @model_validator(mode="after") + def validate_resolution(self) -> ProjectRenderCreate: + if self.width % 2 or self.height % 2: + raise ValueError("Render width and height must be even") + return self + + +class ProjectRenderResponse(BaseModel): + id: str + project_id: str + editor_revision: int + status: Literal["queued", "processing", "completed", "failed", "cancelling", "cancelled"] + render_settings: dict[str, object] + progress: None = None + output_asset_id: str | None + error_code: str | None + error_message: str | None + attempt_count: int + created_at: datetime + started_at: datetime | None + completed_at: datetime | None + cancelled_at: datetime | None + updated_at: datetime + + +class ProjectRenderListResponse(BaseModel): + items: list[ProjectRenderResponse] diff --git a/app/projects/errors.py b/app/projects/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..8c5cf7cb580c4921241934d909f356b7ffab00f8 --- /dev/null +++ b/app/projects/errors.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from app.core.exceptions import MediaAPIError + + +class ProjectNotFoundError(MediaAPIError): + code = "PROJECT_NOT_FOUND" + status_code = 404 + + +class ProjectInvalidNameError(MediaAPIError): + code = "PROJECT_INVALID_NAME" + status_code = 422 + + +class ProjectInvalidStatusError(MediaAPIError): + code = "PROJECT_INVALID_STATUS" + status_code = 422 + + +class ProjectThumbnailInvalidError(MediaAPIError): + code = "PROJECT_THUMBNAIL_INVALID" + status_code = 422 + + +class ProjectAlreadyArchivedError(MediaAPIError): + code = "PROJECT_ALREADY_ARCHIVED" + status_code = 409 + + +class ProjectInvalidCursorError(MediaAPIError): + code = "PROJECT_INVALID_CURSOR" + status_code = 422 + + +class ProjectAssetNotFoundError(MediaAPIError): + code = "PROJECT_ASSET_NOT_FOUND" + status_code = 404 + + +class ProjectAssetConflictError(MediaAPIError): + code = "PROJECT_ASSET_CONFLICT" + status_code = 409 + + +class ProjectJobNotFoundError(MediaAPIError): + code = "PROJECT_JOB_NOT_FOUND" + status_code = 404 + + +class ProjectJobConflictError(MediaAPIError): + code = "PROJECT_JOB_CONFLICT" + status_code = 409 + + +class ProjectEditorNotFoundError(MediaAPIError): + code = "PROJECT_EDITOR_NOT_FOUND" + status_code = 404 + + +class ProjectEditorConflictError(MediaAPIError): + code = "PROJECT_EDITOR_REVISION_CONFLICT" + status_code = 409 + + +class ProjectEditorInvalidError(MediaAPIError): + code = "PROJECT_EDITOR_INVALID" + status_code = 422 + + +class ProjectEditorAssetInvalidError(MediaAPIError): + code = "PROJECT_EDITOR_ASSET_INVALID" + status_code = 422 + + +class ProjectRenderNotFoundError(MediaAPIError): + code = "PROJECT_RENDER_NOT_FOUND" + status_code = 404 + + +class ProjectRenderConflictError(MediaAPIError): + code = "PROJECT_RENDER_CONFLICT" + status_code = 409 + + +class ProjectRenderInvalidError(MediaAPIError): + code = "PROJECT_RENDER_INVALID" + status_code = 422 + + +class ProjectRenderUnsupportedError(MediaAPIError): + code = "PROJECT_RENDER_UNSUPPORTED_FEATURE" + status_code = 422 + + +class ProjectRenderLimitError(MediaAPIError): + code = "PROJECT_RENDER_LIMIT_EXCEEDED" + status_code = 429 + + +class ProjectRenderTransitionError(RuntimeError): + """Internal invariant failure; never accept state transitions from clients.""" + +class CollaborationUnauthorizedError(MediaAPIError): + code = "COLLABORATION_UNAUTHORIZED" + status_code = 403 + + +class TeamNotFoundError(MediaAPIError): + code = "TEAM_NOT_FOUND" + status_code = 404 + + +class TeamMemberNotFoundError(MediaAPIError): + code = "TEAM_MEMBER_NOT_FOUND" + status_code = 404 + + +class ApprovalRequestNotFoundError(MediaAPIError): + code = "APPROVAL_REQUEST_NOT_FOUND" + status_code = 404 + + +class ApprovalWorkflowStateError(MediaAPIError): + code = "APPROVAL_WORKFLOW_STATE_ERROR" + status_code = 409 + + +class ApprovalSeparationOfDutiesError(MediaAPIError): + code = "APPROVAL_SEPARATION_OF_DUTIES_ERROR" + status_code = 403 + + +class InvitationNotFoundError(MediaAPIError): + code = "INVITATION_NOT_FOUND" + status_code = 404 + + +class InvitationExpiredError(MediaAPIError): + code = "INVITATION_EXPIRED" + status_code = 409 diff --git a/app/projects/migrations/0001_projects_foundation.sql b/app/projects/migrations/0001_projects_foundation.sql new file mode 100644 index 0000000000000000000000000000000000000000..96f8a35e0c19ae5e616b7005e84a5c0cced8e6ac --- /dev/null +++ b/app/projects/migrations/0001_projects_foundation.sql @@ -0,0 +1,162 @@ +-- Authoritative workspace-owned Project foundation. +-- Apply after app/security/migrations/0002_authoritative_tenancy_postgres.sql. +-- IDs remain UUID values stored as text to match MediaRouter's existing +-- tenancy and canonical-asset foreign-key types. + +begin; +create extension if not exists pgcrypto; + +create table if not exists projects ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + created_by text not null references users(id) on delete restrict, + name text not null check (char_length(name) between 1 and 200), + description text check (description is null or char_length(description) <= 4000), + status text not null default 'active' check (status in ('active', 'archived')), + thumbnail_asset_id text references media_assets(id) on delete set null, + metadata jsonb not null default '{}'::jsonb + check (jsonb_typeof(metadata) = 'object'), + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + archived_at timestamptz, + constraint ck_projects_archive_timestamp check ( + (status = 'active' and archived_at is null) + or (status = 'archived' and archived_at is not null) + ) +); +create index if not exists ix_projects_workspace on projects(workspace_id); +create index if not exists ix_projects_workspace_status on projects(workspace_id, status); +create index if not exists ix_projects_workspace_updated on projects(workspace_id, updated_at desc); +create index if not exists ix_projects_created_by on projects(created_by); + +-- The shared AuditService owns this generic event table. It is deliberately +-- not project-specific so future domains can reuse the same safe boundary. +create table if not exists audit_events ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + actor_user_id text references users(id) on delete set null, + api_key_id text references api_keys(id) on delete set null, + event_type text not null check (char_length(event_type) between 1 and 100), + entity_type text not null check (char_length(entity_type) between 1 and 64), + entity_id text not null, + request_id text, + metadata jsonb not null default '{}'::jsonb check (jsonb_typeof(metadata) = 'object'), + created_at timestamptz not null default now() +); +create index if not exists ix_audit_events_workspace_created + on audit_events(workspace_id, created_at desc); +create index if not exists ix_audit_events_type on audit_events(event_type); +create index if not exists ix_audit_events_entity on audit_events(entity_type, entity_id); + +drop trigger if exists mediarouter_tenant_touch_updated_at on projects; +create trigger mediarouter_tenant_touch_updated_at +before update on projects +for each row execute function mediarouter_tenant_touch_updated_at(); + +create or replace function mediarouter_assert_project_ownership() +returns trigger language plpgsql as $$ +declare asset_workspace text; +begin + if tg_op = 'UPDATE' then + if new.workspace_id is distinct from old.workspace_id + or new.created_by is distinct from old.created_by then + raise exception 'project ownership fields are immutable' using errcode = '23514'; + end if; + end if; + if new.thumbnail_asset_id is not null then + select workspace_id into asset_workspace from media_assets + where id = new.thumbnail_asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'project thumbnail must belong to its workspace' using errcode = '23503'; + end if; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_project_ownership on projects; +create trigger mediarouter_project_ownership +before insert or update of workspace_id, created_by, thumbnail_asset_id on projects +for each row execute function mediarouter_assert_project_ownership(); + +alter table projects enable row level security; +alter table projects force row level security; + +drop policy if exists projects_select on projects; +create policy projects_select on projects for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = projects.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists projects_insert on projects; +create policy projects_insert on projects for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and created_by = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = projects.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists projects_update on projects; +create policy projects_update on projects for update using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = projects.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +) with check ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = projects.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists projects_delete on projects; +create policy projects_delete on projects for delete using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = projects.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +alter table audit_events enable row level security; +alter table audit_events force row level security; +drop policy if exists audit_events_select on audit_events; +create policy audit_events_select on audit_events for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = audit_events.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); +drop policy if exists audit_events_insert on audit_events; +create policy audit_events_insert on audit_events for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and actor_user_id = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = audit_events.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +commit; + diff --git a/app/projects/migrations/0002_project_resources.sql b/app/projects/migrations/0002_project_resources.sql new file mode 100644 index 0000000000000000000000000000000000000000..55532c38c69ebfbfc4064e0b65e496eba330680b --- /dev/null +++ b/app/projects/migrations/0002_project_resources.sql @@ -0,0 +1,154 @@ +-- Minimum authoritative Project -> Asset / durable Generation Job links. +-- Apply after projects/0001 and security/0003_generation_domain_postgres.sql. + +begin; + +alter table media_assets add column if not exists project_id text; + +do $$ +begin + if not exists ( + select 1 from pg_constraint + where conname = 'fk_media_assets_project' + and conrelid = 'media_assets'::regclass + ) then + alter table media_assets + add constraint fk_media_assets_project + foreign key (project_id) references projects(id) on delete restrict; + end if; +end $$; + +create index if not exists ix_media_assets_workspace_project_created + on media_assets(workspace_id, project_id, created_at desc); + +create or replace function mediarouter_assert_media_asset_project_workspace() +returns trigger language plpgsql as $$ +declare project_workspace text; +declare project_status text; +begin + if tg_op = 'UPDATE' and new.workspace_id is distinct from old.workspace_id then + raise exception 'canonical asset workspace is immutable' using errcode = '23514'; + end if; + if new.project_id is not null then + select workspace_id, status into project_workspace, project_status from projects + where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'project asset must belong to its project workspace' + using errcode = '23503'; + end if; + if project_status is distinct from 'active' then + raise exception 'canonical assets cannot be attached to an archived project' + using errcode = '23514'; + end if; + end if; + return new; +end; +$$; + +drop trigger if exists mediarouter_media_asset_project_workspace on media_assets; +create trigger mediarouter_media_asset_project_workspace +before insert or update of workspace_id, project_id on media_assets +for each row execute function mediarouter_assert_media_asset_project_workspace(); + +create table if not exists project_generation_jobs ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + project_id text not null references projects(id) on delete restrict, + generation_job_id text not null references generation_jobs(id) on delete restrict, + attached_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + constraint uq_project_generation_job unique (generation_job_id) +); + +create index if not exists ix_project_generation_jobs_workspace_project_created + on project_generation_jobs(workspace_id, project_id, created_at desc); + +create or replace function mediarouter_assert_project_generation_job_workspace() +returns trigger language plpgsql as $$ +declare project_workspace text; +declare project_status text; +declare job_workspace text; +begin + if tg_op = 'UPDATE' and ( + new.workspace_id is distinct from old.workspace_id + or new.project_id is distinct from old.project_id + or new.generation_job_id is distinct from old.generation_job_id + or new.attached_by is distinct from old.attached_by + ) then + raise exception 'project generation job ownership is immutable' + using errcode = '23514'; + end if; + + select workspace_id, status into project_workspace, project_status + from projects where id = new.project_id; + select workspace_id into job_workspace + from generation_jobs where id = new.generation_job_id; + + if project_workspace is null + or job_workspace is null + or project_workspace is distinct from new.workspace_id + or job_workspace is distinct from new.workspace_id then + raise exception 'project generation job resources must share a workspace' + using errcode = '23503'; + end if; + if project_status is distinct from 'active' then + raise exception 'generation jobs cannot be attached to an archived project' + using errcode = '23514'; + end if; + if not exists ( + select 1 from workspace_memberships m + where m.workspace_id = new.workspace_id + and m.user_id = new.attached_by + and m.status = 'active' + ) then + raise exception 'generation job attachment requires active workspace membership' + using errcode = '23503'; + end if; + return new; +end; +$$; + +drop trigger if exists mediarouter_project_generation_job_workspace + on project_generation_jobs; +create trigger mediarouter_project_generation_job_workspace +before insert or update on project_generation_jobs +for each row execute function mediarouter_assert_project_generation_job_workspace(); + +alter table project_generation_jobs enable row level security; +alter table project_generation_jobs force row level security; + +drop policy if exists project_generation_jobs_select on project_generation_jobs; +create policy project_generation_jobs_select on project_generation_jobs for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = project_generation_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists project_generation_jobs_insert on project_generation_jobs; +create policy project_generation_jobs_insert on project_generation_jobs for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and attached_by = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = project_generation_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists project_generation_jobs_delete on project_generation_jobs; +create policy project_generation_jobs_delete on project_generation_jobs for delete using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = project_generation_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +commit; diff --git a/app/projects/migrations/0003_editor_persistence_rendering.sql b/app/projects/migrations/0003_editor_persistence_rendering.sql new file mode 100644 index 0000000000000000000000000000000000000000..e1ebca19baca2666f2f7779eff23ea439ecba290 --- /dev/null +++ b/app/projects/migrations/0003_editor_persistence_rendering.sql @@ -0,0 +1,235 @@ +-- Content Studio Phase 2: authoritative editor state and project render jobs. +-- Apply after projects/0001_projects_foundation.sql, projects/0002_project_resources.sql, +-- and security generation migrations. Production migration remains explicit. + +begin; + +create table if not exists project_editor_states ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + project_id text not null references projects(id) on delete restrict, + revision integer not null default 1 check (revision > 0), + schema_version integer not null check (schema_version = 1), + state jsonb not null check (jsonb_typeof(state) = 'object'), + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + updated_by text not null references users(id) on delete restrict, + constraint uq_project_editor_state_project unique (project_id) +); +create index if not exists ix_project_editor_states_workspace_project + on project_editor_states(workspace_id, project_id); + +create table if not exists project_render_jobs ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + project_id text not null references projects(id) on delete restrict, + editor_revision integer not null check (editor_revision > 0), + editor_schema_version integer not null check (editor_schema_version = 1), + editor_state jsonb not null check (jsonb_typeof(editor_state) = 'object'), + render_settings jsonb not null check (jsonb_typeof(render_settings) = 'object'), + request_fingerprint text not null check (char_length(request_fingerprint) = 64), + idempotency_key text not null check (char_length(idempotency_key) between 1 and 255), + requested_by text not null references users(id) on delete restrict, + status text not null default 'queued' + check (status in ('queued', 'processing', 'completed', 'failed', 'cancelling', 'cancelled')), + attempt_count integer not null default 0 check (attempt_count >= 0), + max_attempts integer not null default 3 check (max_attempts between 1 and 10), + next_attempt_at timestamptz, + output_asset_id text references media_assets(id) on delete restrict, + error_code text, + error_message text, + created_at timestamptz not null default now(), + started_at timestamptz, + completed_at timestamptz, + cancelled_at timestamptz, + updated_at timestamptz not null default now(), + constraint uq_project_render_idempotency unique (project_id, editor_revision, idempotency_key) +); +create index if not exists ix_project_render_jobs_workspace_status + on project_render_jobs(workspace_id, status); +create index if not exists ix_project_render_jobs_project_created + on project_render_jobs(project_id, created_at desc); +create index if not exists ix_project_render_jobs_dispatch + on project_render_jobs(status, next_attempt_at); + +drop trigger if exists mediarouter_project_editor_touch_updated_at on project_editor_states; +create trigger mediarouter_project_editor_touch_updated_at +before update on project_editor_states +for each row execute function mediarouter_tenant_touch_updated_at(); +drop trigger if exists mediarouter_project_render_touch_updated_at on project_render_jobs; +create trigger mediarouter_project_render_touch_updated_at +before update on project_render_jobs +for each row execute function mediarouter_tenant_touch_updated_at(); + +create or replace function mediarouter_assert_project_editor_ownership() +returns trigger language plpgsql as $$ +declare project_workspace text; +declare project_status text; +begin + if tg_op = 'INSERT' and new.revision <> 1 then + raise exception 'initial editor revision must be one' using errcode = '23514'; + end if; + if tg_op = 'UPDATE' and ( + new.workspace_id is distinct from old.workspace_id + or new.project_id is distinct from old.project_id + or new.created_at is distinct from old.created_at + ) then + raise exception 'project editor ownership fields are immutable' using errcode = '23514'; + end if; + if tg_op = 'UPDATE' and new.revision <> old.revision + 1 then + raise exception 'editor revision must increment exactly once' using errcode = '23514'; + end if; + if new.state->>'projectId' is distinct from new.project_id + or new.state->>'schemaVersion' is distinct from new.schema_version::text then + raise exception 'editor state identity does not match its row' using errcode = '23514'; + end if; + select workspace_id, status into project_workspace, project_status from projects where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'editor state must belong to its project workspace' using errcode = '23503'; + end if; + if project_status is distinct from 'active' then + raise exception 'archived projects cannot change editor state' using errcode = '23514'; + end if; + if not exists ( + select 1 from workspace_memberships m + where m.workspace_id = new.workspace_id and m.user_id = new.updated_by and m.status = 'active' + ) then + raise exception 'editor state update requires active workspace membership' using errcode = '23503'; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_project_editor_ownership on project_editor_states; +create trigger mediarouter_project_editor_ownership +before insert or update on project_editor_states +for each row execute function mediarouter_assert_project_editor_ownership(); + +create or replace function mediarouter_assert_project_render_ownership() +returns trigger language plpgsql as $$ +declare project_workspace text; +declare project_status text; +declare output_workspace text; +declare output_project text; +begin + if tg_op = 'UPDATE' and ( + new.workspace_id is distinct from old.workspace_id + or new.project_id is distinct from old.project_id + or new.editor_revision is distinct from old.editor_revision + or new.editor_schema_version is distinct from old.editor_schema_version + or new.editor_state is distinct from old.editor_state + or new.render_settings is distinct from old.render_settings + or new.request_fingerprint is distinct from old.request_fingerprint + or new.idempotency_key is distinct from old.idempotency_key + or new.requested_by is distinct from old.requested_by + or new.max_attempts is distinct from old.max_attempts + or new.created_at is distinct from old.created_at + ) then + raise exception 'render identity fields are immutable' using errcode = '23514'; + end if; + if new.editor_state->>'projectId' is distinct from new.project_id + or new.editor_state->>'schemaVersion' is distinct from new.editor_schema_version::text then + raise exception 'render editor snapshot identity does not match its row' using errcode = '23514'; + end if; + if tg_op = 'UPDATE' and new.status is distinct from old.status and not ( + (old.status = 'queued' and new.status in ('processing', 'cancelled', 'failed')) + or (old.status = 'processing' and new.status in ('queued', 'completed', 'failed', 'cancelling', 'cancelled')) + or (old.status = 'cancelling' and new.status in ('completed', 'failed', 'cancelled')) + ) then + raise exception 'invalid render status transition' using errcode = '23514'; + end if; + select workspace_id, status into project_workspace, project_status from projects where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'render must belong to its project workspace' using errcode = '23503'; + end if; + if tg_op = 'INSERT' and project_status is distinct from 'active' then + raise exception 'archived projects cannot be rendered' using errcode = '23514'; + end if; + if new.output_asset_id is not null then + select workspace_id, project_id into output_workspace, output_project + from media_assets where id = new.output_asset_id; + if output_workspace is null + or output_workspace is distinct from new.workspace_id + or output_project is distinct from new.project_id then + raise exception 'render output must belong to its project workspace' using errcode = '23503'; + end if; + end if; + if new.status = 'completed' and (new.output_asset_id is null or new.completed_at is null) then + raise exception 'completed render requires an output asset and timestamp' using errcode = '23514'; + end if; + if new.status = 'cancelled' and (new.cancelled_at is null or new.completed_at is null) then + raise exception 'cancelled render requires cancellation timestamps' using errcode = '23514'; + end if; + if new.status = 'failed' and new.completed_at is null then + raise exception 'failed render requires a completion timestamp' using errcode = '23514'; + end if; + if new.status in ('queued', 'processing', 'cancelling') and new.output_asset_id is not null then + raise exception 'active render cannot expose an output asset' using errcode = '23514'; + end if; + if tg_op = 'INSERT' and not exists ( + select 1 from workspace_memberships m + where m.workspace_id = new.workspace_id and m.user_id = new.requested_by and m.status = 'active' + ) then + raise exception 'render requires active workspace membership' using errcode = '23503'; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_project_render_ownership on project_render_jobs; +create trigger mediarouter_project_render_ownership +before insert or update on project_render_jobs +for each row execute function mediarouter_assert_project_render_ownership(); + +alter table project_editor_states enable row level security; +alter table project_editor_states force row level security; +drop policy if exists project_editor_states_select on project_editor_states; +create policy project_editor_states_select on project_editor_states for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_editor_states.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active') +); +drop policy if exists project_editor_states_insert on project_editor_states; +create policy project_editor_states_insert on project_editor_states for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and updated_by = current_setting('app.user_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_editor_states.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +); +drop policy if exists project_editor_states_update on project_editor_states; +create policy project_editor_states_update on project_editor_states for update using ( + workspace_id = current_setting('app.workspace_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_editor_states.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +) with check ( + workspace_id = current_setting('app.workspace_id', true) + and updated_by = current_setting('app.user_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_editor_states.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +); + +alter table project_render_jobs enable row level security; +alter table project_render_jobs force row level security; +drop policy if exists project_render_jobs_select on project_render_jobs; +create policy project_render_jobs_select on project_render_jobs for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_render_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active') +); +drop policy if exists project_render_jobs_insert on project_render_jobs; +create policy project_render_jobs_insert on project_render_jobs for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and requested_by = current_setting('app.user_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_render_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +); +drop policy if exists project_render_jobs_update on project_render_jobs; +create policy project_render_jobs_update on project_render_jobs for update using ( + workspace_id = current_setting('app.workspace_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_render_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +) with check ( + workspace_id = current_setting('app.workspace_id', true) + and exists (select 1 from workspace_memberships m where m.workspace_id = project_render_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) and m.status = 'active' and m.role <> 'viewer') +); + +commit; diff --git a/app/projects/migrations/0004_ai_studio.sql b/app/projects/migrations/0004_ai_studio.sql new file mode 100644 index 0000000000000000000000000000000000000000..c47daa463ef013e998e3cca424ba89a0c6926c0a --- /dev/null +++ b/app/projects/migrations/0004_ai_studio.sql @@ -0,0 +1,97 @@ +-- AI Studio project context for the existing generation domain. +-- Apply after projects/0003_editor_persistence_rendering.sql. + +begin; + +alter table generation_requests add column if not exists project_id text; +alter table generation_requests add column if not exists product_surface text + not null default 'generation' + check (product_surface in ('generation', 'ai_studio')); + +do $$ +begin + if not exists ( + select 1 from pg_constraint + where conname = 'fk_generation_requests_project' + and conrelid = 'generation_requests'::regclass + ) then + alter table generation_requests + add constraint fk_generation_requests_project + foreign key (project_id) references projects(id) on delete restrict; + end if; +end $$; + +create index if not exists ix_generation_requests_workspace_project_created + on generation_requests(workspace_id, project_id, created_at desc); +create index if not exists ix_generation_requests_workspace_surface_created + on generation_requests(workspace_id, product_surface, created_at desc); + +create or replace function mediarouter_assert_generation_request_project() +returns trigger language plpgsql as $$ +declare project_workspace text; +declare project_status text; +declare asset_workspace text; +declare asset_project text; +begin + if tg_op = 'UPDATE' and new.project_id is distinct from old.project_id then + raise exception 'generation project ownership is immutable' using errcode = '23514'; + end if; + if new.project_id is not null then + select workspace_id, status into project_workspace, project_status + from projects where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'generation project must belong to its request workspace' + using errcode = '23503'; + end if; + if project_status is distinct from 'active' then + raise exception 'generation project must be active' using errcode = '23514'; + end if; + if new.input_asset_id is not null then + select workspace_id, project_id into asset_workspace, asset_project + from media_assets where id = new.input_asset_id; + if asset_workspace is distinct from new.workspace_id + or asset_project is distinct from new.project_id then + raise exception 'generation input must belong to its selected project' + using errcode = '23503'; + end if; + end if; + end if; + return new; +end; +$$; + +drop trigger if exists mediarouter_generation_request_project on generation_requests; +create trigger mediarouter_generation_request_project +before insert or update of workspace_id, project_id, input_asset_id on generation_requests +for each row execute function mediarouter_assert_generation_request_project(); + +create or replace function mediarouter_assert_generation_job_workspace() +returns trigger language plpgsql as $$ +declare request_workspace text; +declare request_project text; +declare asset_workspace text; +declare asset_project text; +begin + select workspace_id, project_id into request_workspace, request_project + from generation_requests where id = new.generation_request_id; + if request_workspace is null or request_workspace is distinct from new.workspace_id then + raise exception 'generation job must belong to its request workspace' + using errcode = '23503'; + end if; + if new.output_asset_id is not null then + select workspace_id, project_id into asset_workspace, asset_project + from media_assets where id = new.output_asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'generation output asset must belong to its job workspace' + using errcode = '23503'; + end if; + if asset_project is distinct from request_project then + raise exception 'generation output asset must belong to its request project' + using errcode = '23503'; + end if; + end if; + return new; +end; +$$; + +commit; diff --git a/app/projects/migrations/0005_ai_copilot.sql b/app/projects/migrations/0005_ai_copilot.sql new file mode 100644 index 0000000000000000000000000000000000000000..c8f0fb5f9f7fda41e3f947d78da4ffb576585e2e --- /dev/null +++ b/app/projects/migrations/0005_ai_copilot.sql @@ -0,0 +1,132 @@ +-- AI Copilot: bounded durable plans and action results. +-- Apply after projects/0004_ai_studio.sql. Production migration remains explicit. + +begin; + +create table if not exists copilot_runs ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + user_id text not null references users(id) on delete restrict, + project_id text references projects(id) on delete restrict, + idempotency_key text not null check (char_length(idempotency_key) between 1 and 255), + request_fingerprint text not null check (char_length(request_fingerprint) = 64), + request_text text not null check (char_length(request_text) between 1 and 4000), + context jsonb not null check (jsonb_typeof(context) = 'object'), + plan jsonb not null check (jsonb_typeof(plan) = 'object'), + status text not null check ( + status in ('plan_ready','blocked','executing','completed','partial','failed','cancelled') + ), + current_action_id text, + results jsonb not null default '[]'::jsonb check (jsonb_typeof(results) = 'array'), + summary text, + error_code text, + error_message text, + confirmed_at timestamptz, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + completed_at timestamptz, + constraint uq_copilot_run_workspace_idempotency unique (workspace_id, idempotency_key) +); + +create index if not exists ix_copilot_runs_workspace_created + on copilot_runs(workspace_id, created_at desc); +create index if not exists ix_copilot_runs_workspace_status + on copilot_runs(workspace_id, status); +create index if not exists ix_copilot_runs_project_created + on copilot_runs(project_id, created_at desc); + +drop trigger if exists mediarouter_copilot_touch_updated_at on copilot_runs; +create trigger mediarouter_copilot_touch_updated_at +before update on copilot_runs +for each row execute function mediarouter_tenant_touch_updated_at(); + +create or replace function mediarouter_assert_copilot_run_ownership() +returns trigger language plpgsql as $$ +declare project_workspace text; +begin + if tg_op = 'UPDATE' and ( + new.workspace_id is distinct from old.workspace_id + or new.user_id is distinct from old.user_id + or new.project_id is distinct from old.project_id + or new.idempotency_key is distinct from old.idempotency_key + or new.request_fingerprint is distinct from old.request_fingerprint + or new.request_text is distinct from old.request_text + or new.context is distinct from old.context + or new.plan is distinct from old.plan + or new.created_at is distinct from old.created_at + ) then + raise exception 'copilot run identity fields are immutable' using errcode = '23514'; + end if; + if new.context->>'workspace_id' is distinct from new.workspace_id then + raise exception 'copilot context workspace does not match its run' using errcode = '23514'; + end if; + if new.project_id is not null then + select workspace_id into project_workspace from projects where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'copilot project must belong to its workspace' using errcode = '23503'; + end if; + end if; + if not exists ( + select 1 from workspace_memberships m + where m.workspace_id = new.workspace_id + and m.user_id = new.user_id + and m.status = 'active' + ) then + raise exception 'copilot run requires active workspace membership' using errcode = '23503'; + end if; + if new.status in ('completed','partial','failed','cancelled') and new.completed_at is null then + raise exception 'terminal copilot runs require completed_at' using errcode = '23514'; + end if; + return new; +end; +$$; + +drop trigger if exists mediarouter_copilot_run_ownership on copilot_runs; +create trigger mediarouter_copilot_run_ownership +before insert or update on copilot_runs +for each row execute function mediarouter_assert_copilot_run_ownership(); + +alter table copilot_runs enable row level security; +alter table copilot_runs force row level security; + +drop policy if exists copilot_runs_select on copilot_runs; +create policy copilot_runs_select on copilot_runs for select using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = copilot_runs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) +); + +drop policy if exists copilot_runs_insert on copilot_runs; +create policy copilot_runs_insert on copilot_runs for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and user_id = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = copilot_runs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + and m.role <> 'viewer' + ) +); + +drop policy if exists copilot_runs_update on copilot_runs; +create policy copilot_runs_update on copilot_runs for update using ( + workspace_id = current_setting('app.workspace_id', true) + and user_id = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = copilot_runs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + and m.role <> 'viewer' + ) +) with check ( + workspace_id = current_setting('app.workspace_id', true) + and user_id = current_setting('app.user_id', true) +); + +commit; diff --git a/app/projects/migrations/0006_template_marketplace.sql b/app/projects/migrations/0006_template_marketplace.sql new file mode 100644 index 0000000000000000000000000000000000000000..c6c7428bc19a689ff0d74556c15d63c0d57398a5 --- /dev/null +++ b/app/projects/migrations/0006_template_marketplace.sql @@ -0,0 +1,152 @@ +-- Template Marketplace: versioned declarative templates and atomic applications. +-- Apply after projects/0005_ai_copilot.sql. Production migration remains explicit. + +begin; + +create table if not exists marketplace_templates ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + slug text not null check (char_length(slug) between 1 and 120), + name text not null check (char_length(name) between 1 and 200), + description text not null check (char_length(description) between 1 and 2000), + status text not null check (status in ('draft','published','archived')), + visibility text not null check (visibility in ('private','workspace','public')), + category text not null, + tags jsonb not null default '[]'::jsonb check (jsonb_typeof(tags) = 'array'), + thumbnail_asset_id text references media_assets(id) on delete set null, + preview_asset_id text references media_assets(id) on delete set null, + duration_ms integer not null check (duration_ms > 0), + aspect_ratio text not null check (aspect_ratio in ('9:16','16:9','1:1','4:5')), + metadata jsonb not null default '{}'::jsonb check (jsonb_typeof(metadata) = 'object'), + created_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint uq_marketplace_template_workspace_slug unique (workspace_id, slug) +); + +create index if not exists ix_marketplace_templates_workspace_status + on marketplace_templates(workspace_id, status); +create index if not exists ix_marketplace_templates_discovery + on marketplace_templates(visibility, status, category); + +create table if not exists marketplace_template_versions ( + id text primary key default gen_random_uuid()::text, + template_id text not null references marketplace_templates(id) on delete restrict, + version integer not null check (version > 0), + schema_version integer not null check (schema_version > 0), + definition jsonb not null check (jsonb_typeof(definition) = 'object' and pg_column_size(definition) <= 1048576), + requirements jsonb not null default '{}'::jsonb check (jsonb_typeof(requirements) = 'object'), + status text not null check (status in ('draft','published','archived')), + created_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + constraint uq_marketplace_template_version unique (template_id, version) +); + +create index if not exists ix_marketplace_template_versions_template + on marketplace_template_versions(template_id, version desc); + +create table if not exists marketplace_template_applications ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + template_id text not null references marketplace_templates(id) on delete restrict, + template_version_id text not null references marketplace_template_versions(id) on delete restrict, + template_version integer not null check (template_version > 0), + project_id text not null references projects(id) on delete restrict, + editor_revision integer not null check (editor_revision > 0), + slot_bindings jsonb not null check (jsonb_typeof(slot_bindings) = 'object'), + idempotency_key text not null check (char_length(idempotency_key) between 1 and 255), + request_fingerprint text not null check (char_length(request_fingerprint) = 64), + created_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + constraint uq_marketplace_template_application_idempotency + unique (workspace_id, idempotency_key) +); + +create index if not exists ix_marketplace_template_applications_project + on marketplace_template_applications(project_id, created_at desc); +create index if not exists ix_marketplace_template_applications_template + on marketplace_template_applications(template_id, template_version_id); + +create or replace function mediarouter_template_version_immutable() +returns trigger language plpgsql as $$ +begin + if old.status = 'published' then + raise exception 'published template versions are immutable' using errcode = '23514'; + end if; + if new.template_id is distinct from old.template_id + or new.version is distinct from old.version + or new.schema_version is distinct from old.schema_version + or new.definition is distinct from old.definition + or new.requirements is distinct from old.requirements + or new.created_by is distinct from old.created_by + or new.created_at is distinct from old.created_at then + raise exception 'template version identity and definition are immutable' using errcode = '23514'; + end if; + return new; +end; +$$; + +drop trigger if exists mediarouter_template_version_immutable + on marketplace_template_versions; +create trigger mediarouter_template_version_immutable +before update on marketplace_template_versions +for each row execute function mediarouter_template_version_immutable(); + +alter table marketplace_templates enable row level security; +alter table marketplace_templates force row level security; +alter table marketplace_template_versions enable row level security; +alter table marketplace_template_versions force row level security; +alter table marketplace_template_applications enable row level security; +alter table marketplace_template_applications force row level security; + +create policy marketplace_templates_select on marketplace_templates for select using ( + workspace_id = current_setting('app.workspace_id', true) + or (visibility = 'public' and status = 'published') +); +create policy marketplace_templates_write on marketplace_templates for all using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = marketplace_templates.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' and m.role <> 'viewer' + ) +) with check (workspace_id = current_setting('app.workspace_id', true)); + +create policy marketplace_template_versions_select +on marketplace_template_versions for select using ( + exists ( + select 1 from marketplace_templates t + where t.id = marketplace_template_versions.template_id + and (t.workspace_id = current_setting('app.workspace_id', true) + or (t.visibility = 'public' and t.status = 'published' + and marketplace_template_versions.status = 'published')) + ) +); +create policy marketplace_template_versions_write +on marketplace_template_versions for all using ( + exists ( + select 1 from marketplace_templates t + where t.id = marketplace_template_versions.template_id + and t.workspace_id = current_setting('app.workspace_id', true) + ) +) with check ( + exists ( + select 1 from marketplace_templates t + where t.id = marketplace_template_versions.template_id + and t.workspace_id = current_setting('app.workspace_id', true) + ) +); + +create policy marketplace_template_applications_select +on marketplace_template_applications for select using ( + workspace_id = current_setting('app.workspace_id', true) + and created_by = current_setting('app.user_id', true) +); +create policy marketplace_template_applications_insert +on marketplace_template_applications for insert with check ( + workspace_id = current_setting('app.workspace_id', true) + and created_by = current_setting('app.user_id', true) +); + +commit; diff --git a/app/projects/migrations/0008_brand_kits.sql b/app/projects/migrations/0008_brand_kits.sql new file mode 100644 index 0000000000000000000000000000000000000000..6bd017e0c37765b13e9315c1486bfef8ea106049 --- /dev/null +++ b/app/projects/migrations/0008_brand_kits.sql @@ -0,0 +1,177 @@ +begin; + +create table if not exists brand_kits ( + id text primary key, + workspace_id text not null references workspaces(id) on delete restrict, + name varchar(200) not null check (char_length(trim(name)) between 1 and 200), + description text check (description is null or char_length(description) <= 4000), + status varchar(16) not null default 'active' + check (status in ('active','archived')), + is_default boolean not null default false, + active_version_id text, + created_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint uq_brand_kits_id_workspace unique (id, workspace_id) +); +create index if not exists ix_brand_kits_workspace_updated + on brand_kits(workspace_id, updated_at desc); +create index if not exists ix_brand_kits_workspace_default + on brand_kits(workspace_id, is_default); +create unique index if not exists uq_brand_kits_one_default + on brand_kits(workspace_id) where is_default and status = 'active'; + +create table if not exists brand_kit_versions ( + id text primary key, + workspace_id text not null references workspaces(id) on delete restrict, + brand_kit_id text not null references brand_kits(id) on delete cascade, + version integer not null check (version > 0), + status varchar(16) not null default 'draft' + check (status in ('draft','published','archived')), + colors jsonb not null default '[]'::jsonb, + typography jsonb not null default '{}'::jsonb, + voice jsonb not null default '{}'::jsonb, + ctas jsonb not null default '{}'::jsonb, + hashtags jsonb not null default '{}'::jsonb, + watermark jsonb not null default '{}'::jsonb, + ai_guidance text check (ai_guidance is null or char_length(ai_guidance) <= 8000), + metadata jsonb not null default '{}'::jsonb, + created_by text not null references users(id) on delete restrict, + created_at timestamptz not null default now(), + published_at timestamptz, + constraint uq_brand_kit_version unique (brand_kit_id, version), + constraint uq_brand_versions_id_kit unique (id, brand_kit_id), + constraint fk_brand_version_workspace_kit + foreign key (brand_kit_id, workspace_id) + references brand_kits(id, workspace_id) on delete cascade +); +create index if not exists ix_brand_versions_kit_created + on brand_kit_versions(brand_kit_id, created_at desc); + +do $$ +begin + if not exists ( + select 1 from pg_constraint where conname = 'fk_brand_kits_active_version' + ) then + alter table brand_kits add constraint fk_brand_kits_active_version + foreign key (active_version_id, id) + references brand_kit_versions(id, brand_kit_id) + on delete restrict; + end if; +end $$; + +create table if not exists brand_kit_assets ( + id text primary key, + workspace_id text not null references workspaces(id) on delete restrict, + brand_kit_version_id text not null references brand_kit_versions(id) on delete cascade, + asset_id text not null references media_assets(id) on delete restrict, + role varchar(32) not null check ( + role in ( + 'logo_primary','logo_secondary','watermark','favicon', + 'social_avatar','social_cover' + ) + ), + created_at timestamptz not null default now(), + constraint uq_brand_asset_role unique (brand_kit_version_id, role) +); +create index if not exists ix_brand_assets_workspace_asset + on brand_kit_assets(workspace_id, asset_id); + +create table if not exists brand_kit_platform_settings ( + id text primary key, + workspace_id text not null references workspaces(id) on delete restrict, + brand_kit_version_id text not null references brand_kit_versions(id) on delete cascade, + provider varchar(32) not null check ( + provider in ( + 'facebook','instagram','tiktok','x','youtube','linkedin','telegram','whatsapp' + ) + ), + settings jsonb not null default '{}'::jsonb, + constraint uq_brand_platform_setting unique (brand_kit_version_id, provider) +); +create index if not exists ix_brand_platform_settings_workspace + on brand_kit_platform_settings(workspace_id); + +create table if not exists brand_governance_settings ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + require_brand_kit_for_publish boolean not null default false, + require_published_brand_version boolean not null default false, + allow_user_override_brand_defaults boolean not null default true, + updated_by text not null references users(id) on delete restrict, + updated_at timestamptz not null default now(), + constraint uq_brand_governance_workspace unique (workspace_id) +); + +alter table projects add column if not exists brand_kit_id text; +do $$ +begin + if not exists ( + select 1 from pg_constraint where conname = 'fk_projects_brand_kit' + ) then + alter table projects add constraint fk_projects_brand_kit + foreign key (brand_kit_id) references brand_kits(id) on delete restrict; + end if; +end $$; +create index if not exists ix_projects_workspace_brand_kit + on projects(workspace_id, brand_kit_id); + +-- Publishing and analytics attribution is explicit. Null means the content +-- did not claim to use a brand version. +alter table social_posts add column if not exists brand_kit_id text; +alter table social_posts add column if not exists brand_kit_version_id text; +alter table analytics_post_metrics add column if not exists brand_kit_id text; +alter table analytics_post_metrics add column if not exists brand_kit_version_id text; +create index if not exists ix_analytics_post_metrics_brand + on analytics_post_metrics(workspace_id, brand_kit_id, brand_kit_version_id); + +create or replace function mediarouter_validate_brand_reference() +returns trigger language plpgsql as $$ +declare reference_workspace text; +begin + if tg_table_name = 'brand_kit_assets' then + select workspace_id into reference_workspace from media_assets where id = new.asset_id; + elsif tg_table_name = 'projects' then + if new.brand_kit_id is null then return new; end if; + select workspace_id into reference_workspace from brand_kits where id = new.brand_kit_id; + else + return new; + end if; + if reference_workspace is null or reference_workspace <> new.workspace_id then + raise exception 'cross-workspace brand reference rejected' + using errcode = '23514'; + end if; + return new; +end $$; + +drop trigger if exists trg_brand_asset_workspace on brand_kit_assets; +create trigger trg_brand_asset_workspace +before insert or update on brand_kit_assets +for each row execute function mediarouter_validate_brand_reference(); + +drop trigger if exists trg_project_brand_workspace on projects; +create trigger trg_project_brand_workspace +before insert or update of brand_kit_id on projects +for each row execute function mediarouter_validate_brand_reference(); + +do $$ +declare table_name text; +begin + foreach table_name in array array[ + 'brand_kits', + 'brand_kit_versions', + 'brand_kit_assets', + 'brand_kit_platform_settings', + 'brand_governance_settings' + ] loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + execute format('drop policy if exists brand_workspace_isolation on %I', table_name); + execute format( + 'create policy brand_workspace_isolation on %I using (workspace_id = current_setting(''app.workspace_id'', true)) with check (workspace_id = current_setting(''app.workspace_id'', true))', + table_name + ); + end loop; +end $$; + +commit; diff --git a/app/projects/migrations/0009_brand_kits_explicit_columns.sql b/app/projects/migrations/0009_brand_kits_explicit_columns.sql new file mode 100644 index 0000000000000000000000000000000000000000..acb06b5ec670261fd885f6e87f17ba470b4e7753 --- /dev/null +++ b/app/projects/migrations/0009_brand_kits_explicit_columns.sql @@ -0,0 +1,30 @@ +begin; + +-- Rename version column to version_number for consistency +alter table brand_kit_versions rename column version to version_number; + +-- Add new explicit columns for branding fields +alter table brand_kit_versions + add column if not exists logo_asset_id text references media_assets(id) on delete set null, + add column if not exists favicon_asset_id text references media_assets(id) on delete set null, + add column if not exists primary_color varchar(7), + add column if not exists secondary_color varchar(7), + add column if not exists accent_color varchar(7), + add column if not exists background_color varchar(7), + add column if not exists text_color varchar(7), + add column if not exists font_family_primary varchar(100), + add column if not exists font_family_secondary varchar(100), + add column if not exists heading_font varchar(100), + add column if not exists body_font varchar(100), + add column if not exists voice_style varchar(50), + add column if not exists tone varchar(50), + add column if not exists default_cta varchar(100), + add column if not exists watermark_asset_id text references media_assets(id) on delete set null, + add column if not exists watermark_position varchar(20), + add column if not exists watermark_opacity numeric(3,2); + +-- Note: We are keeping the old JSONB columns (colors, typography, voice, ctas, hashtags, watermark) +-- for backward compatibility until data migration is complete and the service layer is fully updated. +-- They will be marked as deprecated. + +commit; diff --git a/app/projects/migrations/0010_collaboration_foundation.sql b/app/projects/migrations/0010_collaboration_foundation.sql new file mode 100644 index 0000000000000000000000000000000000000000..9db67893f68b51b24234910d412d23504470ad29 --- /dev/null +++ b/app/projects/migrations/0010_collaboration_foundation.sql @@ -0,0 +1,114 @@ +begin; + +-- Teams +create table if not exists teams ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + name varchar(255) not null, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now() +); + +-- Team Members +create table if not exists team_members ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + team_id text not null references teams(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + role varchar(50) not null, + created_at timestamptz not null default now(), + unique(team_id, user_id) +); + +-- Invitations +create table if not exists invitations ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + email text not null, + role varchar(50) not null, + status varchar(20) not null check (status in ('pending', 'accepted', 'expired', 'revoked')), + expires_at timestamptz not null, + created_at timestamptz not null default now() +); + +-- Project Collaborators +create table if not exists project_collaborators ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + project_id text not null references projects(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + role varchar(50) not null, + created_at timestamptz not null default now(), + unique(project_id, user_id) +); + +-- Approval Workflows +create table if not exists approval_workflows ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + project_id text not null references projects(id) on delete cascade, + name varchar(255) not null, + created_at timestamptz not null default now() +); + +-- Approval Requests +create table if not exists approval_requests ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + workflow_id text not null references approval_workflows(id) on delete cascade, + project_id text not null references projects(id) on delete cascade, + status varchar(20) not null check (status in ('pending', 'approved', 'rejected')), + created_by text not null references users(id) on delete cascade, + created_at timestamptz not null default now() +); + +-- Review Comments +create table if not exists review_comments ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + request_id text not null references approval_requests(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + content text not null, + created_at timestamptz not null default now() +); + +-- Collaboration Activity +create table if not exists collaboration_activity ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + action varchar(100) not null, + entity_id text not null, + entity_type varchar(50) not null, + metadata jsonb not null default '{}', + created_at timestamptz not null default now() +); + +-- Indexes +create index ix_teams_workspace_id on teams(workspace_id); +create index ix_invitations_workspace_id on invitations(workspace_id); +create index ix_team_members_team_id on team_members(team_id); +create index ix_project_collaborators_project_id on project_collaborators(project_id); +create index ix_approval_workflows_project_id on approval_workflows(project_id); +create index ix_approval_requests_workflow_id on approval_requests(workflow_id); +create index ix_review_comments_request_id on review_comments(request_id); +create index ix_collaboration_activity_workspace_id on collaboration_activity(workspace_id); + +-- RLS +do $$ +declare + table_name text; +begin + for table_name in select t.table_name + from information_schema.tables t + where t.table_schema = 'public' + and t.table_name in ('teams', 'team_members', 'invitations', 'project_collaborators', 'approval_workflows', 'approval_requests', 'review_comments', 'collaboration_activity') + loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + execute format('drop policy if exists collaboration_workspace_isolation on %I', table_name); + execute format('create policy collaboration_workspace_isolation on %I using (workspace_id = current_setting(''app.workspace_id'', true)) with check (workspace_id = current_setting(''app.workspace_id'', true))', table_name); + end loop; +end $$; + +commit; diff --git a/app/projects/migrations/0011_collaboration_enhancements.sql b/app/projects/migrations/0011_collaboration_enhancements.sql new file mode 100644 index 0000000000000000000000000000000000000000..ad8811d0a95d6d7293601748b586d932b6fdfca3 --- /dev/null +++ b/app/projects/migrations/0011_collaboration_enhancements.sql @@ -0,0 +1,18 @@ +begin; + +-- Invitation Security: Add token hash +alter table invitations add column if not exists token_hash text; + +-- Collaboration Comments: Add timeline/clip references and resolution state +alter table review_comments + add column if not exists timeline_timestamp float, + add column if not exists clip_id text, + add column if not exists parent_id text references review_comments(id) on delete cascade, + add column if not exists resolution_state varchar(20) not null default 'open' check (resolution_state in ('open', 'resolved', 'ignored')); + +-- Approval Workflows: Update status constraints +alter table approval_requests + drop constraint if exists approval_requests_status_check, + add constraint approval_requests_status_check check (status in ('draft', 'submitted', 'approved', 'rejected', 'revision_required')); + +commit; diff --git a/app/projects/migrations/0012_approval_revision_awareness.sql b/app/projects/migrations/0012_approval_revision_awareness.sql new file mode 100644 index 0000000000000000000000000000000000000000..ba7c607a6a19868f41a70c88ffe52bbdfa7475fd --- /dev/null +++ b/app/projects/migrations/0012_approval_revision_awareness.sql @@ -0,0 +1,6 @@ +begin; + +-- Add editor_revision to approval_requests to bind request to specific project state +alter table approval_requests add column if not exists editor_revision integer not null default 1; + +commit; diff --git a/app/projects/migrations/0013_notification_preferences.sql b/app/projects/migrations/0013_notification_preferences.sql new file mode 100644 index 0000000000000000000000000000000000000000..1f66e1c840627c0b910789a7b76d6cf6c0ffa11d --- /dev/null +++ b/app/projects/migrations/0013_notification_preferences.sql @@ -0,0 +1,25 @@ +begin; + +-- Notification Preferences +create table if not exists notification_preferences ( + id text primary key, + workspace_id text not null references workspaces(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + event_type varchar(50) not null check (event_type in ( + 'collaboration_event', 'approval_request', 'publishing_failure', 'mention', 'project_activity' + )), + enabled boolean not null default true, + created_at timestamptz not null default now(), + unique(workspace_id, user_id, event_type) +); + +create index ix_notification_preferences_workspace_user on notification_preferences(workspace_id, user_id); + +-- RLS +alter table notification_preferences enable row level security; +alter table notification_preferences force row level security; +create policy notification_preferences_workspace_isolation on notification_preferences + using (workspace_id = current_setting('app.workspace_id', true)) + with check (workspace_id = current_setting('app.workspace_id', true)); + +commit; diff --git a/app/projects/models.py b/app/projects/models.py new file mode 100644 index 0000000000000000000000000000000000000000..c6bd8256c61ab4727a1d23f687dd13a54adc76fc --- /dev/null +++ b/app/projects/models.py @@ -0,0 +1,183 @@ +from __future__ import annotations + +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import ( + JSON, + CheckConstraint, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base, utcnow + + +class Project(Base): + """A durable workspace-owned container for future creative resources.""" + + __tablename__ = "projects" + __table_args__ = ( + CheckConstraint("status in ('active', 'archived')", name="ck_projects_status"), + Index("ix_projects_workspace", "workspace_id"), + Index("ix_projects_workspace_status", "workspace_id", "status"), + Index("ix_projects_workspace_updated", "workspace_id", "updated_at"), + Index("ix_projects_created_by", "created_by"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + name: Mapped[str] = mapped_column(String(200), nullable=False) + description: Mapped[str | None] = mapped_column(Text) + status: Mapped[str] = mapped_column(String(16), nullable=False, default="active") + thumbnail_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="SET NULL") + ) + brand_kit_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("brand_kits.id", ondelete="RESTRICT") + ) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + archived_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class ProjectGenerationJob(Base): + """Associates the existing durable generation job with one project.""" + + __tablename__ = "project_generation_jobs" + __table_args__ = ( + UniqueConstraint("generation_job_id", name="uq_project_generation_job"), + Index( + "ix_project_generation_jobs_workspace_project_created", + "workspace_id", + "project_id", + "created_at", + ), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + project_id: Mapped[str] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT"), nullable=False + ) + generation_job_id: Mapped[str] = mapped_column( + String(36), ForeignKey("generation_jobs.id", ondelete="RESTRICT"), nullable=False + ) + attached_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + +class ProjectEditorState(Base): + """The one authoritative current editor document for a project.""" + + __tablename__ = "project_editor_states" + __table_args__ = ( + UniqueConstraint("project_id", name="uq_project_editor_state_project"), + Index("ix_project_editor_states_workspace_project", "workspace_id", "project_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + project_id: Mapped[str] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT"), nullable=False + ) + revision: Mapped[int] = mapped_column(Integer, nullable=False, default=1) + schema_version: Mapped[int] = mapped_column(Integer, nullable=False) + state_json: Mapped[dict[str, object]] = mapped_column( + "state", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + updated_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + + +class ProjectRenderJob(Base): + """A durable render of one immutable editor revision snapshot.""" + + __tablename__ = "project_render_jobs" + __table_args__ = ( + UniqueConstraint( + "project_id", + "editor_revision", + "idempotency_key", + name="uq_project_render_idempotency", + ), + CheckConstraint( + "status in ('queued', 'processing', 'completed', 'failed', 'cancelling', 'cancelled')", + name="ck_project_render_status", + ), + Index("ix_project_render_jobs_workspace_status", "workspace_id", "status"), + Index("ix_project_render_jobs_project_created", "project_id", "created_at"), + Index("ix_project_render_jobs_dispatch", "status", "next_attempt_at"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + project_id: Mapped[str] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT"), nullable=False + ) + editor_revision: Mapped[int] = mapped_column(Integer, nullable=False) + editor_schema_version: Mapped[int] = mapped_column(Integer, nullable=False) + editor_state_json: Mapped[dict[str, object]] = mapped_column( + "editor_state", JSON, nullable=False + ) + render_settings_json: Mapped[dict[str, object]] = mapped_column( + "render_settings", JSON, nullable=False + ) + request_fingerprint: Mapped[str] = mapped_column(String(64), nullable=False) + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + requested_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + max_attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=3) + next_attempt_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + output_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="RESTRICT") + ) + error_code: Mapped[str | None] = mapped_column(String(100)) + error_message: Mapped[str | None] = mapped_column(Text) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + cancelled_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) diff --git a/app/projects/models/collaboration.py b/app/projects/models/collaboration.py new file mode 100644 index 0000000000000000000000000000000000000000..aba71772fee2425682ef5a337229e0819a991467 --- /dev/null +++ b/app/projects/models/collaboration.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import ( + DateTime, + ForeignKey, + Index, + Integer, + String, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base, utcnow + +class Team(Base): + __tablename__ = "teams" + __table_args__ = ( + Index("ix_teams_workspace_id", "workspace_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + name: Mapped[str] = mapped_column(String(255), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + +class TeamMember(Base): + __tablename__ = "team_members" + __table_args__ = ( + UniqueConstraint("team_id", "user_id", name="uq_team_member"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + team_id: Mapped[str] = mapped_column(String(36), ForeignKey("teams.id", ondelete="CASCADE"), nullable=False) + user_id: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False) + role: Mapped[str] = mapped_column(String(50), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) +class Invitation(Base): + __tablename__ = "invitations" + __table_args__ = ( + Index("ix_invitations_workspace_id", "workspace_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + email: Mapped[str] = mapped_column(String(255), nullable=False) + role: Mapped[str] = mapped_column(String(50), nullable=False) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="pending") + token_hash: Mapped[str | None] = mapped_column(String(255), nullable=True) + expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) +class ApprovalWorkflow(Base): + __tablename__ = "approval_workflows" + __table_args__ = ( + Index("ix_approval_workflows_project_id", "project_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + project_id: Mapped[str] = mapped_column(String(36), ForeignKey("projects.id", ondelete="CASCADE"), nullable=False) + name: Mapped[str] = mapped_column(String(255), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + +class ApprovalRequest(Base): + __tablename__ = "approval_requests" + __table_args__ = ( + Index("ix_approval_requests_workflow_id", "workflow_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + workflow_id: Mapped[str] = mapped_column(String(36), ForeignKey("approval_workflows.id", ondelete="CASCADE"), nullable=False) + project_id: Mapped[str] = mapped_column(String(36), ForeignKey("projects.id", ondelete="CASCADE"), nullable=False) + editor_revision: Mapped[int] = mapped_column(Integer, nullable=False, default=1) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="pending") + created_by: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + +class ReviewComment(Base): + __tablename__ = "review_comments" + __table_args__ = ( + Index("ix_review_comments_request_id", "request_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + request_id: Mapped[str] = mapped_column(String(36), ForeignKey("approval_requests.id", ondelete="CASCADE"), nullable=False) + user_id: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False) + content: Mapped[str] = mapped_column(String(2000), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) diff --git a/app/projects/models/notifications.py b/app/projects/models/notifications.py new file mode 100644 index 0000000000000000000000000000000000000000..e76c9d49db0d429ef206faf6c95e3a851e618d3e --- /dev/null +++ b/app/projects/models/notifications.py @@ -0,0 +1,20 @@ +from __future__ import annotations +from datetime import datetime +from uuid import uuid4 +from sqlalchemy import String, ForeignKey, Boolean, DateTime, Index, UniqueConstraint +from sqlalchemy.orm import Mapped, mapped_column +from app.security.models import Base, utcnow + +class NotificationPreference(Base): + __tablename__ = "notification_preferences" + __table_args__ = ( + UniqueConstraint("workspace_id", "user_id", "event_type", name="uq_notification_preference"), + Index("ix_notification_preferences_workspace_user", "workspace_id", "user_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column(String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False) + user_id: Mapped[str] = mapped_column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False) + event_type: Mapped[str] = mapped_column(String(50), nullable=False) + enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) diff --git a/app/projects/render_state.py b/app/projects/render_state.py new file mode 100644 index 0000000000000000000000000000000000000000..5a95c30519922a822d392de5977b8c0b73be8de3 --- /dev/null +++ b/app/projects/render_state.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +from enum import Enum + +from app.projects.errors import ProjectRenderTransitionError + + +class RenderStatus(str, Enum): + QUEUED = "queued" + PROCESSING = "processing" + COMPLETED = "completed" + FAILED = "failed" + CANCELLING = "cancelling" + CANCELLED = "cancelled" + + +TERMINAL_RENDER_STATUSES = frozenset( + {RenderStatus.COMPLETED, RenderStatus.FAILED, RenderStatus.CANCELLED} +) + +ALLOWED_RENDER_TRANSITIONS: dict[RenderStatus, frozenset[RenderStatus]] = { + RenderStatus.QUEUED: frozenset( + {RenderStatus.PROCESSING, RenderStatus.CANCELLED, RenderStatus.FAILED} + ), + RenderStatus.PROCESSING: frozenset( + { + RenderStatus.QUEUED, + RenderStatus.COMPLETED, + RenderStatus.FAILED, + RenderStatus.CANCELLING, + RenderStatus.CANCELLED, + } + ), + RenderStatus.CANCELLING: frozenset( + {RenderStatus.CANCELLED, RenderStatus.COMPLETED, RenderStatus.FAILED} + ), + RenderStatus.COMPLETED: frozenset(), + RenderStatus.FAILED: frozenset(), + RenderStatus.CANCELLED: frozenset(), +} + + +def validate_render_transition(current: str, target: str) -> RenderStatus: + source = RenderStatus(current) + destination = RenderStatus(target) + if destination not in ALLOWED_RENDER_TRANSITIONS[source]: + raise ProjectRenderTransitionError( + f"Cannot transition render from {source.value} to {destination.value}." + ) + return destination diff --git a/app/projects/repositories/__init__.py b/app/projects/repositories/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..246a2571dccc4f240e2ee926e52abb0c6bf45a63 --- /dev/null +++ b/app/projects/repositories/__init__.py @@ -0,0 +1 @@ +"""Project persistence repositories.""" diff --git a/app/projects/repositories/approval_repository.py b/app/projects/repositories/approval_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..57a087072e442cabaf6c184bc3f76111b8554ac0 --- /dev/null +++ b/app/projects/repositories/approval_repository.py @@ -0,0 +1,69 @@ +from __future__ import annotations +from app.security.database import SecurityDatabase +from sqlalchemy import select +from app.projects.models.collaboration import ApprovalWorkflow, ApprovalRequest, ReviewComment + +class ApprovalRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def create_workflow(self, workspace_id: str, project_id: str, name: str) -> ApprovalWorkflow: + async with self.database.session() as session: + workflow = ApprovalWorkflow(workspace_id=workspace_id, project_id=project_id, name=name) + session.add(workflow) + await session.commit() + await session.refresh(workflow) + return workflow + + async def create_request(self, workspace_id: str, workflow_id: str, project_id: str, user_id: str, editor_revision: int) -> ApprovalRequest: + async with self.database.session() as session: + request = ApprovalRequest( + workspace_id=workspace_id, + workflow_id=workflow_id, + project_id=project_id, + editor_revision=editor_revision, + created_by=user_id, + status="pending" + ) + session.add(request) + await session.commit() + await session.refresh(request) + return request + + async def update_request_status(self, request_id: str, status: str) -> ApprovalRequest: + async with self.database.session() as session: + request = await session.get(ApprovalRequest, request_id) + if not request: + raise Exception("Request not found") + request.status = status + await session.commit() + await session.refresh(request) + return request + + async def get_request(self, request_id: str) -> ApprovalRequest: + async with self.database.session() as session: + request = await session.get(ApprovalRequest, request_id) + if not request: + raise Exception("Request not found") + return request + + async def add_review_comment(self, request_id: str, user_id: str, workspace_id: str, content: str) -> ReviewComment: + async with self.database.session() as session: + comment = ReviewComment( + request_id=request_id, + user_id=user_id, + workspace_id=workspace_id, + content=content + ) + session.add(comment) + await session.commit() + await session.refresh(comment) + return comment + + async def list_requests(self, workflow_id: str) -> list[ApprovalRequest]: + async with self.database.session() as session: + result = await session.scalars( + select(ApprovalRequest).where(ApprovalRequest.workflow_id == workflow_id) + ) + return list(result.all()) + diff --git a/app/projects/repositories/collaboration_repository.py b/app/projects/repositories/collaboration_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..29fb2ca9f96b8b6cb836007134e490e1ad2983d4 --- /dev/null +++ b/app/projects/repositories/collaboration_repository.py @@ -0,0 +1,181 @@ +from __future__ import annotations +from typing import Any + +from app.security.database import SecurityDatabase +from sqlalchemy import select, func +from app.projects.models.collaboration import Team, Invitation, ProjectCollaborator, CollaborationActivity +from app.security.models import WorkspaceMembership + +class CollaborationRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def list_teams(self, workspace_id: str) -> list[Team]: + async with self.database.session() as session: + result = await session.scalars(select(Team).where(Team.workspace_id == workspace_id)) + return list(result.all()) + + async def create_team(self, workspace_id: str, name: str) -> Team: + async with self.database.session() as session: + team = Team(workspace_id=workspace_id, name=name) + session.add(team) + await session.commit() + await session.refresh(team) + return team + + async def update_team(self, workspace_id: str, team_id: str, name: str) -> Team: + async with self.database.session() as session: + team = await session.scalar( + select(Team).where( + Team.workspace_id == workspace_id, + Team.id == team_id + ) + ) + if not team: + raise Exception("Team not found") + team.name = name + await session.commit() + await session.refresh(team) + return team + + async def archive_team(self, workspace_id: str, team_id: str) -> None: + async with self.database.session() as session: + team = await session.scalar( + select(Team).where( + Team.workspace_id == workspace_id, + Team.id == team_id + ) + ) + if team: + await session.delete(team) + await session.commit() + + async def invite_member_with_hash(self, workspace_id: str, email: str, role: str, token_hash: str) -> Invitation: + # Simplified expiration for now + from datetime import datetime, timedelta, timezone + expires_at = datetime.now(timezone.utc) + timedelta(days=7) + async with self.database.session() as session: + invitation = Invitation( + workspace_id=workspace_id, + email=email, + role=role, + token_hash=token_hash, + expires_at=expires_at + ) + session.add(invitation) + await session.commit() + await session.refresh(invitation) + return invitation + + async def get_membership(self, workspace_id: str, user_id: str) -> WorkspaceMembership | None: + async with self.database.session() as session: + result = await session.execute( + select(WorkspaceMembership).where( + WorkspaceMembership.workspace_id == workspace_id, + WorkspaceMembership.user_id == user_id + ) + ) + return result.scalar_one_or_none() + + async def remove_member(self, workspace_id: str, user_id: str) -> None: + async with self.database.session() as session: + membership = await session.scalar( + select(WorkspaceMembership).where( + WorkspaceMembership.workspace_id == workspace_id, + WorkspaceMembership.user_id == user_id + ) + ) + if membership: + await session.delete(membership) + await session.commit() + + async def update_member_role(self, workspace_id: str, user_id: str, role: str) -> None: + async with self.database.session() as session: + membership = await session.scalar( + select(WorkspaceMembership).where( + WorkspaceMembership.workspace_id == workspace_id, + WorkspaceMembership.user_id == user_id + ) + ) + if membership: + membership.role = role + await session.commit() + + async def list_members(self, workspace_id: str) -> list[WorkspaceMembership]: + async with self.database.session() as session: + result = await session.scalars( + select(WorkspaceMembership).where(WorkspaceMembership.workspace_id == workspace_id) + ) + return list(result.all()) + + async def list_invitations(self, workspace_id: str) -> list[Invitation]: + async with self.database.session() as session: + result = await session.scalars( + select(Invitation).where(Invitation.workspace_id == workspace_id) + ) + return list(result.all()) + + async def count_workspace_admins(self, workspace_id: str) -> int: + async with self.database.session() as session: + result = await session.execute( + select(func.count(WorkspaceMembership.id)).where( + WorkspaceMembership.workspace_id == workspace_id, + WorkspaceMembership.role == 'admin' + ) + ) + return result.scalar_one() or 0 + + async def record_activity(self, workspace_id: str, user_id: str, action: str, entity_id: str, entity_type: str, metadata: dict[str, Any]) -> None: + async with self.database.session() as session: + activity = CollaborationActivity( + workspace_id=workspace_id, + user_id=user_id, + action=action, + entity_id=entity_id, + entity_type=entity_type, + metadata=metadata + ) + session.add(activity) + await session.commit() + + async def list_activity(self, workspace_id: str) -> list[CollaborationActivity]: + async with self.database.session() as session: + result = await session.scalars( + select(CollaborationActivity) + .where(CollaborationActivity.workspace_id == workspace_id) + .order_by(CollaborationActivity.created_at.desc()) + .limit(50) + ) + return list(result.all()) + + async def list_project_collaborators(self, project_id: str) -> list[ProjectCollaborator]: + async with self.database.session() as session: + result = await session.scalars( + select(ProjectCollaborator).where(ProjectCollaborator.project_id == project_id) + ) + return list(result.all()) + + async def add_project_collaborator(self, workspace_id: str, project_id: str, user_id: str, role: str) -> ProjectCollaborator: + async with self.database.session() as session: + collaborator = ProjectCollaborator( + workspace_id=workspace_id, + project_id=project_id, + user_id=user_id, + role=role + ) + session.add(collaborator) + await session.commit() + await session.refresh(collaborator) + return collaborator + + async def remove_project_collaborator(self, project_id: str, user_id: str) -> None: + async with self.database.session() as session: + collaborator = await session.scalar( + select(ProjectCollaborator).where( + ProjectCollaborator.project_id == project_id, + ProjectCollaborator.user_id == user_id + ) + ) + if collaborator: + await session.delete(collaborator) + await session.commit() diff --git a/app/projects/repositories/editor_repository.py b/app/projects/repositories/editor_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..a41f7833c9c23105c679582a95a51bdb400cee16 --- /dev/null +++ b/app/projects/repositories/editor_repository.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError + +from app.projects.editor_schemas import EditorDocument +from app.projects.errors import ( + ProjectAlreadyArchivedError, + ProjectEditorConflictError, + ProjectEditorNotFoundError, + ProjectNotFoundError, +) +from app.projects.models import Project, ProjectEditorState +from app.security.database import SecurityDatabase + + +class ProjectEditorRepository: + """Tenant-scoped persistence for the one current editor document.""" + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def get(self, workspace_id: str, project_id: str, *, user_id: str) -> ProjectEditorState: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project).where( + Project.id == project_id, Project.workspace_id == workspace_id + ) + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + state = await session.scalar( + select(ProjectEditorState).where( + ProjectEditorState.project_id == project_id, + ProjectEditorState.workspace_id == workspace_id, + ) + ) + if state is None: + raise ProjectEditorNotFoundError("This project has no saved editor state.") + return state + + async def save( + self, + workspace_id: str, + project_id: str, + *, + user_id: str, + expected_revision: int, + document: EditorDocument, + ) -> tuple[ProjectEditorState, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project) + .where(Project.id == project_id, Project.workspace_id == workspace_id) + .with_for_update() + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + if project.status == "archived": + raise ProjectAlreadyArchivedError("Archived projects cannot change editor state.") + current = await session.scalar( + select(ProjectEditorState) + .where( + ProjectEditorState.project_id == project_id, + ProjectEditorState.workspace_id == workspace_id, + ) + .with_for_update() + ) + if current is None: + if expected_revision != 0: + raise ProjectEditorConflictError( + "Editor state does not exist at the expected revision.", + details={"current_revision": 0, "expected_revision": expected_revision}, + ) + current = ProjectEditorState( + workspace_id=workspace_id, + project_id=project_id, + revision=1, + schema_version=document.schema_version, + state_json=document.model_dump(by_alias=True), + updated_by=user_id, + ) + session.add(current) + try: + await session.commit() + except IntegrityError as exc: + await session.rollback() + winner = await session.scalar( + select(ProjectEditorState).where( + ProjectEditorState.project_id == project_id + ) + ) + raise ProjectEditorConflictError( + "Editor state was created by another client.", + details={ + "current_revision": winner.revision if winner else 1, + "expected_revision": 0, + }, + ) from exc + await session.refresh(current) + return current, True + if current.revision != expected_revision: + raise ProjectEditorConflictError( + "Editor state was changed by another client.", + details={ + "current_revision": current.revision, + "expected_revision": expected_revision, + "updated_at": current.updated_at.isoformat(), + }, + ) + current.revision += 1 + current.schema_version = document.schema_version + current.state_json = document.model_dump(by_alias=True) + current.updated_by = user_id + current.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(current) + return current, False diff --git a/app/projects/repositories/notification_repository.py b/app/projects/repositories/notification_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..4a6a06aa47ca81431f720b88315423908ec8744f --- /dev/null +++ b/app/projects/repositories/notification_repository.py @@ -0,0 +1,43 @@ +from __future__ import annotations +from app.security.database import SecurityDatabase +from sqlalchemy import select +from app.projects.models.notifications import NotificationPreference + +class NotificationRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def get_preferences(self, workspace_id: str, user_id: str) -> list[NotificationPreference]: + async with self.database.session() as session: + result = await session.scalars( + select(NotificationPreference).where( + NotificationPreference.workspace_id == workspace_id, + NotificationPreference.user_id == user_id + ) + ) + return list(result.all()) + + async def update_preference(self, workspace_id: str, user_id: str, event_type: str, enabled: bool) -> NotificationPreference: + async with self.database.session() as session: + # Upsert logic + result = await session.execute( + select(NotificationPreference).where( + NotificationPreference.workspace_id == workspace_id, + NotificationPreference.user_id == user_id, + NotificationPreference.event_type == event_type + ) + ) + pref = result.scalar_one_or_none() + if pref: + pref.enabled = enabled + else: + pref = NotificationPreference( + workspace_id=workspace_id, + user_id=user_id, + event_type=event_type, + enabled=enabled + ) + session.add(pref) + await session.commit() + await session.refresh(pref) + return pref diff --git a/app/projects/repositories/project_repository.py b/app/projects/repositories/project_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..f7fec0cc9a96363b5a1b141c0efa5b6b5fb686c9 --- /dev/null +++ b/app/projects/repositories/project_repository.py @@ -0,0 +1,347 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import func, or_, select +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from app.generation.models import GenerationJob +from app.projects.errors import ( + ProjectAlreadyArchivedError, + ProjectAssetConflictError, + ProjectAssetNotFoundError, + ProjectJobConflictError, + ProjectJobNotFoundError, + ProjectNotFoundError, +) +from app.projects.models import Project, ProjectGenerationJob +from app.projects.schemas import ProjectStatus +from app.security.database import SecurityDatabase +from app.security.models import CanonicalMediaAsset + + +class ProjectRepository: + """Tenant-scoped persistence for authoritative project records.""" + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def create(self, project: Project, *, user_id: str) -> Project: + async with self.database.tenant_session( + workspace_id=project.workspace_id, user_id=user_id + ) as session: + session.add(project) + await session.commit() + await session.refresh(project) + return project + + async def get(self, workspace_id: str, project_id: str, *, user_id: str) -> Project: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project).where( + Project.id == project_id, + Project.workspace_id == workspace_id, + ) + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + return project + + async def list( + self, + workspace_id: str, + *, + user_id: str, + status: ProjectStatus | None, + search: str | None, + limit: int, + cursor_updated_at: datetime | None, + cursor_id: str | None, + ) -> tuple[list[Project], bool]: + query = select(Project).where(Project.workspace_id == workspace_id) + if status is not None: + query = query.where(Project.status == status.value) + if search: + term = search.casefold() + query = query.where( + or_( + func.lower(Project.name).contains(term, autoescape=True), + func.lower(Project.description).contains(term, autoescape=True), + ) + ) + if cursor_updated_at is not None and cursor_id is not None: + query = query.where( + or_( + Project.updated_at < cursor_updated_at, + (Project.updated_at == cursor_updated_at) & (Project.id < cursor_id), + ) + ) + query = query.order_by(Project.updated_at.desc(), Project.id.desc()).limit(limit + 1) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + projects = list((await session.scalars(query)).all()) + return projects[:limit], len(projects) > limit + + async def update( + self, + workspace_id: str, + project_id: str, + *, + user_id: str, + fields: dict[str, object], + ) -> Project: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project) + .where(Project.id == project_id, Project.workspace_id == workspace_id) + .with_for_update() + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + if project.status == ProjectStatus.ARCHIVED.value: + raise ProjectAlreadyArchivedError("Project is already archived.") + now = datetime.now(timezone.utc) + for field, value in fields.items(): + setattr(project, field, value) + if fields.get("status") == ProjectStatus.ARCHIVED.value: + project.archived_at = now + project.updated_at = now + await session.commit() + await session.refresh(project) + return project + + async def archive(self, workspace_id: str, project_id: str, *, user_id: str) -> Project: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project) + .where(Project.id == project_id, Project.workspace_id == workspace_id) + .with_for_update() + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + if project.status == ProjectStatus.ARCHIVED.value: + raise ProjectAlreadyArchivedError("Project is already archived.") + now = datetime.now(timezone.utc) + project.status = ProjectStatus.ARCHIVED.value + project.archived_at = now + project.updated_at = now + await session.commit() + await session.refresh(project) + return project + + @staticmethod + async def _project_for_resources( + session: AsyncSession, + workspace_id: str, + project_id: str, + *, + mutable: bool, + ) -> Project: + query = select(Project).where( + Project.id == project_id, + Project.workspace_id == workspace_id, + ) + if mutable: + query = query.with_for_update() + project = await session.scalar(query) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + if mutable and project.status == ProjectStatus.ARCHIVED.value: + raise ProjectAlreadyArchivedError("Archived projects cannot change resources.") + return project + + async def list_assets( + self, workspace_id: str, project_id: str, *, user_id: str + ) -> list[CanonicalMediaAsset]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=False) + return list( + ( + await session.scalars( + select(CanonicalMediaAsset) + .where( + CanonicalMediaAsset.workspace_id == workspace_id, + CanonicalMediaAsset.project_id == project_id, + ) + .order_by(CanonicalMediaAsset.created_at.desc()) + ) + ).all() + ) + + async def attach_asset( + self, + workspace_id: str, + project_id: str, + asset_id: str, + *, + user_id: str, + ) -> tuple[CanonicalMediaAsset, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=True) + asset = await session.scalar( + select(CanonicalMediaAsset) + .where( + CanonicalMediaAsset.id == asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + ) + .with_for_update() + ) + if asset is None: + raise ProjectAssetNotFoundError("Canonical asset was not found in this workspace.") + if asset.project_id not in (None, project_id): + raise ProjectAssetConflictError( + "Canonical asset is already attached to another project." + ) + if asset.project_id == project_id: + return asset, False + asset.project_id = project_id + await session.commit() + await session.refresh(asset) + return asset, True + + async def detach_asset( + self, + workspace_id: str, + project_id: str, + asset_id: str, + *, + user_id: str, + ) -> CanonicalMediaAsset: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=True) + asset = await session.scalar( + select(CanonicalMediaAsset) + .where( + CanonicalMediaAsset.id == asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + CanonicalMediaAsset.project_id == project_id, + ) + .with_for_update() + ) + if asset is None: + raise ProjectAssetNotFoundError("Project asset was not found in this workspace.") + asset.project_id = None + await session.commit() + await session.refresh(asset) + return asset + + async def list_generation_jobs( + self, workspace_id: str, project_id: str, *, user_id: str + ) -> list[tuple[ProjectGenerationJob, GenerationJob]]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=False) + rows = await session.execute( + select(ProjectGenerationJob, GenerationJob) + .join( + GenerationJob, + GenerationJob.id == ProjectGenerationJob.generation_job_id, + ) + .where( + ProjectGenerationJob.workspace_id == workspace_id, + ProjectGenerationJob.project_id == project_id, + GenerationJob.workspace_id == workspace_id, + ) + .order_by(GenerationJob.updated_at.desc()) + ) + return [(association, job) for association, job in rows.all()] + + async def attach_generation_job( + self, + workspace_id: str, + project_id: str, + generation_job_id: str, + *, + user_id: str, + ) -> tuple[ProjectGenerationJob, GenerationJob, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=True) + job = await session.scalar( + select(GenerationJob).where( + GenerationJob.id == generation_job_id, + GenerationJob.workspace_id == workspace_id, + ) + ) + if job is None: + raise ProjectJobNotFoundError("Generation job was not found in this workspace.") + existing = await session.scalar( + select(ProjectGenerationJob).where( + ProjectGenerationJob.generation_job_id == generation_job_id + ) + ) + if existing is not None: + if existing.project_id == project_id: + return existing, job, False + raise ProjectJobConflictError( + "Generation job is already attached to another project." + ) + association = ProjectGenerationJob( + workspace_id=workspace_id, + project_id=project_id, + generation_job_id=generation_job_id, + attached_by=user_id, + ) + session.add(association) + try: + await session.commit() + except IntegrityError as exc: + await session.rollback() + raise ProjectJobConflictError( + "Generation job could not be attached to this project." + ) from exc + await session.refresh(association) + return association, job, True + + async def detach_generation_job( + self, + workspace_id: str, + project_id: str, + generation_job_id: str, + *, + user_id: str, + ) -> tuple[ProjectGenerationJob, GenerationJob]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + await self._project_for_resources(session, workspace_id, project_id, mutable=True) + row = ( + await session.execute( + select(ProjectGenerationJob, GenerationJob) + .join( + GenerationJob, + GenerationJob.id == ProjectGenerationJob.generation_job_id, + ) + .where( + ProjectGenerationJob.workspace_id == workspace_id, + ProjectGenerationJob.project_id == project_id, + ProjectGenerationJob.generation_job_id == generation_job_id, + GenerationJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + ).one_or_none() + if row is None: + raise ProjectJobNotFoundError( + "Project generation job was not found in this workspace." + ) + association, job = row + await session.delete(association) + await session.commit() + return association, job diff --git a/app/projects/repositories/render_repository.py b/app/projects/repositories/render_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..703655389de4d1f83b9a10f2a3e4d3d8844be353 --- /dev/null +++ b/app/projects/repositories/render_repository.py @@ -0,0 +1,401 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone + +from sqlalchemy import func, or_, select, update +from sqlalchemy.exc import IntegrityError + +from app.projects.editor_schemas import EditorDocument +from app.projects.errors import ( + ProjectAlreadyArchivedError, + ProjectNotFoundError, + ProjectRenderConflictError, + ProjectRenderLimitError, + ProjectRenderNotFoundError, +) +from app.projects.models import Project, ProjectRenderJob +from app.projects.render_state import RenderStatus, validate_render_transition +from app.security.database import SecurityDatabase + + +class ProjectRenderRepository: + """Durable render queue persistence shared by API and worker processes.""" + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def create( + self, + *, + workspace_id: str, + project_id: str, + user_id: str, + editor_revision: int, + document: EditorDocument, + render_settings: dict[str, object], + request_fingerprint: str, + idempotency_key: str, + max_attempts: int, + max_active_jobs: int, + ) -> tuple[ProjectRenderJob, bool]: + job = ProjectRenderJob( + workspace_id=workspace_id, + project_id=project_id, + editor_revision=editor_revision, + editor_schema_version=document.schema_version, + editor_state_json=document.model_dump(by_alias=True), + render_settings_json=render_settings, + request_fingerprint=request_fingerprint, + idempotency_key=idempotency_key, + requested_by=user_id, + max_attempts=max_attempts, + ) + try: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + project = await session.scalar( + select(Project) + .where(Project.id == project_id, Project.workspace_id == workspace_id) + .with_for_update() + ) + if project is None: + raise ProjectNotFoundError("Project was not found in this workspace.") + if project.status == "archived": + raise ProjectAlreadyArchivedError("Archived projects cannot be rendered.") + existing = await session.scalar( + select(ProjectRenderJob).where( + ProjectRenderJob.workspace_id == workspace_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.editor_revision == editor_revision, + ProjectRenderJob.idempotency_key == idempotency_key, + ) + ) + if existing is not None: + if existing.request_fingerprint != request_fingerprint: + raise ProjectRenderConflictError( + "The idempotency key was already used for different render settings." + ) + return existing, False + active = int( + await session.scalar( + select(func.count(ProjectRenderJob.id)).where( + ProjectRenderJob.workspace_id == workspace_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.status.in_( + [ + RenderStatus.QUEUED.value, + RenderStatus.PROCESSING.value, + RenderStatus.CANCELLING.value, + ] + ), + ) + ) + or 0 + ) + if active >= max_active_jobs: + raise ProjectRenderLimitError( + "This project already has the maximum number of active renders." + ) + session.add(job) + await session.commit() + await session.refresh(job) + return job, True + except IntegrityError: + existing = await self.get_by_idempotency( + workspace_id, project_id, editor_revision, idempotency_key, user_id=user_id + ) + if existing is None: + raise + if existing.request_fingerprint != request_fingerprint: + raise ProjectRenderConflictError( + "The idempotency key was already used for different render settings." + ) + return existing, False + + async def get_by_idempotency( + self, + workspace_id: str, + project_id: str, + editor_revision: int, + idempotency_key: str, + *, + user_id: str, + ) -> ProjectRenderJob | None: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return await session.scalar( + select(ProjectRenderJob).where( + ProjectRenderJob.workspace_id == workspace_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.editor_revision == editor_revision, + ProjectRenderJob.idempotency_key == idempotency_key, + ) + ) + + async def get( + self, workspace_id: str, project_id: str, render_id: str, *, user_id: str + ) -> ProjectRenderJob: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(ProjectRenderJob).where( + ProjectRenderJob.id == render_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.workspace_id == workspace_id, + ) + ) + if job is None: + raise ProjectRenderNotFoundError("Render job was not found in this project.") + return job + + async def list( + self, workspace_id: str, project_id: str, *, user_id: str, limit: int = 50 + ) -> list[ProjectRenderJob]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return list( + ( + await session.scalars( + select(ProjectRenderJob) + .where( + ProjectRenderJob.workspace_id == workspace_id, + ProjectRenderJob.project_id == project_id, + ) + .order_by(ProjectRenderJob.created_at.desc()) + .limit(limit) + ) + ).all() + ) + + async def active_count(self, workspace_id: str, project_id: str, *, user_id: str) -> int: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + return int( + await session.scalar( + select(func.count(ProjectRenderJob.id)).where( + ProjectRenderJob.workspace_id == workspace_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.status.in_( + [ + RenderStatus.QUEUED.value, + RenderStatus.PROCESSING.value, + RenderStatus.CANCELLING.value, + ] + ), + ) + ) + or 0 + ) + + async def request_cancel( + self, + workspace_id: str, + project_id: str, + render_id: str, + *, + user_id: str, + ) -> tuple[ProjectRenderJob, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + job = await session.scalar( + select(ProjectRenderJob) + .where( + ProjectRenderJob.id == render_id, + ProjectRenderJob.project_id == project_id, + ProjectRenderJob.workspace_id == workspace_id, + ) + .with_for_update() + ) + if job is None: + raise ProjectRenderNotFoundError("Render job was not found in this project.") + changed = False + if job.status == RenderStatus.QUEUED.value: + validate_render_transition(job.status, RenderStatus.CANCELLED.value) + job.status = RenderStatus.CANCELLED.value + job.cancelled_at = datetime.now(timezone.utc) + job.completed_at = job.cancelled_at + job.next_attempt_at = None + job.error_code = None + job.error_message = None + changed = True + elif job.status == RenderStatus.PROCESSING.value: + validate_render_transition(job.status, RenderStatus.CANCELLING.value) + job.status = RenderStatus.CANCELLING.value + changed = True + if changed: + await session.commit() + await session.refresh(job) + return job, changed + + async def claim_next(self, *, stale_after_seconds: int) -> ProjectRenderJob | None: + now = datetime.now(timezone.utc) + stale = now - timedelta(seconds=stale_after_seconds) + async with self.database.session() as session: + await session.execute( + update(ProjectRenderJob) + .where( + ProjectRenderJob.status == RenderStatus.PROCESSING.value, + ProjectRenderJob.updated_at < stale, + ProjectRenderJob.attempt_count >= ProjectRenderJob.max_attempts, + ) + .values( + status=RenderStatus.FAILED.value, + error_code="RENDER_WORKER_LEASE_EXPIRED", + error_message="Render worker lease expired.", + completed_at=now, + updated_at=now, + ) + ) + await session.commit() + job = await session.scalar( + select(ProjectRenderJob) + .where( + or_( + (ProjectRenderJob.status == RenderStatus.QUEUED.value) + & ( + ProjectRenderJob.next_attempt_at.is_(None) + | (ProjectRenderJob.next_attempt_at <= now) + ), + (ProjectRenderJob.status == RenderStatus.PROCESSING.value) + & (ProjectRenderJob.updated_at < stale) + & (ProjectRenderJob.attempt_count < ProjectRenderJob.max_attempts), + ), + ) + .order_by(ProjectRenderJob.created_at) + .with_for_update(skip_locked=True) + .limit(1) + ) + if job is None: + return None + job.status = RenderStatus.PROCESSING.value + job.attempt_count += 1 + job.started_at = now + job.next_attempt_at = None + job.error_code = None + job.error_message = None + await session.commit() + await session.refresh(job) + return job + + async def heartbeat(self, render_id: str) -> ProjectRenderJob | None: + async with self.database.session() as session: + job = await session.scalar( + select(ProjectRenderJob).where(ProjectRenderJob.id == render_id).with_for_update() + ) + if job is not None and job.status == RenderStatus.PROCESSING.value: + job.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(job) + return job + + async def get_worker(self, render_id: str) -> ProjectRenderJob | None: + async with self.database.session() as session: + return await session.scalar( + select(ProjectRenderJob).where(ProjectRenderJob.id == render_id) + ) + + async def ensure_project_active(self, render_id: str) -> None: + """Fail a queued worker job before spending resources on an archive.""" + + async with self.database.session() as session: + project = await session.scalar( + select(Project) + .join(ProjectRenderJob, ProjectRenderJob.project_id == Project.id) + .where(ProjectRenderJob.id == render_id) + ) + if project is None: + raise ProjectRenderNotFoundError("Render job project was not found.") + if project.status == "archived": + raise ProjectAlreadyArchivedError("Archived projects cannot be rendered.") + + async def complete(self, render_id: str, output_asset_id: str) -> ProjectRenderJob: + async with self.database.session() as session: + job = await session.scalar( + select(ProjectRenderJob).where(ProjectRenderJob.id == render_id).with_for_update() + ) + if job is None: + raise ProjectRenderNotFoundError("Render job was not found.") + if job.status == RenderStatus.CANCELLING.value: + job.status = RenderStatus.CANCELLED.value + job.cancelled_at = datetime.now(timezone.utc) + job.output_asset_id = None + else: + validate_render_transition(job.status, RenderStatus.COMPLETED.value) + job.status = RenderStatus.COMPLETED.value + job.output_asset_id = output_asset_id + job.completed_at = datetime.now(timezone.utc) + job.next_attempt_at = None + job.error_code = None + job.error_message = None + await session.commit() + await session.refresh(job) + return job + + async def cancel_worker(self, render_id: str) -> None: + async with self.database.session() as session: + job = await session.scalar( + select(ProjectRenderJob).where(ProjectRenderJob.id == render_id).with_for_update() + ) + if job is None: + return + if job.status not in { + RenderStatus.CANCELLED.value, + RenderStatus.COMPLETED.value, + RenderStatus.FAILED.value, + }: + job.status = RenderStatus.CANCELLED.value + job.cancelled_at = datetime.now(timezone.utc) + job.completed_at = job.cancelled_at + job.next_attempt_at = None + job.output_asset_id = None + job.error_code = None + job.error_message = None + await session.commit() + + async def fail( + self, render_id: str, *, code: str, message: str, retryable: bool + ) -> ProjectRenderJob | None: + async with self.database.session() as session: + job = await session.scalar( + select(ProjectRenderJob).where(ProjectRenderJob.id == render_id).with_for_update() + ) + if job is None: + return None + if job.status in { + RenderStatus.CANCELLED.value, + RenderStatus.COMPLETED.value, + RenderStatus.FAILED.value, + }: + return job + if job.status == RenderStatus.CANCELLING.value: + validate_render_transition(job.status, RenderStatus.CANCELLED.value) + job.status = RenderStatus.CANCELLED.value + job.cancelled_at = datetime.now(timezone.utc) + job.completed_at = job.cancelled_at + job.next_attempt_at = None + job.output_asset_id = None + elif retryable and job.attempt_count < job.max_attempts: + validate_render_transition(job.status, RenderStatus.QUEUED.value) + job.status = RenderStatus.QUEUED.value + job.next_attempt_at = datetime.now(timezone.utc) + timedelta( + seconds=2 ** max(0, job.attempt_count - 1) + ) + else: + validate_render_transition(job.status, RenderStatus.FAILED.value) + job.status = RenderStatus.FAILED.value + job.completed_at = datetime.now(timezone.utc) + if job.status == RenderStatus.CANCELLED.value: + job.error_code = None + job.error_message = None + else: + job.error_code = code[:100] + job.error_message = message[:1_000] + await session.commit() + await session.refresh(job) + return job diff --git a/app/projects/schemas.py b/app/projects/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..a508fdbee540e20f9ca2b2ba8000e87915b2ad0e --- /dev/null +++ b/app/projects/schemas.py @@ -0,0 +1,158 @@ +from __future__ import annotations + +import json +import re +from datetime import datetime +from enum import Enum +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +PROJECT_NAME_MAX_LENGTH = 200 +PROJECT_DESCRIPTION_MAX_LENGTH = 4000 +PROJECT_METADATA_MAX_BYTES = 16 * 1024 + + +class ProjectStatus(str, Enum): + ACTIVE = "active" + ARCHIVED = "archived" + + +def normalize_project_name(value: str) -> str: + normalized = re.sub(r"\s+", " ", value).strip() + if not normalized: + raise ValueError("Project name cannot be blank") + if len(normalized) > PROJECT_NAME_MAX_LENGTH: + raise ValueError(f"Project name cannot exceed {PROJECT_NAME_MAX_LENGTH} characters") + return normalized + + +def validate_metadata_size(value: dict[str, Any]) -> dict[str, Any]: + try: + encoded = json.dumps( + value, ensure_ascii=False, separators=(",", ":"), allow_nan=False + ).encode("utf-8") + except (TypeError, ValueError) as exc: + raise ValueError("Project metadata must be valid JSON") from exc + if len(encoded) > PROJECT_METADATA_MAX_BYTES: + raise ValueError(f"Project metadata cannot exceed {PROJECT_METADATA_MAX_BYTES} bytes") + return value + + +class ProjectCreate(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str + description: str | None = Field(default=None, max_length=PROJECT_DESCRIPTION_MAX_LENGTH) + thumbnail_asset_id: str | None = Field(default=None, min_length=1, max_length=36) + metadata: dict[str, Any] = Field(default_factory=dict) + + _normalize_name = field_validator("name")(normalize_project_name) + _validate_metadata = field_validator("metadata")(validate_metadata_size) + + +class ProjectUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str | None = None + description: str | None = Field(default=None, max_length=PROJECT_DESCRIPTION_MAX_LENGTH) + thumbnail_asset_id: str | None = Field(default=None, min_length=1, max_length=36) + metadata: dict[str, Any] | None = None + status: ProjectStatus | None = None + + @field_validator("name") + @classmethod + def normalize_name(cls, value: str | None) -> str | None: + if value is None: + return value + return normalize_project_name(value) + + @field_validator("metadata") + @classmethod + def validate_metadata(cls, value: dict[str, Any] | None) -> dict[str, Any] | None: + return None if value is None else validate_metadata_size(value) + + @model_validator(mode="after") + def validate_patch(self) -> ProjectUpdate: + if not self.model_fields_set: + raise ValueError("At least one project field must be supplied") + if "name" in self.model_fields_set and self.name is None: + raise ValueError("Project name cannot be null") + if "metadata" in self.model_fields_set and self.metadata is None: + raise ValueError("Project metadata cannot be null") + if "status" in self.model_fields_set and self.status is None: + raise ValueError("Project status cannot be null") + return self + + +class ProjectResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + workspace_id: str + name: str + description: str | None + status: ProjectStatus + thumbnail_asset_id: str | None + brand_kit_id: str | None = None + brand_kit_version_id: str | None = None + metadata: dict[str, Any] + created_by: str + created_at: datetime + updated_at: datetime + archived_at: datetime | None + + +class ProjectListResponse(BaseModel): + items: list[ProjectResponse] + next_cursor: str | None = None + limit: int = Field(ge=1, le=100) + + +class ProjectAssetAttach(BaseModel): + model_config = ConfigDict(extra="forbid") + + asset_id: str = Field(min_length=1, max_length=36) + + +class ProjectAssetResponse(BaseModel): + id: str + project_id: str + request_id: str + filename: str + mime_type: str + file_size: int + metadata: dict[str, Any] + created_by_user_id: str | None + created_at: datetime + updated_at: datetime + + +class ProjectAssetListResponse(BaseModel): + items: list[ProjectAssetResponse] + + +class ProjectGenerationJobAttach(BaseModel): + model_config = ConfigDict(extra="forbid") + + generation_job_id: str = Field(min_length=1, max_length=36) + + +class ProjectGenerationJobResponse(BaseModel): + id: str + project_id: str + job_type: Literal["generation"] = "generation" + generation_request_id: str + provider: str + status: str + output_asset_id: str | None + error_code: str | None + created_at: datetime + started_at: datetime | None + completed_at: datetime | None + updated_at: datetime + attached_at: datetime + + +class ProjectGenerationJobListResponse(BaseModel): + items: list[ProjectGenerationJobResponse] diff --git a/app/projects/schemas/approval.py b/app/projects/schemas/approval.py new file mode 100644 index 0000000000000000000000000000000000000000..cf590a65ee384a49b1095004154a0b58f2f1e3c2 --- /dev/null +++ b/app/projects/schemas/approval.py @@ -0,0 +1,17 @@ +from pydantic import BaseModel + +class ApprovalRequest(BaseModel): + id: str + workflow_id: str + project_id: str + editor_revision: int + status: str + created_by: str + created_at: str + +class ReviewComment(BaseModel): + id: str + request_id: str + user_id: str + content: str + created_at: str diff --git a/app/projects/schemas/collaboration.py b/app/projects/schemas/collaboration.py new file mode 100644 index 0000000000000000000000000000000000000000..3003944d449d86d786378200a33d8021db5668e1 --- /dev/null +++ b/app/projects/schemas/collaboration.py @@ -0,0 +1,31 @@ +from __future__ import annotations +from pydantic import BaseModel, Field + +class MemberBase(BaseModel): + user_id: str + role: str + +class MemberResponse(MemberBase): + id: str + workspace_id: str + created_at: str + +class InvitationCreate(BaseModel): + email: str = Field(..., email=True) + role: str + +class InvitationResponse(InvitationCreate): + id: str + workspace_id: str + status: str + expires_at: str + created_at: str + token: str | None = None + +class TeamBase(BaseModel): + name: str + +class TeamResponse(TeamBase): + id: str + workspace_id: str + created_at: str diff --git a/app/projects/services/__init__.py b/app/projects/services/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..229370b2a3961bc05027a27704d92532a9a19be9 --- /dev/null +++ b/app/projects/services/__init__.py @@ -0,0 +1 @@ +"""Project application services.""" diff --git a/app/projects/services/approval_service.py b/app/projects/services/approval_service.py new file mode 100644 index 0000000000000000000000000000000000000000..16cec16292e8f3d2b3e4f162462d7fb862180953 --- /dev/null +++ b/app/projects/services/approval_service.py @@ -0,0 +1,37 @@ +from __future__ import annotations +from app.projects.repositories.approval_repository import ApprovalRepository +from app.projects.repositories.editor_repository import ProjectEditorRepository +from app.projects.models.collaboration import ApprovalWorkflow, ApprovalRequest, ReviewComment + +class ApprovalService: + def __init__(self, repository: ApprovalRepository, editor_repository: ProjectEditorRepository) -> None: + self.repository = repository + self.editor_repository = editor_repository + + async def create_workflow(self, workspace_id: str, project_id: str, name: str) -> ApprovalWorkflow: + return await self.repository.create_workflow(workspace_id, project_id, name) + + async def create_request(self, workspace_id: str, workflow_id: str, project_id: str, user_id: str) -> ApprovalRequest: + state = await self.editor_repository.get_state(workspace_id, project_id) + return await self.repository.create_request(workspace_id, workflow_id, project_id, user_id, state.revision) + + async def approve_request(self, request_id: str, actor_user_id: str) -> ApprovalRequest: + request = await self.repository.get_request(request_id) + if request.created_by == actor_user_id: + raise Exception("Separation of duties: Cannot approve your own submission.") + + # Verify revision hasn't changed + state = await self.editor_repository.get_state(request.workspace_id, request.project_id) + if state.revision != request.editor_revision: + raise Exception("Content has changed since request; approval is stale.") + + return await self.repository.update_request_status(request_id, "approved") + + async def reject_request(self, request_id: str, actor_user_id: str) -> ApprovalRequest: + return await self.repository.update_request_status(request_id, "rejected") + + async def add_comment(self, request_id: str, user_id: str, workspace_id: str, content: str) -> ReviewComment: + return await self.repository.add_review_comment(request_id, user_id, workspace_id, content) + + async def list_requests(self, workflow_id: str) -> list[ApprovalRequest]: + return await self.repository.list_requests(workflow_id) diff --git a/app/projects/services/collaboration_service.py b/app/projects/services/collaboration_service.py new file mode 100644 index 0000000000000000000000000000000000000000..db3f29930a68333ec0b1d57ecdf910c398ddc13e --- /dev/null +++ b/app/projects/services/collaboration_service.py @@ -0,0 +1,134 @@ +from __future__ import annotations +import secrets +import hashlib +from typing import Any +from app.projects.repositories.collaboration_repository import CollaborationRepository +from app.projects.schemas.collaboration import TeamResponse, InvitationResponse, MemberResponse +from app.security.models import WorkspaceMembership +from app.projects.errors import ( + CollaborationUnauthorizedError, + TeamNotFoundError, + TeamMemberNotFoundError, +) + +class CollaborationService: + def __init__(self, repository: CollaborationRepository) -> None: + self.repository = repository + + async def list_teams(self, workspace_id: str) -> list[TeamResponse]: + teams = await self.repository.list_teams(workspace_id) + return [TeamResponse(id=t.id, workspace_id=t.workspace_id, name=t.name, created_at=t.created_at.isoformat()) for t in teams] + + async def create_team(self, workspace_id: str, name: str) -> TeamResponse: + team = await self.repository.create_team(workspace_id, name) + return TeamResponse(id=team.id, workspace_id=team.workspace_id, name=team.name, created_at=team.created_at.isoformat()) + + async def update_team(self, workspace_id: str, team_id: str, name: str) -> TeamResponse: + try: + team = await self.repository.update_team(workspace_id, team_id, name) + except Exception: + raise TeamNotFoundError(f"Team {team_id} not found.") + return TeamResponse(id=team.id, workspace_id=team.workspace_id, name=team.name, created_at=team.created_at.isoformat()) + + async def archive_team(self, workspace_id: str, team_id: str) -> None: + await self.repository.archive_team(workspace_id, team_id) + + async def invite_member(self, workspace_id: str, email: str, role: str) -> InvitationResponse: + if role not in {"admin", "member"}: + raise ValueError("Invitation role must be admin or member.") + token = secrets.token_urlsafe(32) + token_hash = hashlib.sha256(token.encode()).hexdigest() + + invitation = await self.repository.invite_member_with_hash(workspace_id, email, role, token_hash) + + response = InvitationResponse( + id=invitation.id, + workspace_id=invitation.workspace_id, + email=invitation.email, + role=invitation.role, + status=invitation.status, + expires_at=invitation.expires_at.isoformat(), + created_at=invitation.created_at.isoformat() + ) + setattr(response, "token", token) + return response + + async def list_members(self, workspace_id: str) -> list[MemberResponse]: + members = await self.repository.list_members(workspace_id) + return [MemberResponse(id=m.id, workspace_id=m.workspace_id, user_id=m.user_id, role=m.role, created_at=m.created_at.isoformat()) for m in members] + + async def list_invitations(self, workspace_id: str) -> list[InvitationResponse]: + invitations = await self.repository.list_invitations(workspace_id) + return [InvitationResponse( + id=i.id, + workspace_id=i.workspace_id, + email=i.email, + role=i.role, + status=i.status, + expires_at=i.expires_at.isoformat(), + created_at=i.created_at.isoformat() + ) for i in invitations] + + async def remove_member(self, workspace_id: str, actor_user_id: str, target_user_id: str) -> None: + # Check permissions: actor must be admin + actor_membership = await self.repository.get_membership(workspace_id, actor_user_id) + if not actor_membership or actor_membership.role != 'admin': + raise CollaborationUnauthorizedError("Unauthorized: Only admins can remove members.") + + # Prevent removing last admin + if actor_membership.user_id == target_user_id: + raise CollaborationUnauthorizedError("Cannot remove yourself.") + + target_membership = await self.repository.get_membership(workspace_id, target_user_id) + if not target_membership: + raise TeamMemberNotFoundError("Member not found.") + + if target_membership.role == 'admin': + admin_count = await self.repository.count_workspace_admins(workspace_id) + if admin_count <= 1: + raise CollaborationUnauthorizedError("Cannot remove the last administrator.") + + # Perform removal + await self.repository.remove_member(workspace_id, target_user_id) + + async def update_member_role(self, workspace_id: str, actor_user_id: str, target_user_id: str, new_role: str) -> None: + # Check permissions: actor must be admin + actor_membership = await self.repository.get_membership(workspace_id, actor_user_id) + if not actor_membership or actor_membership.role != 'admin': + raise CollaborationUnauthorizedError("Unauthorized: Only admins can update member roles.") + + # Prevent self-elevation + if actor_user_id == target_user_id and new_role == 'admin': + raise CollaborationUnauthorizedError("Cannot elevate your own privileges.") + + # Perform update + await self.repository.update_member_role(workspace_id, target_user_id, new_role) + + async def get_membership(self, workspace_id: str, user_id: str) -> WorkspaceMembership | None: + return await self.repository.get_membership(workspace_id, user_id) + + async def list_project_collaborators(self, project_id: str) -> list[MemberResponse]: + collaborators = await self.repository.list_project_collaborators(project_id) + return [MemberResponse(id=c.id, workspace_id=c.workspace_id, user_id=c.user_id, role=c.role, created_at=c.created_at.isoformat()) for c in collaborators] + + async def add_project_collaborator(self, workspace_id: str, project_id: str, user_id: str, role: str) -> MemberResponse: + collaborator = await self.repository.add_project_collaborator(workspace_id, project_id, user_id, role) + return MemberResponse(id=collaborator.id, workspace_id=collaborator.workspace_id, user_id=collaborator.user_id, role=collaborator.role, created_at=collaborator.created_at.isoformat()) + + async def record_activity(self, workspace_id: str, user_id: str, action: str, entity_id: str, entity_type: str, metadata: dict[str, Any]) -> None: + await self.repository.record_activity(workspace_id, user_id, action, entity_id, entity_type, metadata) + + async def list_activity(self, workspace_id: str) -> list[dict[str, Any]]: + activities = await self.repository.list_activity(workspace_id) + return [ + { + "id": a.id, + "user_id": a.user_id, + "action": a.action, + "entity_id": a.entity_id, + "entity_type": a.entity_type, + "metadata": a.metadata, + "created_at": a.created_at.isoformat(), + } + for a in activities + ] diff --git a/app/projects/services/editor_service.py b/app/projects/services/editor_service.py new file mode 100644 index 0000000000000000000000000000000000000000..3bb5537d5bc9bcf8330d2fea6f0735386094123a --- /dev/null +++ b/app/projects/services/editor_service.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +import time + +from app.core.logger import get_logger +from app.projects.editor_schemas import ( + EDITOR_STATE_MAX_BYTES, + AudioClip, + EditorDocument, + EditorSaveRequest, + EditorStateResponse, + MediaClip, +) +from app.projects.errors import ProjectEditorAssetInvalidError, ProjectEditorInvalidError +from app.projects.models import ProjectEditorState +from app.projects.repositories.editor_repository import ProjectEditorRepository +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService + +logger = get_logger(__name__) + + +class ProjectEditorService: + def __init__( + self, + repository: ProjectEditorRepository, + assets: CanonicalAssetService, + audit: AuditService, + *, + state_max_bytes: int = EDITOR_STATE_MAX_BYTES, + max_tracks: int = 32, + max_clips: int = 500, + max_duration_seconds: int = 3600, + ) -> None: + self.repository = repository + self.assets = assets + self.audit = audit + self.state_max_bytes = state_max_bytes + self.max_tracks = max_tracks + self.max_clips = max_clips + self.max_duration_ms = max_duration_seconds * 1000 + + async def get(self, *, workspace_id: str, user_id: str, project_id: str) -> EditorStateResponse: + return self._response(await self.repository.get(workspace_id, project_id, user_id=user_id)) + + async def save( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + payload: EditorSaveRequest, + ) -> EditorStateResponse: + started = time.monotonic() + document = self._validate_document(payload.state, project_id) + await self._validate_assets(workspace_id, user_id, document) + try: + state, created = await self.repository.save( + workspace_id, + project_id, + user_id=user_id, + expected_revision=payload.expected_revision, + document=document, + ) + except Exception as exc: + if getattr(exc, "code", None) == "PROJECT_EDITOR_REVISION_CONFLICT": + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.editor_conflict", + entity_type="project", + entity_id=project_id, + metadata={"expected_revision": payload.expected_revision}, + ) + logger.info( + "project editor conflict", + extra={ + "operation": "editor.conflict", + "project_id": project_id, + "workspace_id": workspace_id, + "expected_revision": payload.expected_revision, + }, + ) + raise + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.editor_created" if created else "project.editor_updated", + entity_type="project", + entity_id=project_id, + metadata={"revision": state.revision, "schema_version": state.schema_version}, + ) + logger.info( + "project editor saved", + extra={ + "operation": "editor.save", + "project_id": project_id, + "workspace_id": workspace_id, + "revision": state.revision, + "duration_ms": round((time.monotonic() - started) * 1000), + }, + ) + return self._response(state) + + def _validate_document(self, document: EditorDocument, project_id: str) -> EditorDocument: + if document.project_id != project_id: + raise ProjectEditorInvalidError( + "Editor document project does not match the URL project." + ) + if len(document.timeline.tracks) > self.max_tracks: + raise ProjectEditorInvalidError("Editor document contains too many tracks.") + clip_count = sum(len(track.clips) for track in document.timeline.tracks) + if clip_count > self.max_clips: + raise ProjectEditorInvalidError("Editor document contains too many clips.") + if document.duration_ms() > self.max_duration_ms: + raise ProjectEditorInvalidError( + "Editor timeline exceeds the configured duration limit." + ) + try: + encoded = document.json_bytes() + except (TypeError, ValueError) as exc: + raise ProjectEditorInvalidError("Editor document contains invalid JSON.") from exc + if len(encoded) > self.state_max_bytes: + raise ProjectEditorInvalidError("Editor document exceeds the configured size limit.") + return document + + async def _validate_assets( + self, workspace_id: str, user_id: str, document: EditorDocument + ) -> None: + owned = {} + for asset_id in document.asset_ids(): + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, user_id=user_id, asset_id=asset_id + ) + if asset.project_id != document.project_id: + raise CanonicalAssetNotFoundError("Asset is not attached to this project.") + owned[asset_id] = asset + except CanonicalAssetNotFoundError as exc: + raise ProjectEditorAssetInvalidError( + "Every editor asset reference must belong to the authenticated workspace." + ) from exc + for track in document.timeline.tracks: + for clip in track.clips: + if not isinstance(clip, (MediaClip, AudioClip)): + continue + asset = owned[clip.asset_id] + if isinstance(clip, AudioClip) and not asset.mime_type.startswith("audio/"): + raise ProjectEditorAssetInvalidError( + "An audio clip must reference an audio asset." + ) + if isinstance(clip, MediaClip) and not asset.mime_type.startswith( + f"{clip.media_type}/" + ): + raise ProjectEditorAssetInvalidError( + f"A {clip.media_type} clip must reference a {clip.media_type} asset." + ) + metadata = asset.metadata_json or {} + duration_ms = metadata.get("duration_ms") + if duration_ms is None and isinstance(metadata.get("duration"), (int, float)): + duration_ms = float(metadata["duration"]) * 1000 + if isinstance(duration_ms, (int, float)) and ( + clip.source_start_ms + clip.source_duration_ms > duration_ms + 1 + ): + raise ProjectEditorAssetInvalidError( + "A clip source range exceeds its asset duration." + ) + + @staticmethod + def _response(state: ProjectEditorState) -> EditorStateResponse: + return EditorStateResponse( + project_id=state.project_id, + revision=state.revision, + schema_version=state.schema_version, + state=EditorDocument.model_validate(state.state_json), + created_at=state.created_at, + updated_at=state.updated_at, + updated_by=state.updated_by, + ) diff --git a/app/projects/services/notification_service.py b/app/projects/services/notification_service.py new file mode 100644 index 0000000000000000000000000000000000000000..b58edbbd65980709851757490f07f6cc5270985a --- /dev/null +++ b/app/projects/services/notification_service.py @@ -0,0 +1,14 @@ +from __future__ import annotations +from app.projects.repositories.notification_repository import NotificationRepository + +class NotificationService: + def __init__(self, repository: NotificationRepository) -> None: + self.repository = repository + + async def get_preferences(self, workspace_id: str, user_id: str) -> list[dict]: + prefs = await self.repository.get_preferences(workspace_id, user_id) + return [{"event_type": p.event_type, "enabled": p.enabled} for p in prefs] + + async def update_preference(self, workspace_id: str, user_id: str, event_type: str, enabled: bool) -> dict: + pref = await self.repository.update_preference(workspace_id, user_id, event_type, enabled) + return {"event_type": pref.event_type, "enabled": pref.enabled} diff --git a/app/projects/services/project_service.py b/app/projects/services/project_service.py new file mode 100644 index 0000000000000000000000000000000000000000..5b28f25ce825a349921554e1944eeef9fb52953c --- /dev/null +++ b/app/projects/services/project_service.py @@ -0,0 +1,471 @@ +from __future__ import annotations + +import base64 +import binascii +import json +import time +from datetime import datetime, timezone +from typing import Any +from uuid import UUID + +from app.core.logger import get_logger +from app.generation.models import GenerationJob +from app.projects.errors import ( + ProjectInvalidCursorError, + ProjectInvalidNameError, + ProjectInvalidStatusError, + ProjectThumbnailInvalidError, +) +from app.projects.models import Project, ProjectGenerationJob +from app.projects.repositories.project_repository import ProjectRepository +from app.projects.schemas import ( + ProjectAssetListResponse, + ProjectAssetResponse, + ProjectCreate, + ProjectGenerationJobListResponse, + ProjectGenerationJobResponse, + ProjectListResponse, + ProjectResponse, + ProjectStatus, + ProjectUpdate, + normalize_project_name, +) +from app.brand.repositories.brand_repository import BrandKitRepository +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService +from app.security.models import CanonicalMediaAsset + +logger = get_logger(__name__) + + +class ProjectService: + """Coordinates validation, tenant ownership, persistence, and audit.""" + + def __init__( + self, + repository: ProjectRepository, + assets: CanonicalAssetService, + audit: AuditService, + brand_kits: BrandKitRepository, + ) -> None: + self.repository = repository + self.assets = assets + self.audit = audit + self.brand_kits = brand_kits + + async def get_active_brand_kit(self, *, workspace_id: str, project_id: str, user_id: str) -> Any | None: + project = await self.repository.get(workspace_id, project_id, user_id=user_id) + if not project.brand_kit_id: + return None + return await self.brand_kits.get_latest_version(workspace_id, project.brand_kit_id, user_id=user_id) + + async def create( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + payload: ProjectCreate, + ) -> ProjectResponse: + started = time.monotonic() + try: + name = normalize_project_name(payload.name) + except ValueError as exc: + raise ProjectInvalidNameError(str(exc)) from exc + await self._validate_thumbnail(workspace_id, user_id, payload.thumbnail_asset_id) + project = await self.repository.create( + Project( + workspace_id=workspace_id, + created_by=user_id, + name=name, + description=payload.description, + status=ProjectStatus.ACTIVE.value, + thumbnail_asset_id=payload.thumbnail_asset_id, + metadata_json=dict(payload.metadata), + ), + user_id=user_id, + ) + await self._audit( + project, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.created", + metadata={"has_thumbnail": project.thumbnail_asset_id is not None}, + ) + self._log("create", project, started, "success") + return self._response(project) + + async def list( + self, + *, + workspace_id: str, + user_id: str, + status: ProjectStatus | None, + search: str | None, + limit: int, + cursor: str | None, + ) -> ProjectListResponse: + cursor_updated_at, cursor_id = self._decode_cursor(cursor) + projects, has_more = await self.repository.list( + workspace_id, + user_id=user_id, + status=status, + search=search.strip() if search else None, + limit=limit, + cursor_updated_at=cursor_updated_at, + cursor_id=cursor_id, + ) + next_cursor = self._encode_cursor(projects[-1]) if has_more and projects else None + + # NOTE: For listing, we might want to optimize this to avoid N+1 queries. + # Keeping it simple for now as requested. + responses = [] + for project in projects: + active_version = None + if project.brand_kit_id: + active_version = await self.brand_kits.get_latest_version(workspace_id, project.brand_kit_id, user_id=user_id) + responses.append(self._response(project, active_version.id if active_version else None)) + + return ProjectListResponse( + items=responses, + next_cursor=next_cursor, + limit=limit, + ) + + async def get(self, *, workspace_id: str, user_id: str, project_id: str) -> ProjectResponse: + project = await self.repository.get(workspace_id, project_id, user_id=user_id) + active_version = None + if project.brand_kit_id: + active_version = await self.brand_kits.get_latest_version(workspace_id, project.brand_kit_id, user_id=user_id) + return self._response(project, active_version.id if active_version else None) + + async def update( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + payload: ProjectUpdate, + ) -> ProjectResponse: + started = time.monotonic() + fields = payload.model_dump(exclude_unset=True) + if "name" in fields: + try: + fields["name"] = normalize_project_name(str(fields["name"])) + except ValueError as exc: + raise ProjectInvalidNameError(str(exc)) from exc + if "thumbnail_asset_id" in fields: + await self._validate_thumbnail(workspace_id, user_id, fields["thumbnail_asset_id"]) + if "metadata" in fields: + fields["metadata_json"] = fields.pop("metadata") + if "status" in fields: + status = fields["status"] + try: + fields["status"] = ProjectStatus(status).value + except ValueError as exc: + raise ProjectInvalidStatusError("Project status is invalid.") from exc + project = await self.repository.update( + workspace_id, + project_id, + user_id=user_id, + fields=fields, + ) + archived = fields.get("status") == ProjectStatus.ARCHIVED.value + await self._audit( + project, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.archived" if archived else "project.updated", + metadata={"fields": sorted(payload.model_fields_set)}, + ) + self._log("update", project, started, "success") + return self._response(project) + + async def delete( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + ) -> None: + started = time.monotonic() + project = await self.repository.archive(workspace_id, project_id, user_id=user_id) + await self._audit( + project, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.deleted", + metadata={"disposition": "archived"}, + ) + self._log("delete", project, started, "success") + + async def list_assets( + self, *, workspace_id: str, user_id: str, project_id: str + ) -> ProjectAssetListResponse: + assets = await self.repository.list_assets(workspace_id, project_id, user_id=user_id) + return ProjectAssetListResponse(items=[self._asset_response(asset) for asset in assets]) + + async def attach_asset( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + asset_id: str, + ) -> ProjectAssetResponse: + asset, attached = await self.repository.attach_asset( + workspace_id, project_id, asset_id, user_id=user_id + ) + if attached: + await self._audit_resource( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=project_id, + event_type="project.asset_attached", + resource_id=asset.id, + ) + return self._asset_response(asset) + + async def detach_asset( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + asset_id: str, + ) -> None: + asset = await self.repository.detach_asset( + workspace_id, project_id, asset_id, user_id=user_id + ) + await self._audit_resource( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=project_id, + event_type="project.asset_detached", + resource_id=asset.id, + ) + + async def list_generation_jobs( + self, *, workspace_id: str, user_id: str, project_id: str + ) -> ProjectGenerationJobListResponse: + jobs = await self.repository.list_generation_jobs(workspace_id, project_id, user_id=user_id) + return ProjectGenerationJobListResponse( + items=[self._job_response(association, job) for association, job in jobs] + ) + + async def attach_generation_job( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + generation_job_id: str, + ) -> ProjectGenerationJobResponse: + association, job, attached = await self.repository.attach_generation_job( + workspace_id, project_id, generation_job_id, user_id=user_id + ) + if attached: + await self._audit_resource( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=project_id, + event_type="project.job_attached", + resource_id=job.id, + ) + return self._job_response(association, job) + + async def detach_generation_job( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + generation_job_id: str, + ) -> None: + _, job = await self.repository.detach_generation_job( + workspace_id, project_id, generation_job_id, user_id=user_id + ) + await self._audit_resource( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + project_id=project_id, + event_type="project.job_detached", + resource_id=job.id, + ) + + async def _validate_thumbnail(self, workspace_id: str, user_id: str, asset_id: Any) -> None: + if asset_id is None: + return + if not isinstance(asset_id, str): + raise ProjectThumbnailInvalidError("Project thumbnail reference is invalid.") + try: + await self.assets.get_owned_by_id( + workspace_id=workspace_id, user_id=user_id, asset_id=asset_id + ) + except CanonicalAssetNotFoundError as exc: + raise ProjectThumbnailInvalidError( + "Project thumbnail is not an owned canonical asset." + ) from exc + + async def _audit( + self, + project: Project, + *, + user_id: str, + api_key_id: str, + request_id: str, + event_type: str, + metadata: dict[str, object], + ) -> None: + await self.audit.record_event( + workspace_id=project.workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type=event_type, + entity_type="project", + entity_id=project.id, + metadata=metadata, + ) + + async def _audit_resource( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + event_type: str, + resource_id: str, + ) -> None: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type=event_type, + entity_type="project", + entity_id=project_id, + metadata={"resource_id": resource_id}, + ) + + @staticmethod + def _response(project: Project, active_version_id: str | None = None) -> ProjectResponse: + return ProjectResponse( + id=project.id, + workspace_id=project.workspace_id, + name=project.name, + description=project.description, + status=ProjectStatus(project.status), + thumbnail_asset_id=project.thumbnail_asset_id, + brand_kit_id=project.brand_kit_id, + brand_kit_version_id=active_version_id, + metadata=dict(project.metadata_json or {}), + created_by=project.created_by, + created_at=project.created_at, + updated_at=project.updated_at, + archived_at=project.archived_at, + ) + + @staticmethod + def _asset_response(asset: CanonicalMediaAsset) -> ProjectAssetResponse: + return ProjectAssetResponse( + id=asset.id, + project_id=str(asset.project_id), + request_id=asset.request_id, + filename=asset.filename, + mime_type=asset.mime_type, + file_size=asset.file_size, + metadata=dict(asset.metadata_json or {}), + created_by_user_id=asset.created_by_user_id, + created_at=asset.created_at, + updated_at=asset.updated_at, + ) + + @staticmethod + def _job_response( + association: ProjectGenerationJob, job: GenerationJob + ) -> ProjectGenerationJobResponse: + return ProjectGenerationJobResponse( + id=job.id, + project_id=association.project_id, + generation_request_id=job.generation_request_id, + provider=job.provider, + status=job.status, + output_asset_id=job.output_asset_id, + error_code=job.error_code, + created_at=job.created_at, + started_at=job.started_at, + completed_at=job.completed_at, + updated_at=job.updated_at, + attached_at=association.created_at, + ) + + @staticmethod + def _encode_cursor(project: Project) -> str: + payload = json.dumps( + {"updated_at": project.updated_at.isoformat(), "id": project.id}, + separators=(",", ":"), + ).encode("utf-8") + return base64.urlsafe_b64encode(payload).decode("ascii").rstrip("=") + + @staticmethod + def _decode_cursor(cursor: str | None) -> tuple[datetime | None, str | None]: + if cursor is None: + return None, None + try: + padding = "=" * (-len(cursor) % 4) + payload = json.loads(base64.urlsafe_b64decode(cursor + padding)) + updated_at = datetime.fromisoformat(payload["updated_at"]) + if updated_at.tzinfo is None: + updated_at = updated_at.replace(tzinfo=timezone.utc) + project_id = payload["id"] + if not isinstance(project_id, str): + raise ValueError + UUID(project_id) + return updated_at, project_id + except (KeyError, TypeError, ValueError, binascii.Error, json.JSONDecodeError) as exc: + raise ProjectInvalidCursorError("Project pagination cursor is invalid.") from exc + + @staticmethod + def _log( + operation: str, + project: Project, + started: float, + result: str, + ) -> None: + logger.info( + "project operation completed", + extra={ + "operation": operation, + "project_id": project.id, + "workspace_id": project.workspace_id, + "duration_ms": max(0, round((time.monotonic() - started) * 1000)), + "result": result, + }, + ) diff --git a/app/projects/services/render_compiler.py b/app/projects/services/render_compiler.py new file mode 100644 index 0000000000000000000000000000000000000000..67d7881fe839aa2aef724225cd4c170c9eaf610e --- /dev/null +++ b/app/projects/services/render_compiler.py @@ -0,0 +1,213 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +from app.projects.editor_schemas import AudioClip, EditorDocument, MediaClip +from app.projects.errors import ProjectRenderInvalidError, ProjectRenderUnsupportedError + + +@dataclass(frozen=True, slots=True) +class RenderInput: + asset_id: str + path: Path + kind: str + source_start_ms: int + duration_ms: int + input_index: int + + +@dataclass(frozen=True, slots=True) +class RenderPlan: + args: tuple[str | Path, ...] + inputs: tuple[RenderInput, ...] + duration_ms: int + output_extension: str + + +def validate_renderable_document(document: EditorDocument) -> None: + if document.timeline.transitions: + raise ProjectRenderUnsupportedError( + "Timeline transitions are not supported by the render compiler yet." + ) + for track in document.timeline.tracks: + if not track.visible: + continue + for clip in track.clips: + if clip.visible and clip.kind in {"caption", "effect"}: + raise ProjectRenderUnsupportedError( + f"The render compiler does not support {clip.kind} clips yet." + ) + + +def compile_render( + document: EditorDocument, + *, + asset_paths: dict[str, tuple[Path, str]], + width: int, + height: int, + frame_rate: float, + output_format: str, + quality: str, + preset: str, +) -> RenderPlan: + """Compile the currently supported editor primitives into FFmpeg args. + + The compiler is deterministic: all inputs, offsets, trims, canvas settings, + and codecs are derived solely from the immutable editor snapshot/settings. + Captions, effects, transitions, and unsupported clip types fail closed. + """ + if output_format not in {"mp4", "webm"}: + raise ProjectRenderInvalidError("Unsupported render output format.") + if width % 2 or height % 2 or width < 2 or height < 2: + raise ProjectRenderInvalidError("Render resolution must be positive and even.") + validate_renderable_document(document) + inputs: list[RenderInput] = [] + video_clips: list[tuple[MediaClip, RenderInput]] = [] + audio_clips: list[tuple[AudioClip, RenderInput]] = [] + for track in sorted(document.timeline.tracks, key=lambda item: item.order): + if not track.visible: + continue + for clip in track.clips: + if not clip.visible: + continue + if clip.kind in {"caption", "effect"}: + raise ProjectRenderUnsupportedError( + f"The render compiler does not support {clip.kind} clips yet." + ) + if clip.kind == "audio" and track.muted: + continue + if clip.asset_id not in asset_paths: + raise ProjectRenderInvalidError("A referenced media asset is unavailable.") + path, mime_type = asset_paths[clip.asset_id] + if not path.is_file(): + raise ProjectRenderInvalidError("A referenced media asset file is unavailable.") + input_index = len(inputs) + 1 # 0 is the generated black canvas. + item = RenderInput( + asset_id=clip.asset_id, + path=path, + kind=clip.kind, + source_start_ms=clip.source_start_ms, + duration_ms=clip.duration_ms, + input_index=input_index, + ) + inputs.append(item) + if clip.kind == "media": + if clip.media_type == "image" and not mime_type.startswith("image/"): + raise ProjectRenderInvalidError("An image clip references a non-image asset.") + if clip.media_type == "video" and not mime_type.startswith("video/"): + raise ProjectRenderInvalidError("A video clip references a non-video asset.") + video_clips.append((clip, item)) + elif clip.kind == "audio": + if not mime_type.startswith("audio/"): + raise ProjectRenderInvalidError("An audio clip references a non-audio asset.") + audio_clips.append((clip, item)) + if not video_clips and not audio_clips: + raise ProjectRenderInvalidError("The editor timeline contains no renderable media.") + # Hidden tracks/clips are non-rendering state and must not extend the + # generated canvas with a long black tail. + duration_ms = max( + (clip.start_ms + clip.duration_ms for clip, _ in [*video_clips, *audio_clips]), + default=0, + ) + + args: list[str | Path] = [ + "-f", + "lavfi", + "-t", + f"{duration_ms / 1000:.6f}", + "-i", + f"color=c=black:s={width}x{height}:r={frame_rate:.6g}", + ] + for item in inputs: + if item.kind == "media" and any( + clip.media_type == "image" and ref is item for clip, ref in video_clips + ): + args.extend(["-loop", "1"]) + if item.source_start_ms: + args.extend(["-ss", f"{item.source_start_ms / 1000:.6f}"]) + args.extend(["-t", f"{item.duration_ms / 1000:.6f}", "-i", item.path]) + + filters: list[str] = [] + current_video = "base0" + filters.append(f"[0:v]setpts=PTS-STARTPTS[{current_video}]") + for number, (clip, item) in enumerate(video_clips, start=1): + label = f"clipv{number}" + scale_x = f"{clip.transform.scale_x:.6g}" + scale_y = f"{clip.transform.scale_y:.6g}" + chain = f"setpts=PTS-STARTPTS,scale=trunc(iw*{scale_x}/2)*2:trunc(ih*{scale_y}/2)*2" + if abs(clip.transform.rotation) > 1e-9: + chain += f",rotate={clip.transform.rotation:.6g}*PI/180:c=none:ow=rotw(iw):oh=roth(ih)" + if clip.opacity < 1: + chain += f",format=rgba,colorchannelmixer=aa={clip.opacity:.6g}" + chain += f",setpts=PTS+{clip.start_ms / 1000:.6f}/TB" + filters.append(f"[{item.input_index}:v]{chain}[{label}]") + next_video = f"mixv{number}" + x = f"(W-w)/2+{clip.transform.x:.6g}" + y = f"(H-h)/2+{clip.transform.y:.6g}" + filters.append( + f"[{current_video}][{label}]overlay=x={x}:y={y}:eof_action=pass:shortest=0[{next_video}]" + ) + current_video = next_video + + audio_labels: list[str] = [] + for number, (clip, item) in enumerate(audio_clips, start=1): + label = f"clipa{number}" + chain = f"atrim=duration={clip.duration_ms / 1000:.6f},asetpts=PTS-STARTPTS,volume={clip.volume:.6g}" + if clip.fade_in_ms: + chain += f",afade=t=in:st=0:d={clip.fade_in_ms / 1000:.6f}" + if clip.fade_out_ms: + chain += f",afade=t=out:st={(clip.duration_ms - clip.fade_out_ms) / 1000:.6f}:d={clip.fade_out_ms / 1000:.6f}" + chain += f",adelay={clip.start_ms}|{clip.start_ms}" + filters.append(f"[{item.input_index}:a]{chain}[{label}]") + audio_labels.append(label) + if audio_labels: + filters.append( + "".join(f"[{label}]" for label in audio_labels) + + f"amix=inputs={len(audio_labels)}:duration=longest:dropout_transition=0[aout]" + ) + + args.extend(["-filter_complex", ";".join(filters), "-map", f"[{current_video}]"]) + if audio_labels: + args.extend(["-map", "[aout]"]) + else: + args.append("-an") + if output_format == "webm": + args.extend( + [ + "-c:v", + "libvpx-vp9", + "-crf", + {"draft": "30", "standard": "24", "high": "18"}[quality], + "-b:v", + "0", + "-deadline", + "good", + "-cpu-used", + {"fast": "5", "balanced": "2", "quality": "0"}[preset], + "-c:a", + "libopus", + ] + ) + else: + crf = {"draft": "28", "standard": "23", "high": "18"}[quality] + args.extend( + [ + "-c:v", + "libx264", + "-pix_fmt", + "yuv420p", + "-preset", + {"fast": "veryfast", "balanced": "medium", "quality": "slow"}[preset], + "-crf", + crf, + "-c:a", + "aac", + "-b:a", + "128k", + "-movflags", + "+faststart", + ] + ) + args.extend(["-r", f"{frame_rate:.6g}", "-t", f"{duration_ms / 1000:.6f}"]) + return RenderPlan(tuple(args), tuple(inputs), duration_ms, output_format) diff --git a/app/projects/services/render_service.py b/app/projects/services/render_service.py new file mode 100644 index 0000000000000000000000000000000000000000..4693a3885656d632c6068eb04b5e166ca3bc4804 --- /dev/null +++ b/app/projects/services/render_service.py @@ -0,0 +1,273 @@ +from __future__ import annotations + +import hashlib +import json +from pathlib import Path + +from app.core.config import Settings +from app.core.logger import get_logger +from app.projects.editor_schemas import ( + AudioClip, + EditorDocument, + MediaClip, + ProjectRenderCreate, + ProjectRenderResponse, +) +from app.projects.errors import ( + ProjectRenderConflictError, + ProjectRenderInvalidError, + ProjectRenderLimitError, +) +from app.projects.models import ProjectRenderJob +from app.projects.repositories.editor_repository import ProjectEditorRepository +from app.projects.repositories.render_repository import ProjectRenderRepository +from app.projects.services.render_compiler import validate_renderable_document +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService +from app.services.cleanup import CleanupService + +logger = get_logger(__name__) + + +class ProjectRenderService: + def __init__( + self, + settings: Settings, + repository: ProjectRenderRepository, + editor_repository: ProjectEditorRepository, + assets: CanonicalAssetService, + cleanup: CleanupService, + audit: AuditService, + ) -> None: + self.settings = settings + self.repository = repository + self.editor_repository = editor_repository + self.assets = assets + self.cleanup = cleanup + self.audit = audit + + @staticmethod + def _document_from_job(job) -> EditorDocument: + return EditorDocument.model_validate(job.editor_state_json) + + async def create( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + payload: ProjectRenderCreate, + idempotency_key: str, + ) -> ProjectRenderResponse: + if len(idempotency_key.strip()) > 255 or not idempotency_key.strip(): + raise ProjectRenderInvalidError("An Idempotency-Key header is required for rendering.") + editor = await self.editor_repository.get(workspace_id, project_id, user_id=user_id) + if editor.revision != payload.editor_revision: + raise ProjectRenderInvalidError( + "The requested editor revision is not current or available." + ) + document = EditorDocument.model_validate(editor.state_json) + if document.project_id != project_id: + raise ProjectRenderInvalidError("Saved editor state belongs to another project.") + validate_renderable_document(document) + if document.duration_ms() > self.settings.render_max_duration_seconds * 1000: + raise ProjectRenderLimitError("The editor timeline exceeds the render duration limit.") + if len(document.timeline.tracks) > self.settings.render_max_tracks: + raise ProjectRenderLimitError("The editor timeline exceeds the render track limit.") + if ( + sum(len(track.clips) for track in document.timeline.tracks) + > self.settings.render_max_clips + ): + raise ProjectRenderLimitError("The editor timeline exceeds the render clip limit.") + if payload.width * payload.height > self.settings.max_resolution_pixels: + raise ProjectRenderLimitError( + "The requested render resolution exceeds the configured limit." + ) + await self._validate_assets(workspace_id, user_id, document) + render_settings = { + "format": payload.output_format, + "width": payload.width, + "height": payload.height, + "frameRate": payload.frame_rate, + "quality": payload.quality, + "preset": payload.preset, + } + fingerprint = hashlib.sha256( + json.dumps( + { + "revision": editor.revision, + "state": document.model_dump(by_alias=True), + "settings": render_settings, + }, + sort_keys=True, + separators=(",", ":"), + ).encode() + ).hexdigest() + existing = await self.repository.get_by_idempotency( + workspace_id, + project_id, + editor.revision, + idempotency_key.strip(), + user_id=user_id, + ) + if existing is not None: + if existing.request_fingerprint != fingerprint: + raise ProjectRenderConflictError( + "The idempotency key was already used for different render settings." + ) + return self._response(existing) + job, created = await self.repository.create( + workspace_id=workspace_id, + project_id=project_id, + user_id=user_id, + editor_revision=editor.revision, + document=document, + render_settings=render_settings, + request_fingerprint=fingerprint, + idempotency_key=idempotency_key.strip(), + max_attempts=max(1, self.settings.render_job_retry_limit + 1), + max_active_jobs=self.settings.render_max_active_jobs_per_project, + ) + if created: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.render_requested", + entity_type="project_render", + entity_id=job.id, + metadata={"project_id": project_id, "editor_revision": editor.revision}, + ) + logger.info( + "project render requested", + extra={ + "operation": "render.create", + "project_id": project_id, + "render_id": job.id, + "revision": editor.revision, + }, + ) + return self._response(job) + + async def get( + self, *, workspace_id: str, user_id: str, project_id: str, render_id: str + ) -> ProjectRenderResponse: + return self._response( + await self.repository.get(workspace_id, project_id, render_id, user_id=user_id) + ) + + async def list( + self, *, workspace_id: str, user_id: str, project_id: str + ) -> list[ProjectRenderResponse]: + return [ + self._response(item) + for item in await self.repository.list(workspace_id, project_id, user_id=user_id) + ] + + async def cancel( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + project_id: str, + render_id: str, + ) -> ProjectRenderResponse: + job, changed = await self.repository.request_cancel( + workspace_id, project_id, render_id, user_id=user_id + ) + # Queued jobs are cancelled immediately. Processing jobs are audited + # only after the worker has actually stopped FFmpeg. + if changed and job.status == "cancelled": + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type="project.render_cancelled", + entity_type="project_render", + entity_id=render_id, + metadata={"project_id": project_id}, + ) + logger.info( + "project render cancellation handled", + extra={ + "operation": "render.cancel", + "project_id": project_id, + "render_id": render_id, + "status": job.status, + "changed": changed, + }, + ) + return self._response(job) + + async def resolve_asset_paths(self, job: ProjectRenderJob) -> dict[str, tuple[Path, str]]: + document = EditorDocument.model_validate(job.editor_state_json) + paths: dict[str, tuple[Path, str]] = {} + for asset_id in document.asset_ids(): + asset = await self.assets.get_owned_by_id( + workspace_id=job.workspace_id, user_id=job.requested_by, asset_id=asset_id + ) + if asset.project_id != job.project_id: + raise CanonicalAssetNotFoundError("Asset is no longer attached to this project.") + path = self.cleanup.resolve_download(asset.request_id, asset.filename) + await self.assets.verify_file(asset, path) + paths[asset_id] = (path, asset.mime_type) + return paths + + async def _validate_assets( + self, workspace_id: str, user_id: str, document: EditorDocument + ) -> None: + total_bytes = 0 + owned = {} + for asset_id in document.asset_ids(): + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, user_id=user_id, asset_id=asset_id + ) + if asset.project_id != document.project_id: + raise CanonicalAssetNotFoundError("Asset is not attached to this project.") + total_bytes += asset.file_size + owned[asset_id] = asset + except CanonicalAssetNotFoundError as exc: + raise ProjectRenderInvalidError( + "A referenced asset is not accessible in this workspace." + ) from exc + if total_bytes > self.settings.render_max_input_bytes: + raise ProjectRenderLimitError( + "The referenced render inputs exceed the configured byte limit." + ) + for track in document.timeline.tracks: + for clip in track.clips: + if not isinstance(clip, (MediaClip, AudioClip)): + continue + mime = owned[clip.asset_id].mime_type + if isinstance(clip, AudioClip) and not mime.startswith("audio/"): + raise ProjectRenderInvalidError("An audio clip must reference an audio asset.") + if isinstance(clip, MediaClip) and not mime.startswith(f"{clip.media_type}/"): + raise ProjectRenderInvalidError( + f"A {clip.media_type} clip must reference a {clip.media_type} asset." + ) + + @staticmethod + def _response(job: ProjectRenderJob) -> ProjectRenderResponse: + return ProjectRenderResponse( + id=job.id, + project_id=job.project_id, + editor_revision=job.editor_revision, + status=job.status, + render_settings=dict(job.render_settings_json or {}), + output_asset_id=job.output_asset_id, + error_code=job.error_code, + error_message=job.error_message, + attempt_count=job.attempt_count, + created_at=job.created_at, + started_at=job.started_at, + completed_at=job.completed_at, + cancelled_at=job.cancelled_at, + updated_at=job.updated_at, + ) diff --git a/app/projects/workers/render_worker.py b/app/projects/workers/render_worker.py new file mode 100644 index 0000000000000000000000000000000000000000..a925380f5f3c7b07504487ca2311d16c3fc7a44a --- /dev/null +++ b/app/projects/workers/render_worker.py @@ -0,0 +1,249 @@ +from __future__ import annotations + +import asyncio +from pathlib import Path + +from app.core.config import Settings +from app.core.exceptions import MediaAPIError, ProcessingError +from app.core.logger import get_logger +from app.projects.repositories.render_repository import ProjectRenderRepository +from app.projects.services.render_compiler import compile_render +from app.projects.services.render_service import ProjectRenderService +from app.security.assets import CanonicalAssetNotFoundError +from app.security.audit import AuditService +from app.services.cleanup import CleanupService +from app.services.ffmpeg_service import FFmpegService + +logger = get_logger(__name__) + + +class ProjectRenderWorker: + """Small cooperative worker; durable state remains the source of truth.""" + + def __init__( + self, + settings: Settings, + repository: ProjectRenderRepository, + renders: ProjectRenderService, + ffmpeg: FFmpegService, + cleanup: CleanupService, + audit: AuditService, + ) -> None: + self.settings = settings + self.repository = repository + self.renders = renders + self.ffmpeg = ffmpeg + self.cleanup = cleanup + self.audit = audit + self._task: asyncio.Task[None] | None = None + self._stopping = asyncio.Event() + + async def start(self) -> None: + if not self.settings.render_worker_enabled or self._task is not None: + return + self._stopping.clear() + self._task = asyncio.create_task(self._run(), name="project-render-worker") + + async def stop(self) -> None: + self._stopping.set() + if self._task is not None: + self._task.cancel() + await asyncio.gather(self._task, return_exceptions=True) + self._task = None + + async def _run(self) -> None: + while not self._stopping.is_set(): + job = await self.repository.claim_next( + stale_after_seconds=self.settings.render_job_stale_after_seconds + ) + if job is None: + try: + await asyncio.wait_for( + self._stopping.wait(), self.settings.render_worker_interval_seconds + ) + except asyncio.TimeoutError: + pass + continue + try: + await self._execute(job) + except Exception: + logger.exception( + "project render worker execution failed", + extra={"render_id": job.id, "project_id": job.project_id}, + ) + + async def _execute(self, job) -> None: + workspace = await self.cleanup.create_workspace(job.id) + staging = ( + workspace.outputs + / f"render-{job.id}.staging.{job.render_settings_json.get('format', 'mp4')}" + ) + cancel_event = asyncio.Event() + render_task: asyncio.Task[None] | None = None + try: + await self.repository.ensure_project_active(job.id) + paths = await self.renders.resolve_asset_paths(job) + document = self.renders._document_from_job(job) + settings = job.render_settings_json + plan = compile_render( + document, + asset_paths=paths, + width=int(settings["width"]), + height=int(settings["height"]), + frame_rate=float(settings["frameRate"]), + output_format=str(settings["format"]), + quality=str(settings["quality"]), + preset=str(settings["preset"]), + ) + logger.info( + "project render started", + extra={ + "operation": "render.started", + "render_id": job.id, + "project_id": job.project_id, + "workspace_id": job.workspace_id, + "attempt": job.attempt_count, + }, + ) + render_task = asyncio.create_task( + self.ffmpeg.run( + [*plan.args, staging], + operation="project.render", + timeout=self.settings.render_job_timeout_seconds, + cancel_event=cancel_event, + ) + ) + while not render_task.done(): + latest = await self.repository.heartbeat(job.id) + if latest is not None and latest.status == "cancelling": + cancel_event.set() + break + await asyncio.sleep(1) + try: + await render_task + except ProcessingError as exc: + if cancel_event.is_set(): + await self.repository.cancel_worker(job.id) + await self._audit(job, "project.render_cancelled", {}) + return + failed = await self.repository.fail( + job.id, + code="RENDER_PROCESSING_FAILED", + message=str(exc), + retryable=True, + ) + if failed is not None and failed.status == "cancelled": + await self._audit(job, "project.render_cancelled", {}) + elif failed is not None and failed.status == "failed": + await self._audit( + job, + "project.render_failed", + {"error_code": "RENDER_PROCESSING_FAILED"}, + ) + return + latest = await self.repository.get_worker(job.id) + if latest is not None and latest.status == "cancelling": + await self.repository.cancel_worker(job.id) + await self._audit(job, "project.render_cancelled", {}) + return + filename = f"render-{job.id}.{plan.output_extension}" + published = await self.cleanup.publish_new(job.id, staging, filename) + if published is None: + # A prior attempt may have published before losing its worker + # lease or database connection. Reconcile that deterministic + # path instead of overwriting it or failing a safe retry. + published = self.cleanup.resolve_download(job.id, filename) + asset = await self.renders.assets.register_output( + workspace_id=job.workspace_id, + user_id=job.requested_by, + request_id=job.id, + path=published, + mime_type="video/webm" if plan.output_extension == "webm" else "video/mp4", + metadata={"project_render_id": job.id, "editor_revision": job.editor_revision}, + project_id=job.project_id, + ) + completed = await self.repository.complete(job.id, asset.id) + if completed.status == "cancelled": + await self.renders.assets.discard_output( + workspace_id=job.workspace_id, + asset_id=asset.id, + request_id=job.id, + filename=published.name, + ) + await self.cleanup.remove_request(job.id) + await self._audit(job, "project.render_cancelled", {}) + return + await self._audit(job, "project.render_completed", {"output_asset_id": asset.id}) + except asyncio.CancelledError: + cancel_event.set() + if render_task is not None: + await asyncio.gather(render_task, return_exceptions=True) + stopped = await self.repository.fail( + job.id, + code="RENDER_WORKER_STOPPED", + message="Render worker stopped before completion.", + retryable=True, + ) + if stopped is not None and stopped.status == "cancelled": + await self._audit(job, "project.render_cancelled", {}) + raise + except (MediaAPIError, CanonicalAssetNotFoundError, ValueError) as exc: + code = getattr(exc, "code", "RENDER_INVALID") + failed = await self.repository.fail( + job.id, code=code, message=str(exc), retryable=False + ) + if failed is not None and failed.status == "cancelled": + await self._audit(job, "project.render_cancelled", {}) + else: + await self._audit(job, "project.render_failed", {"error_code": code}) + except Exception: + logger.exception( + "project render failed unexpectedly", + extra={"render_id": job.id, "project_id": job.project_id}, + ) + failed = await self.repository.fail( + job.id, + code="RENDER_INTERNAL_ERROR", + message="Render processing failed.", + retryable=True, + ) + if failed is not None and failed.status == "cancelled": + await self._audit(job, "project.render_cancelled", {}) + elif failed is not None and failed.status == "failed": + await self._audit( + job, "project.render_failed", {"error_code": "RENDER_INTERNAL_ERROR"} + ) + finally: + await self._remove(staging) + await self.cleanup.remove_temporary_request(job.id) + await self.cleanup.complete(job.id) + + async def _audit(self, job, event_type: str, metadata: dict[str, object]) -> None: + await self.audit.record_event( + workspace_id=job.workspace_id, + user_id=job.requested_by, + api_key_id=None, + request_id=job.id, + event_type=event_type, + entity_type="project_render", + entity_id=job.id, + metadata={"project_id": job.project_id, **metadata}, + ) + logger.info( + "project render lifecycle event", + extra={ + "operation": event_type, + "render_id": job.id, + "project_id": job.project_id, + "workspace_id": job.workspace_id, + "error_code": metadata.get("error_code"), + "result": event_type.rsplit(".", 1)[-1], + }, + ) + + @staticmethod + async def _remove(path: Path) -> None: + try: + await asyncio.to_thread(path.unlink, missing_ok=True) + except OSError: + pass diff --git a/app/security/assets.py b/app/security/assets.py new file mode 100644 index 0000000000000000000000000000000000000000..8e126e959ee82d46620a73eb7a99f115698a4308 --- /dev/null +++ b/app/security/assets.py @@ -0,0 +1,170 @@ +from __future__ import annotations + +import asyncio +import hashlib +import hmac +from pathlib import Path +from typing import Any + +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError + +from app.security.database import SecurityDatabase +from app.security.models import CanonicalMediaAsset + + +class CanonicalAssetNotFoundError(Exception): + """A requested output was never issued to the caller's workspace.""" + + +class CanonicalAssetService: + """Persists workspace ownership for MediaRouter-produced files. + + The database stores an immutable, validated locator rather than a client + filesystem path. File resolution remains the responsibility of + ``CleanupService`` so every consumer receives the same traversal checks. + """ + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def register_output( + self, + *, + workspace_id: str, + user_id: str | None, + request_id: str, + path: Path, + mime_type: str, + metadata: dict[str, Any] | None = None, + project_id: str | None = None, + ) -> CanonicalMediaAsset: + if path.name != str(path.name) or not path.is_file(): + raise CanonicalAssetNotFoundError("Generated output is unavailable.") + digest = await asyncio.to_thread(self._sha256, path) + record = CanonicalMediaAsset( + workspace_id=workspace_id, + request_id=request_id, + filename=path.name, + mime_type=mime_type, + file_size=path.stat().st_size, + sha256=digest, + metadata_json=dict(metadata or {}), + created_by_user_id=user_id, + project_id=project_id, + ) + try: + async with self.database.session() as session: + session.add(record) + await session.commit() + await session.refresh(record) + return record + except IntegrityError: + async with self.database.session() as session: + existing = await session.scalar( + select(CanonicalMediaAsset).where( + CanonicalMediaAsset.request_id == request_id, + CanonicalMediaAsset.filename == path.name, + ) + ) + if existing is None: + raise + # Output IDs are globally unique. A second workspace must + # never be allowed to claim the same path after a race. + if ( + existing.workspace_id != workspace_id + or existing.project_id != project_id + or existing.mime_type != mime_type + ): + raise CanonicalAssetNotFoundError( + "Generated output is not owned by this workspace." + ) + await self.verify_file(existing, path) + return existing + + async def discard_output( + self, + *, + workspace_id: str, + asset_id: str, + request_id: str, + filename: str, + ) -> bool: + """Remove a just-created canonical output after cancellation wins. + + Immutable locator fields must all match so this internal compensation + cannot delete an unrelated asset selected only by an opaque ID. + """ + + async with self.database.session() as session: + record = await session.scalar( + select(CanonicalMediaAsset) + .where( + CanonicalMediaAsset.id == asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + CanonicalMediaAsset.request_id == request_id, + CanonicalMediaAsset.filename == filename, + ) + .with_for_update() + ) + if record is None: + return False + await session.delete(record) + await session.commit() + return True + + async def get_owned( + self, *, workspace_id: str, request_id: str, filename: str + ) -> CanonicalMediaAsset: + async with self.database.session() as session: + record = await session.scalar( + select(CanonicalMediaAsset).where( + CanonicalMediaAsset.workspace_id == workspace_id, + CanonicalMediaAsset.request_id == request_id, + CanonicalMediaAsset.filename == filename, + ) + ) + if record is None: + raise CanonicalAssetNotFoundError("Media asset was not found in this workspace.") + return record + + async def get_owned_by_id( + self, *, workspace_id: str, user_id: str, asset_id: str + ) -> CanonicalMediaAsset: + """Resolve a canonical asset reference without accepting a path. + + Generation (and future first-party services) receive only the opaque + canonical asset ID. The workspace predicate remains mandatory even + though the table is also protected by PostgreSQL RLS. + """ + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + record = await session.scalar( + select(CanonicalMediaAsset).where( + CanonicalMediaAsset.id == asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + ) + ) + if record is None: + raise CanonicalAssetNotFoundError("Media asset was not found in this workspace.") + return record + + async def verify_file(self, record: CanonicalMediaAsset, path: Path) -> None: + if not path.is_file() or path.name != record.filename: + raise CanonicalAssetNotFoundError("Media asset is no longer readable.") + stat = path.stat() + if stat.st_size != record.file_size: + raise CanonicalAssetNotFoundError("Media asset changed after it was registered.") + digest = await asyncio.to_thread(self._sha256, path) + if not hmac.compare_digest(digest, record.sha256): + raise CanonicalAssetNotFoundError("Media asset changed after it was registered.") + + @staticmethod + def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + while chunk := stream.read(1024 * 1024): + digest.update(chunk) + return digest.hexdigest() diff --git a/app/security/audit.py b/app/security/audit.py index 5037ce9a3f183dd22e6cca1405ddb861a3a7edbd..4228aaee306c779a9f11c0f50e1cdba13573de3a 100644 --- a/app/security/audit.py +++ b/app/security/audit.py @@ -1,9 +1,11 @@ from __future__ import annotations +from typing import Any + from sqlalchemy import select from app.security.database import SecurityDatabase -from app.security.models import AuditLog +from app.security.models import AuditEvent, AuditLog class AuditService: @@ -55,3 +57,56 @@ class AuditService: ) ).all() ) + + async def record_event( + self, + *, + workspace_id: str, + user_id: str, + event_type: str, + entity_type: str, + entity_id: str, + api_key_id: str | None = None, + request_id: str | None = None, + metadata: dict[str, Any] | None = None, + ) -> None: + """Persist a safe domain event through the existing audit boundary.""" + + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + session.add( + AuditEvent( + workspace_id=workspace_id, + actor_user_id=user_id, + api_key_id=api_key_id, + event_type=event_type[:100], + entity_type=entity_type[:64], + entity_id=entity_id, + request_id=request_id[:64] if request_id else None, + metadata_json=self._safe_metadata(metadata or {}), + ) + ) + await session.commit() + + @staticmethod + def _safe_metadata(metadata: dict[str, Any]) -> dict[str, object]: + """Keep audit metadata flat, bounded, and free of credential-like keys.""" + + blocked = ("token", "secret", "credential", "authorization", "password", "key") + result: dict[str, object] = {} + for raw_key, raw_value in list(metadata.items())[:32]: + key = str(raw_key)[:64] + if any(fragment in key.casefold() for fragment in blocked): + continue + if raw_value is None or isinstance(raw_value, (bool, int, float)): + result[key] = raw_value + elif isinstance(raw_value, str): + result[key] = raw_value[:256] + elif isinstance(raw_value, (list, tuple)): + result[key] = [ + value[:128] if isinstance(value, str) else value + for value in raw_value[:32] + if value is None or isinstance(value, (bool, int, float, str)) + ] + return result diff --git a/app/security/cli.py b/app/security/cli.py index 998a3accaa222e1c77d29e4295a2101c31d4ad24..fa4a047889e3c94af096a7b163a1ec1aedbc0c0d 100644 --- a/app/security/cli.py +++ b/app/security/cli.py @@ -8,12 +8,13 @@ from app.core.config import get_settings from app.security.database import SecurityDatabase from app.security.schemas import APIKeyCreate, APIKeyView from app.security.service import APIKeyService +from app.security.tenancy import TenantService async def _create(arguments: argparse.Namespace) -> None: settings = get_settings() - database = SecurityDatabase(settings.database_url) - service = APIKeyService(database, settings) + database = SecurityDatabase(settings.database_url, auto_migrate=settings.security_auto_migrate) + service = APIKeyService(database, settings, TenantService(database)) await database.initialize() try: record, secret = await service.create( diff --git a/app/security/context.py b/app/security/context.py index e0a878641d4b7501b3525543e5765bd5e1883938..01b63eaed02a3c626f4bcb332a90b01a59849ac6 100644 --- a/app/security/context.py +++ b/app/security/context.py @@ -18,9 +18,17 @@ class AuthContext: uploads_per_hour: int processing_bytes_per_day: int expires_at: datetime | None + # Resolved exclusively by the server from APIKeyPrincipal. These are not + # accepted from REST, MCP, SDK, or n8n request payloads. + workspace_id: str | None = None + user_id: str | None = None + membership_id: str | None = None + membership_role: str | None = None def allows(self, required_scope: str) -> bool: - return "admin" in self.scopes or required_scope in self.scopes + if required_scope in self.scopes or "admin" in self.scopes: + return True + return False auth_context: contextvars.ContextVar[AuthContext | None] = contextvars.ContextVar( diff --git a/app/security/database.py b/app/security/database.py index 1eb182e165a1f826100548bd694a9977b2d6bde5..412d19972e2c1ba304f0486362e7fba25689e4b1 100644 --- a/app/security/database.py +++ b/app/security/database.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import AsyncIterator from contextlib import asynccontextmanager -from sqlalchemy import event +from sqlalchemy import event, inspect, text from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, @@ -11,13 +11,91 @@ from sqlalchemy.ext.asyncio import ( create_async_engine, ) +# Importing the generation models here guarantees that local SQLite metadata +# and production schema checks see the full security-owned domain even when a +# caller constructs SecurityDatabase without first building the application +# container. +__import__("app.generation.models") +__import__("app.projects.models") +__import__("app.copilot.models") +__import__("app.templates.marketplace_models") +__import__("app.brand.models") +from app.core.database_url import normalize_async_database_url from app.security.models import Base +REQUIRED_SECURITY_TABLES = frozenset( + { + "api_keys", + "audit_logs", + "rate_limits", + "users", + "workspaces", + "workspace_memberships", + "api_key_principals", + "media_assets", + "media_asset_variants", + "audit_events", + "projects", + "project_generation_jobs", + "generation_requests", + "generation_jobs", + "generation_job_attempts", + "project_editor_states", + "project_render_jobs", + "copilot_runs", + "marketplace_templates", + "marketplace_template_versions", + "marketplace_template_applications", + "brand_kits", + "brand_kit_versions", + } +) +REQUIRED_POSTGRES_SECURITY_INDEXES = frozenset( + { + "uq_generation_job_provider_external", + "ix_projects_workspace", + "ix_projects_workspace_status", + "ix_projects_workspace_updated", + "ix_projects_created_by", + "ix_media_assets_workspace_project_created", + "ix_project_generation_jobs_workspace_project_created", + "uq_project_generation_job", + "uq_project_editor_state_project", + "ix_project_editor_states_workspace_project", + "uq_project_render_idempotency", + "ix_project_render_jobs_workspace_status", + "ix_project_render_jobs_project_created", + "ix_project_render_jobs_dispatch", + "ix_generation_requests_workspace_project_created", + "ix_generation_requests_workspace_surface_created", + "uq_copilot_run_workspace_idempotency", + "ix_copilot_runs_workspace_created", + "ix_copilot_runs_workspace_status", + "ix_copilot_runs_project_created", + "uq_marketplace_template_workspace_slug", + "ix_marketplace_templates_workspace_status", + "ix_marketplace_templates_discovery", + "uq_marketplace_template_version", + "ix_marketplace_template_versions_template", + "uq_marketplace_template_application_idempotency", + "ix_marketplace_template_applications_project", + "ix_marketplace_template_applications_template", + "ix_brand_kits_workspace_updated", + "ix_brand_kits_workspace_default", + "uq_brand_kits_one_default", + "uq_brand_kit_version", + "ix_brand_versions_kit_created", + } +) + class SecurityDatabase: """Owns the authentication database engine and short-lived async sessions.""" - def __init__(self, database_url: str) -> None: + def __init__(self, database_url: str, *, auto_migrate: bool = False) -> None: + database_url = normalize_async_database_url(database_url) + self.database_url = database_url + self.auto_migrate = auto_migrate or database_url.startswith("sqlite") self.engine: AsyncEngine = create_async_engine( database_url, pool_pre_ping=True, @@ -37,13 +115,93 @@ class SecurityDatabase: cursor.close() async def initialize(self) -> None: + if not self.auto_migrate: + return async with self.engine.begin() as connection: await connection.run_sync(Base.metadata.create_all) async def close(self) -> None: await self.engine.dispose() + @property + def is_postgres(self) -> bool: + return self.database_url.startswith(("postgresql", "postgres")) + + async def verify_execution_boundary(self, *, expected_role: str, enforce_rls: bool) -> None: + """Ensure tenancy administration never uses the public tenant role.""" + if not self.is_postgres or not enforce_rls: + return + if not expected_role.strip(): + raise RuntimeError( + "SECURITY_DATABASE_ROLE is required when SECURITY_ENFORCE_RLS is enabled." + ) + async with self.engine.connect() as connection: + row = ( + ( + await connection.execute( + text( + "select current_user as role, r.rolbypassrls as bypass_rls, r.rolsuper as superuser " + "from pg_roles r where r.rolname = current_user" + ) + ) + ) + .mappings() + .one_or_none() + ) + if row is None or row["role"] != expected_role.strip(): + raise RuntimeError("DATABASE_URL is not connected as SECURITY_DATABASE_ROLE.") + if not row["bypass_rls"] and not row["superuser"]: + raise RuntimeError( + "SECURITY_DATABASE_ROLE must be backend-only and able to administer forced-RLS tenancy tables." + ) + + async def schema_ready(self) -> bool: + return not await self.missing_schema_objects() + + async def missing_schema_objects(self) -> list[str]: + """Return missing production schema requirements, including safety indexes.""" + + async with self.engine.connect() as connection: + tables = await connection.run_sync(lambda sync: set(inspect(sync).get_table_names())) + if not self.is_postgres: + indexes: set[str] = set() + else: + rows = await connection.execute( + text("select indexname from pg_indexes " "where schemaname = current_schema()") + ) + indexes = set(rows.scalars().all()) + missing = sorted(REQUIRED_SECURITY_TABLES - tables) + if self.is_postgres: + missing.extend( + f"index:{name}" for name in sorted(REQUIRED_POSTGRES_SECURITY_INDEXES - indexes) + ) + return missing + + async def missing_tables(self) -> list[str]: + async with self.engine.connect() as connection: + tables = await connection.run_sync(lambda sync: set(inspect(sync).get_table_names())) + return sorted(REQUIRED_SECURITY_TABLES - tables) + @asynccontextmanager async def session(self) -> AsyncIterator[AsyncSession]: async with self.session_factory() as session: yield session + + @asynccontextmanager + async def tenant_session( + self, *, workspace_id: str, user_id: str + ) -> AsyncIterator[AsyncSession]: + """Explicit tenant context for RLS verification and scoped services.""" + if not workspace_id or not user_id: + raise RuntimeError("A security tenant session requires workspace and user IDs.") + async with self.session_factory() as session: + if not self.engine.url.drivername.startswith("sqlite"): + await session.execute( + text("select set_config('app.workspace_id', :workspace_id, true)"), + {"workspace_id": workspace_id}, + ) + await session.execute( + text("select set_config('app.user_id', :user_id, true)"), + {"user_id": user_id}, + ) + yield session diff --git a/app/security/middleware.py b/app/security/middleware.py index a1d450267a47fb8b5a3b49c6f73b00f400234888..830a5114d7a31099cd138cbe46caf12ae6f9f4c4 100644 --- a/app/security/middleware.py +++ b/app/security/middleware.py @@ -40,9 +40,7 @@ class APIKeyAuthenticationMiddleware(BaseHTTPMiddleware): self.audit = audit self.policy = ScopePolicy() - async def dispatch( - self, request: Request, call_next: RequestResponseEndpoint - ) -> Response: + async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response: if not getattr(request.state, "request_id", None): request.state.request_id = str(uuid4()) if not self.settings.auth_enabled or self.policy.is_public(request): @@ -63,10 +61,37 @@ class APIKeyAuthenticationMiddleware(BaseHTTPMiddleware): required_scope = await self.policy.required_scope(request) lease = await self.rate_limiter.acquire( context, - is_job=self.policy.is_job(required_scope), + is_job=self.policy.is_job(required_scope) + or request.url.path.startswith("/v1/projects/") + and "/renders" in request.url.path + and request.method == "POST", is_upload=self.policy.is_upload(request, required_scope), uploaded_bytes=bytes_uploaded, ) + if context.membership_role == "viewer" and required_scope not in { + "templates:read", + "operations:read", + "jobs:read", + "assets:read", + "mcp:read", + "system:read", + "social:accounts:read", + "social:posts:read", + "social:schedules:read", + "social:analytics:read", + "analytics:read", + "generation:providers:read", + "generation:requests:read", + "ai:read", + "copilot:read", + "projects:read", + "members:read", + "teams:read", + "projects:collaborate", + "comments:read", + "approvals:read", + }: + raise ForbiddenError self.api_keys.authorize(context, required_scope) await self._apply_social_rate_limit(request, context) await self.api_keys.mark_used(context) @@ -87,9 +112,7 @@ class APIKeyAuthenticationMiddleware(BaseHTTPMiddleware): ) except ForbiddenError: response_code = 403 - response = self._error( - 403, "Forbidden", "Missing required scope.", request - ) + response = self._error(403, "Forbidden", "Missing required scope.", request) bytes_downloaded = len(response.body) return response except RateLimitError as exc: @@ -121,10 +144,22 @@ class APIKeyAuthenticationMiddleware(BaseHTTPMiddleware): bytes_downloaded, ) - async def _apply_social_rate_limit( - self, request: Request, context: AuthContext - ) -> None: + async def _apply_social_rate_limit(self, request: Request, context: AuthContext) -> None: path = request.url.path + if path.startswith("/v1/analytics"): + category = "analytics_sync" if path.endswith(("/sync", "/cancel")) else "analytics_read" + limit = ( + max(1, self.settings.social_analytics_requests_per_minute // 6) + if category == "analytics_sync" + else self.settings.social_analytics_requests_per_minute + ) + await self.rate_limiter.acquire_category( + context, + category, + limit=limit, + window_seconds=60, + ) + return if not path.startswith("/v1/social"): return if "/analytics" in path: @@ -155,6 +190,20 @@ class APIKeyAuthenticationMiddleware(BaseHTTPMiddleware): limit=self.settings.social_schedule_requests_per_minute, window_seconds=60, ) + elif path.endswith("/reschedule"): + await self.rate_limiter.acquire_category( + context, + "social_schedule", + limit=self.settings.social_schedule_requests_per_minute, + window_seconds=60, + ) + elif path.endswith("/bulk"): + await self.rate_limiter.acquire_category( + context, + "social_bulk", + limit=max(1, self.settings.social_schedule_requests_per_minute // 4), + window_seconds=60, + ) @staticmethod def _bearer_token(header: str | None) -> str: diff --git a/app/security/migrations/0001_api_key_security.sql b/app/security/migrations/0001_api_key_security.sql index d897dd2d8520998bd1e3293b7c87fdb6082ec9b7..d654daac839ab327d449d4c57d5d9d530b4f89f2 100644 --- a/app/security/migrations/0001_api_key_security.sql +++ b/app/security/migrations/0001_api_key_security.sql @@ -1,5 +1,6 @@ --- MediaRouter security schema migration 0001 (SQLite). --- Production startup applies the equivalent SQLAlchemy metadata transactionally. +-- MediaRouter security schema migration 0001 (SQLite and PostgreSQL). +-- Use migration management in production; application metadata creation is +-- retained only for local development/backward compatibility. CREATE TABLE IF NOT EXISTS api_keys ( id VARCHAR(36) PRIMARY KEY, name VARCHAR(120) NOT NULL, @@ -9,10 +10,10 @@ CREATE TABLE IF NOT EXISTS api_keys ( status VARCHAR(16) NOT NULL, role VARCHAR(64), scopes JSON NOT NULL, - created_at DATETIME NOT NULL, - last_used_at DATETIME, - expires_at DATETIME, - grace_expires_at DATETIME, + created_at TIMESTAMPTZ NOT NULL, + last_used_at TIMESTAMPTZ, + expires_at TIMESTAMPTZ, + grace_expires_at TIMESTAMPTZ, created_by VARCHAR(120), notes TEXT, rotated_from_id VARCHAR(36) REFERENCES api_keys(id) ON DELETE SET NULL, @@ -39,7 +40,7 @@ CREATE TABLE IF NOT EXISTS audit_logs ( processing_time_ms INTEGER NOT NULL, bytes_uploaded BIGINT NOT NULL DEFAULT 0, bytes_downloaded BIGINT NOT NULL DEFAULT 0, - created_at DATETIME NOT NULL + created_at TIMESTAMPTZ NOT NULL ); CREATE INDEX IF NOT EXISTS ix_audit_logs_api_key_id ON audit_logs(api_key_id); CREATE INDEX IF NOT EXISTS ix_audit_logs_created_at ON audit_logs(created_at); @@ -48,10 +49,10 @@ CREATE INDEX IF NOT EXISTS ix_audit_logs_request_id ON audit_logs(request_id); CREATE TABLE IF NOT EXISTS rate_limits ( api_key_id VARCHAR(36) NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE, bucket_type VARCHAR(32) NOT NULL, - bucket_start DATETIME NOT NULL, + bucket_start TIMESTAMPTZ NOT NULL, count BIGINT NOT NULL DEFAULT 0, units BIGINT NOT NULL DEFAULT 0, - updated_at DATETIME NOT NULL, + updated_at TIMESTAMPTZ NOT NULL, PRIMARY KEY(api_key_id, bucket_type, bucket_start) ); CREATE INDEX IF NOT EXISTS ix_rate_limits_api_key_id ON rate_limits(api_key_id); diff --git a/app/security/migrations/0002_authoritative_tenancy_postgres.sql b/app/security/migrations/0002_authoritative_tenancy_postgres.sql new file mode 100644 index 0000000000000000000000000000000000000000..1e5848d167068562bab9ffca03c32cf6edfb6c03 --- /dev/null +++ b/app/security/migrations/0002_authoritative_tenancy_postgres.sql @@ -0,0 +1,172 @@ +-- Authoritative MediaRouter tenancy and canonical media ownership. +-- Apply after 0001 with the normal production migration process. +-- API keys remain credentials; api_key_principals binds them to a persisted +-- active membership. Do not grant these tables directly to browser clients. + +begin; +create extension if not exists pgcrypto; + +create table if not exists users ( + id text primary key default gen_random_uuid()::text, + subject text not null unique, + display_name text, + status text not null default 'active' check (status in ('active', 'disabled')), + metadata jsonb not null default '{}'::jsonb, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now() +); + +create table if not exists workspaces ( + id text primary key default gen_random_uuid()::text, + slug text not null unique, + name text not null, + status text not null default 'active' check (status in ('active', 'suspended', 'disabled')), + metadata jsonb not null default '{}'::jsonb, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now() +); + +create table if not exists workspace_memberships ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + role text not null default 'member' check (role in ('owner', 'admin', 'member', 'viewer', 'service')), + status text not null default 'active' check (status in ('active', 'disabled')), + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint uq_workspace_membership unique (workspace_id, user_id) +); +create index if not exists ix_workspace_memberships_user on workspace_memberships(user_id); +create index if not exists ix_workspace_memberships_workspace on workspace_memberships(workspace_id); + +create table if not exists api_key_principals ( + id text primary key default gen_random_uuid()::text, + api_key_id text not null unique references api_keys(id) on delete cascade, + workspace_id text not null references workspaces(id) on delete cascade, + user_id text not null references users(id) on delete cascade, + membership_id text not null references workspace_memberships(id) on delete restrict, + status text not null default 'active' check (status in ('active', 'disabled')), + created_at timestamptz not null default now(), + updated_at timestamptz not null default now() +); +create index if not exists ix_api_key_principals_workspace on api_key_principals(workspace_id); +create index if not exists ix_api_key_principals_user on api_key_principals(user_id); + +create table if not exists media_assets ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + request_id text not null, + filename text not null, + mime_type text not null, + file_size bigint not null check (file_size >= 0), + sha256 char(64) not null check (sha256 ~ '^[0-9a-f]{64}$'), + metadata jsonb not null default '{}'::jsonb, + created_by_user_id text references users(id) on delete set null, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint uq_media_asset_output unique (request_id, filename) +); +create index if not exists ix_media_assets_workspace_created on media_assets(workspace_id, created_at desc); +create index if not exists ix_media_assets_workspace_request on media_assets(workspace_id, request_id); + +create table if not exists media_asset_variants ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + asset_id text not null references media_assets(id) on delete cascade, + request_id text not null, + filename text not null, + mime_type text not null, + file_size bigint not null check (file_size >= 0), + sha256 char(64) not null check (sha256 ~ '^[0-9a-f]{64}$'), + metadata jsonb not null default '{}'::jsonb, + created_at timestamptz not null default now(), + constraint uq_media_asset_variant_output unique (asset_id, request_id, filename) +); +create index if not exists ix_media_asset_variants_workspace_asset on media_asset_variants(workspace_id, asset_id); + +create or replace function mediarouter_tenant_touch_updated_at() +returns trigger language plpgsql as $$ +begin + new.updated_at = now(); + return new; +end; +$$; + +do $$ +declare table_name text; +begin + foreach table_name in array array['users','workspaces','workspace_memberships','api_key_principals','media_assets'] loop + execute format('drop trigger if exists mediarouter_tenant_touch_updated_at on %I', table_name); + execute format('create trigger mediarouter_tenant_touch_updated_at before update on %I for each row execute function mediarouter_tenant_touch_updated_at()', table_name); + end loop; +end $$; + +-- A tenant API connection must set both values transaction-locally. Backend +-- authentication/administration and service workers use separate credentials, +-- never an exposed browser or SDK connection. +alter table users enable row level security; +alter table users force row level security; +drop policy if exists users_self on users; +create policy users_self on users + using (id = current_setting('app.user_id', true)) + with check (id = current_setting('app.user_id', true)); + +alter table workspaces enable row level security; +alter table workspaces force row level security; +drop policy if exists workspace_membership_read on workspaces; +create policy workspace_membership_read on workspaces + using (id = current_setting('app.workspace_id', true) and exists ( + select 1 from workspace_memberships m + where m.workspace_id = workspaces.id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + )) + with check (false); + +alter table workspace_memberships enable row level security; +alter table workspace_memberships force row level security; +drop policy if exists workspace_membership_self on workspace_memberships; +create policy workspace_membership_self on workspace_memberships + using (workspace_id = current_setting('app.workspace_id', true) + and user_id = current_setting('app.user_id', true)) + with check (false); + +alter table api_key_principals enable row level security; +alter table api_key_principals force row level security; +drop policy if exists api_key_principal_self on api_key_principals; +create policy api_key_principal_self on api_key_principals + using (workspace_id = current_setting('app.workspace_id', true) + and user_id = current_setting('app.user_id', true)) + with check (false); + +do $$ +declare table_name text; +begin + foreach table_name in array array['media_assets','media_asset_variants'] loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + execute format('drop policy if exists canonical_asset_workspace_isolation on %I', table_name); + execute format( + 'create policy canonical_asset_workspace_isolation on %I using (workspace_id = current_setting(''app.workspace_id'', true)) with check (workspace_id = current_setting(''app.workspace_id'', true))', + table_name + ); + end loop; +end $$; + +create or replace function mediarouter_assert_canonical_variant_workspace() +returns trigger language plpgsql as $$ +declare asset_workspace text; +begin + select workspace_id into asset_workspace from media_assets where id = new.asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'media variant must belong to its source asset workspace' using errcode = '23503'; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_canonical_variant_workspace on media_asset_variants; +create trigger mediarouter_canonical_variant_workspace +before insert or update of workspace_id, asset_id on media_asset_variants +for each row execute function mediarouter_assert_canonical_variant_workspace(); + +commit; diff --git a/app/security/migrations/0003_generation_domain_postgres.sql b/app/security/migrations/0003_generation_domain_postgres.sql new file mode 100644 index 0000000000000000000000000000000000000000..6e16f9515c700032b90e13aa35c9000db033ce82 --- /dev/null +++ b/app/security/migrations/0003_generation_domain_postgres.sql @@ -0,0 +1,195 @@ +-- Provider-neutral generation foundation. +-- +-- Apply after 0002_authoritative_tenancy_postgres.sql. This migration adds +-- durable generation intent and execution records only; it does not configure +-- WAN, FLUX, external worker endpoints, or any credentials. + +begin; + +create table if not exists generation_requests ( + id text primary key default gen_random_uuid()::text, + workspace_id text not null references workspaces(id) on delete restrict, + created_by_user_id text not null references users(id) on delete restrict, + provider text not null check (provider ~ '^[a-z][a-z0-9_-]{0,63}$'), + model_id text not null, + modality text not null check (modality in ('image', 'video')), + input_asset_id text references media_assets(id) on delete restrict, + spec jsonb not null default '{}'::jsonb, + request_fingerprint char(64) not null check (request_fingerprint ~ '^[0-9a-f]{64}$'), + idempotency_key text not null, + status text not null default 'queued' check (status in ( + 'queued', 'submitting', 'running', 'retrying', 'succeeded', 'failed', + 'cancel_requested', 'cancelled' + )), + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + completed_at timestamptz, + constraint uq_generation_request_workspace_idempotency unique (workspace_id, idempotency_key) +); +create index if not exists ix_generation_requests_workspace_created + on generation_requests(workspace_id, created_at desc); +create index if not exists ix_generation_requests_workspace_status + on generation_requests(workspace_id, status); + +create table if not exists generation_jobs ( + id text primary key default gen_random_uuid()::text, + generation_request_id text not null unique references generation_requests(id) on delete cascade, + workspace_id text not null references workspaces(id) on delete restrict, + provider text not null check (provider ~ '^[a-z][a-z0-9_-]{0,63}$'), + status text not null default 'queued' check (status in ( + 'queued', 'submitting', 'running', 'retrying', 'succeeded', 'failed', + 'cancel_requested', 'cancelled' + )), + attempt_count integer not null default 0 check (attempt_count >= 0), + max_attempts integer not null default 3 check (max_attempts >= 0), + next_attempt_at timestamptz, + external_job_id text, + provider_metadata jsonb not null default '{}'::jsonb, + output_asset_id text references media_assets(id) on delete restrict, + error_code text, + error_message text, + created_at timestamptz not null default now(), + started_at timestamptz, + completed_at timestamptz, + updated_at timestamptz not null default now() +); +create index if not exists ix_generation_jobs_workspace_status + on generation_jobs(workspace_id, status); +create index if not exists ix_generation_jobs_next_attempt + on generation_jobs(status, next_attempt_at); +create index if not exists ix_generation_jobs_external + on generation_jobs(provider, external_job_id) where external_job_id is not null; + +create table if not exists generation_job_attempts ( + id text primary key default gen_random_uuid()::text, + generation_job_id text not null references generation_jobs(id) on delete cascade, + attempt_number integer not null check (attempt_number > 0), + status text not null, + error_code text, + error_message text, + provider_request_id text, + started_at timestamptz not null default now(), + completed_at timestamptz, + constraint uq_generation_job_attempt_number unique (generation_job_id, attempt_number) +); +create index if not exists ix_generation_job_attempts_job + on generation_job_attempts(generation_job_id, attempt_number); + +-- Keep denormalised workspace IDs and canonical input/output references +-- internally consistent even when a trusted service role writes the rows. +create or replace function mediarouter_assert_generation_request_asset_workspace() +returns trigger language plpgsql as $$ +declare asset_workspace text; +begin + if new.input_asset_id is not null then + select workspace_id into asset_workspace from media_assets where id = new.input_asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'generation input asset must belong to its request workspace' using errcode = '23503'; + end if; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_generation_request_asset_workspace on generation_requests; +create trigger mediarouter_generation_request_asset_workspace +before insert or update of workspace_id, input_asset_id on generation_requests +for each row execute function mediarouter_assert_generation_request_asset_workspace(); + +create or replace function mediarouter_assert_generation_job_workspace() +returns trigger language plpgsql as $$ +declare request_workspace text; +declare asset_workspace text; +begin + select workspace_id into request_workspace from generation_requests where id = new.generation_request_id; + if request_workspace is null or request_workspace is distinct from new.workspace_id then + raise exception 'generation job must belong to its request workspace' using errcode = '23503'; + end if; + if new.output_asset_id is not null then + select workspace_id into asset_workspace from media_assets where id = new.output_asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'generation output asset must belong to its job workspace' using errcode = '23503'; + end if; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_generation_job_workspace on generation_jobs; +create trigger mediarouter_generation_job_workspace +before insert or update of workspace_id, generation_request_id, output_asset_id on generation_jobs +for each row execute function mediarouter_assert_generation_job_workspace(); + +do $$ +declare table_name text; +begin + foreach table_name in array array['generation_requests', 'generation_jobs'] loop + execute format('drop trigger if exists mediarouter_tenant_touch_updated_at on %I', table_name); + execute format('create trigger mediarouter_tenant_touch_updated_at before update on %I for each row execute function mediarouter_tenant_touch_updated_at()', table_name); + end loop; +end $$; + +-- API paths set app.workspace_id and app.user_id through SecurityDatabase's +-- tenant session. The security/service role is backend-only; these policies +-- remain defence in depth and independently verifiable with a tenant role. +alter table generation_requests enable row level security; +alter table generation_requests force row level security; +drop policy if exists generation_requests_workspace_isolation on generation_requests; +create policy generation_requests_workspace_isolation on generation_requests + using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = generation_requests.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) + ) + with check ( + workspace_id = current_setting('app.workspace_id', true) + and created_by_user_id = current_setting('app.user_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = generation_requests.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) + ); + +alter table generation_jobs enable row level security; +alter table generation_jobs force row level security; +drop policy if exists generation_jobs_workspace_isolation on generation_jobs; +create policy generation_jobs_workspace_isolation on generation_jobs + using ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = generation_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) + ) + with check ( + workspace_id = current_setting('app.workspace_id', true) + and exists ( + select 1 from workspace_memberships m + where m.workspace_id = generation_jobs.workspace_id + and m.user_id = current_setting('app.user_id', true) + and m.status = 'active' + ) + ); + +alter table generation_job_attempts enable row level security; +alter table generation_job_attempts force row level security; +drop policy if exists generation_job_attempts_workspace_isolation on generation_job_attempts; +create policy generation_job_attempts_workspace_isolation on generation_job_attempts + using (exists ( + select 1 from generation_jobs j + where j.id = generation_job_attempts.generation_job_id + and j.workspace_id = current_setting('app.workspace_id', true) + )) + with check (exists ( + select 1 from generation_jobs j + where j.id = generation_job_attempts.generation_job_id + and j.workspace_id = current_setting('app.workspace_id', true) + )); + +commit; diff --git a/app/security/migrations/0004_generation_provider_runtime_postgres.sql b/app/security/migrations/0004_generation_provider_runtime_postgres.sql new file mode 100644 index 0000000000000000000000000000000000000000..889c11a38af77fa067ca06503c3b6d565effdfed --- /dev/null +++ b/app/security/migrations/0004_generation_provider_runtime_postgres.sql @@ -0,0 +1,30 @@ +-- Provider-runtime safety constraints for the generation foundation. +-- +-- Apply after 0003_generation_domain_postgres.sql. This migration does not +-- register a generation provider, add a worker URL, or introduce credentials. + +begin; + +-- An opaque remote worker job may be bound to exactly one logical +-- MediaRouter job for a provider. PostgreSQL permits multiple NULL values, so +-- queued local jobs remain unaffected until a trusted dispatcher binds them. +do $$ +begin + if exists ( + select 1 + from generation_jobs + where external_job_id is not null + group by provider, external_job_id + having count(*) > 1 + ) then + raise exception + 'cannot add provider worker-job uniqueness: duplicate generation_jobs external IDs exist' + using errcode = '23505'; + end if; +end $$; + +create unique index if not exists uq_generation_job_provider_external + on generation_jobs(provider, external_job_id) + where external_job_id is not null; + +commit; diff --git a/app/security/models.py b/app/security/models.py index 93c62d5e6d68d743a3bb539a89a879d377ef3e5c..5914fc6f484f3d478afd77721a4c909d56656a4e 100644 --- a/app/security/models.py +++ b/app/security/models.py @@ -3,7 +3,17 @@ from __future__ import annotations from datetime import datetime, timezone from uuid import uuid4 -from sqlalchemy import BigInteger, DateTime, ForeignKey, Index, Integer, JSON, String, Text +from sqlalchemy import ( + JSON, + BigInteger, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + UniqueConstraint, +) from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column @@ -15,6 +25,83 @@ class Base(DeclarativeBase): pass +class User(Base): + """An authoritative MediaRouter actor, independent from API credentials.""" + + __tablename__ = "users" + __table_args__ = ( + UniqueConstraint("subject", name="uq_users_subject"), + Index("ix_users_status", "status"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + # `subject` is an immutable backend identity (for example an IdP subject), + # never an API-key secret or a user-supplied workspace selector. + subject: Mapped[str] = mapped_column(String(255), nullable=False) + display_name: Mapped[str | None] = mapped_column(String(255)) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="active") + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class Workspace(Base): + """A first-class tenant. API keys must be bound through a membership.""" + + __tablename__ = "workspaces" + __table_args__ = ( + UniqueConstraint("slug", name="uq_workspaces_slug"), + Index("ix_workspaces_status", "status"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + slug: Mapped[str] = mapped_column(String(120), nullable=False) + name: Mapped[str] = mapped_column(String(255), nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="active") + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class WorkspaceMembership(Base): + """Authorizes a user to act inside one workspace.""" + + __tablename__ = "workspace_memberships" + __table_args__ = ( + UniqueConstraint("workspace_id", "user_id", name="uq_workspace_membership"), + Index("ix_workspace_memberships_user", "user_id"), + Index("ix_workspace_memberships_workspace", "workspace_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False + ) + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False + ) + role: Mapped[str] = mapped_column(String(32), nullable=False, default="member") + status: Mapped[str] = mapped_column(String(32), nullable=False, default="active") + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + class APIKey(Base): __tablename__ = "api_keys" __table_args__ = ( @@ -51,6 +138,114 @@ class APIKey(Base): ) +class APIKeyPrincipal(Base): + """Server-side API-key-to-membership binding. + + Keeping this association outside the opaque API-key record prevents a + credential identifier from accidentally becoming a tenant identifier. + """ + + __tablename__ = "api_key_principals" + __table_args__ = ( + UniqueConstraint("api_key_id", name="uq_api_key_principal_key"), + Index("ix_api_key_principals_workspace", "workspace_id"), + Index("ix_api_key_principals_user", "user_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + api_key_id: Mapped[str] = mapped_column( + String(36), ForeignKey("api_keys.id", ondelete="CASCADE"), nullable=False + ) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="CASCADE"), nullable=False + ) + user_id: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False + ) + membership_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspace_memberships.id", ondelete="RESTRICT"), nullable=False + ) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="active") + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class CanonicalMediaAsset(Base): + """Workspace-owned output locator issued only by the MediaRouter pipeline.""" + + __tablename__ = "media_assets" + __table_args__ = ( + UniqueConstraint("request_id", "filename", name="uq_media_asset_output"), + Index("ix_media_assets_workspace_created", "workspace_id", "created_at"), + Index("ix_media_assets_workspace_request", "workspace_id", "request_id"), + Index( + "ix_media_assets_workspace_project_created", "workspace_id", "project_id", "created_at" + ), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + # Assets may remain workspace-level. Once attached, one canonical asset + # has one project parent; PostgreSQL also verifies matching workspaces. + project_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT") + ) + request_id: Mapped[str] = mapped_column(String(36), nullable=False) + filename: Mapped[str] = mapped_column(String(255), nullable=False) + mime_type: Mapped[str] = mapped_column(String(255), nullable=False) + file_size: Mapped[int] = mapped_column(BigInteger, nullable=False) + sha256: Mapped[str] = mapped_column(String(64), nullable=False) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_by_user_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("users.id", ondelete="SET NULL") + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class CanonicalMediaVariant(Base): + """An immutable derivative of a canonical asset, owned by the same tenant.""" + + __tablename__ = "media_asset_variants" + __table_args__ = ( + UniqueConstraint( + "asset_id", "request_id", "filename", name="uq_media_asset_variant_output" + ), + Index("ix_media_asset_variants_workspace_asset", "workspace_id", "asset_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + asset_id: Mapped[str] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="CASCADE"), nullable=False + ) + request_id: Mapped[str] = mapped_column(String(36), nullable=False) + filename: Mapped[str] = mapped_column(String(255), nullable=False) + mime_type: Mapped[str] = mapped_column(String(255), nullable=False) + file_size: Mapped[int] = mapped_column(BigInteger, nullable=False) + sha256: Mapped[str] = mapped_column(String(64), nullable=False) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + class AuditLog(Base): __tablename__ = "audit_logs" __table_args__ = ( @@ -78,6 +273,38 @@ class AuditLog(Base): ) +class AuditEvent(Base): + """Safe workspace-scoped domain event recorded by the shared AuditService.""" + + __tablename__ = "audit_events" + __table_args__ = ( + Index("ix_audit_events_workspace_created", "workspace_id", "created_at"), + Index("ix_audit_events_type", "event_type"), + Index("ix_audit_events_entity", "entity_type", "entity_id"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + actor_user_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("users.id", ondelete="SET NULL") + ) + api_key_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("api_keys.id", ondelete="SET NULL") + ) + event_type: Mapped[str] = mapped_column(String(100), nullable=False) + entity_type: Mapped[str] = mapped_column(String(64), nullable=False) + entity_id: Mapped[str] = mapped_column(String(36), nullable=False) + request_id: Mapped[str | None] = mapped_column(String(64)) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + class RateLimit(Base): __tablename__ = "rate_limits" __table_args__ = ( diff --git a/app/security/policy.py b/app/security/policy.py index c48e1bb628e4c426e5f2e4ed10f7cae808686b7f..c756a3300cb44a9e8f2275a794efc83e1e7a1294 100644 --- a/app/security/policy.py +++ b/app/security/policy.py @@ -34,12 +34,39 @@ class ScopePolicy: return None if path.startswith("/v1/social"): return self._social_scope(path, method) + if path.startswith("/v1/analytics"): + if path == "/v1/analytics/sync" and method == "POST": + return "analytics:sync" + if path.endswith("/cancel") and method == "POST": + return "analytics:sync" + return "analytics:read" + if path.startswith("/v1/generation"): + return self._generation_scope(path, method) + if path.startswith("/v1/ai"): + return self._ai_scope(path, method) + if path.startswith("/v1/copilot"): + return self._copilot_scope(path, method) + if path.startswith("/v1/projects"): + return self._project_scope(path, method) if path.startswith("/mcp"): return await self._mcp_scope(request) if path.startswith("/v1/api-keys") or path.startswith("/v1/audit-logs"): return "admin" - if path == "/v1/templates" or path == "/v1/templates/categories" or ( - path.startswith("/v1/templates/") and path != "/v1/templates/run" + if path.startswith("/v1/templates/catalog"): + if method == "GET": + return "templates:read" + if path.endswith(("/apply", "/instantiate")): + return "templates:apply" + if method == "POST": + return "templates:create" + if method == "PATCH": + return "templates:update" + if method == "DELETE": + return "templates:delete" + if ( + path == "/v1/templates" + or path == "/v1/templates/categories" + or (path.startswith("/v1/templates/") and path != "/v1/templates/run") ): return "templates:read" if method == "GET" else "admin" if path == "/v1/templates/run": @@ -58,9 +85,10 @@ class ScopePolicy: return "assets:write" if path.startswith("/v1/media/") and method == "GET": return "operations:read" - if path.startswith( - ("/v1/video", "/v1/audio", "/v1/image", "/v1/whisper", "/v1/ytdlp") - ) or path == "/v1/probe": + if ( + path.startswith(("/v1/video", "/v1/audio", "/v1/image", "/v1/whisper", "/v1/ytdlp")) + or path == "/v1/probe" + ): return "operations:execute" if method == "GET": return "system:read" @@ -74,6 +102,16 @@ class ScopePolicy: return "social:analytics:read" if "/jobs" in path: return "social:posts:read" + if path.endswith("/calendar"): + return "social:schedules:read" + if path.endswith("/publishing-context"): + return "social:posts:read" + if path.endswith("/queue"): + return "social:posts:read" + if path.endswith("/bulk"): + return "social:schedules:write" + if "/drafts" in path: + return "social:posts:read" if method == "GET" else "social:posts:write" if "/accounts" in path: return "social:accounts:read" if method == "GET" else "social:accounts:write" if "/posts" in path: @@ -83,11 +121,75 @@ class ScopePolicy: return "social:posts:publish" if path.endswith("/schedule"): return "social:schedules:write" + if path.endswith("/reschedule"): + return "social:schedules:write" if path.endswith("/cancel"): return "social:posts:write" return "social:posts:write" return "social:accounts:read" + @staticmethod + def _generation_scope(path: str, method: str) -> str: + if "/providers" in path: + return "generation:providers:read" + if path.endswith("/cancel"): + return "generation:jobs:cancel" + if "/jobs/" in path or method == "GET": + return "generation:requests:read" + return "generation:requests:create" + + @staticmethod + def _ai_scope(path: str, method: str) -> str: + if path.endswith("/capabilities") or method == "GET": + return "ai:read" + return "ai:generate" + + @staticmethod + def _copilot_scope(path: str, method: str) -> str: + if method == "GET" or path.endswith("/capabilities"): + return "copilot:read" + return "copilot:execute" + + @staticmethod + def _project_scope(path: str, method: str) -> str: + if "/editor" in path: + return "projects:read" if method == "GET" else "projects:update" + if "/renders" in path: + return "projects:read" if method == "GET" else "projects:update" + if "/assets" in path or "/jobs" in path: + return "projects:read" if method == "GET" else "projects:update" + if "/workspace/teams" in path: + return { + "GET": "teams:read", + "POST": "teams:create", + "PATCH": "teams:update", + "DELETE": "teams:delete", + }.get(method, "admin") + if "/workspace/invitations" in path: + return "members:invite" + if "/workspace/members" in path: + return { + "DELETE": "members:remove", + "PATCH": "members:update", + "GET": "members:read" + }.get(method, "admin") + if "/workspace/workflows" in path or "/workspace/requests" in path: + if method == "POST": + if "approve" in path or "reject" in path: + return "approvals:review" + if "comments" in path: + return "comments:create" + return "approvals:create" + return "approvals:read" + if "/collaborators" in path: + return "projects:collaborate" + return { + "GET": "projects:read", + "POST": "projects:create", + "PATCH": "projects:update", + "DELETE": "projects:delete", + }.get(method, "admin") + @staticmethod async def _mcp_scope(request: Request) -> str: if request.method != "POST": @@ -98,9 +200,7 @@ class ScopePolicy: return "mcp:read" messages = payload if isinstance(payload, list) else [payload] methods = { - str(message.get("method", "")) - for message in messages - if isinstance(message, dict) + str(message.get("method", "")) for message in messages if isinstance(message, dict) } return "mcp:execute" if "tools/call" in methods else "mcp:read" @@ -112,6 +212,10 @@ class ScopePolicy: "jobs:create", "mcp:execute", "social:posts:publish", + "generation:requests:create", + "analytics:sync", + "ai:generate", + "copilot:execute", } @staticmethod diff --git a/app/security/schemas.py b/app/security/schemas.py index a9fe86dc0a1dc9b8305c5e3cb38b1865dc92ca0f..a68d709c59652ed37438003e344ddffc76b30c9a 100644 --- a/app/security/schemas.py +++ b/app/security/schemas.py @@ -131,6 +131,9 @@ class AuthContextView(BaseModel): role: str | None scopes: list[str] expires_at: datetime | None + workspace_id: str | None = None + user_id: str | None = None + membership_role: str | None = None class AuditLogView(BaseModel): diff --git a/app/security/scopes.py b/app/security/scopes.py index 890ddc0ad73d9bf99fc9e562b32c301692346cd1..c06da3468893afa242dc4f33ab24a8e6fbb1f566 100644 --- a/app/security/scopes.py +++ b/app/security/scopes.py @@ -6,6 +6,11 @@ ALL_SCOPES = frozenset( { "templates:read", "templates:run", + "templates:create", + "templates:update", + "templates:delete", + "templates:apply", + "templates:publish", "operations:read", "operations:execute", "jobs:read", @@ -25,6 +30,43 @@ ALL_SCOPES = frozenset( "social:schedules:read", "social:schedules:write", "social:analytics:read", + "analytics:read", + "analytics:sync", + "analytics:export", + "generation:providers:read", + "generation:requests:read", + "generation:requests:create", + "generation:jobs:cancel", + "ai:read", + "ai:generate", + "ai:transform", + "ai:analyze", + "ai:create", + "copilot:read", + "copilot:execute", + "projects:read", + "projects:create", + "projects:update", + "projects:delete", + # Frontend Architecture v2 uses this aggregate solely for capability + # rendering. Enforcement remains method-specific in ScopePolicy. + "projects:write", + "members:read", + "members:invite", + "members:update", + "members:remove", + "teams:read", + "teams:create", + "teams:update", + "teams:delete", + "projects:share", + "projects:collaborate", + "comments:read", + "comments:create", + "approvals:read", + "approvals:create", + "approvals:review", + "approvals:manage", "admin", } ) @@ -35,6 +77,11 @@ DEFAULT_ROLE_SCOPES: dict[str, frozenset[str]] = { { "templates:read", "templates:run", + "templates:create", + "templates:update", + "templates:delete", + "templates:apply", + "templates:publish", "operations:read", "operations:execute", "jobs:read", @@ -53,12 +100,52 @@ DEFAULT_ROLE_SCOPES: dict[str, frozenset[str]] = { "social:schedules:read", "social:schedules:write", "social:analytics:read", + "analytics:read", + "analytics:sync", + "analytics:export", + "generation:providers:read", + "generation:requests:read", + "generation:requests:create", + "generation:jobs:cancel", + "ai:read", + "ai:generate", + "ai:transform", + "ai:analyze", + "ai:create", + "copilot:read", + "copilot:execute", + "projects:read", + "projects:create", + "projects:update", + "projects:delete", + "projects:write", + "members:read", + "members:invite", + "members:update", + "members:remove", + "teams:read", + "teams:create", + "teams:update", + "teams:delete", + "projects:share", + "projects:collaborate", + "comments:read", + "comments:create", + "approvals:read", + "approvals:create", + "approvals:review", + "approvals:manage", } ), "operator": frozenset( { "templates:read", "templates:run", + "templates:create", + "templates:update", + "templates:delete", + "templates:apply", + "templates:publish", "operations:read", "operations:execute", "jobs:read", @@ -76,6 +163,32 @@ DEFAULT_ROLE_SCOPES: dict[str, frozenset[str]] = { "social:schedules:read", "social:schedules:write", "social:analytics:read", + "analytics:read", + "analytics:sync", + "generation:providers:read", + "generation:requests:read", + "generation:requests:create", + "generation:jobs:cancel", + "ai:read", + "ai:generate", + "ai:transform", + "ai:analyze", + "ai:create", + "copilot:read", + "copilot:execute", + "projects:read", + "projects:create", + "projects:update", + "projects:delete", + "projects:write", + "members:read", + "teams:read", + "projects:share", + "projects:collaborate", + "comments:read", + "comments:create", + "approvals:read", + "approvals:review", } ), "viewer": frozenset( @@ -90,6 +203,17 @@ DEFAULT_ROLE_SCOPES: dict[str, frozenset[str]] = { "social:posts:read", "social:schedules:read", "social:analytics:read", + "analytics:read", + "generation:providers:read", + "generation:requests:read", + "ai:read", + "copilot:read", + "projects:read", + "members:read", + "teams:read", + "projects:collaborate", + "comments:read", + "approvals:read", } ), } @@ -117,6 +241,11 @@ def effective_scopes( scopes = set(scope.strip().lower() for scope in explicit_scopes) if role: scopes.update(configured_roles(overrides).get(role.strip().lower(), ())) + project_mutations = {"projects:create", "projects:update", "projects:delete"} + if "projects:write" in scopes: + scopes.update(project_mutations) + if project_mutations.issubset(scopes): + scopes.add("projects:write") if "admin" in scopes: return ALL_SCOPES return frozenset(scopes) diff --git a/app/security/service.py b/app/security/service.py index 8109b2a9557ce4539104c7e34dfd895448969a05..5c4318c4423096e67ed8dde1bcf0574e10c5e8fe 100644 --- a/app/security/service.py +++ b/app/security/service.py @@ -23,6 +23,7 @@ from app.security.errors import ( from app.security.models import APIKey from app.security.schemas import APIKeyCreate, APIKeyPatch from app.security.scopes import configured_roles, effective_scopes +from app.security.tenancy import TenantService KEY_PATTERN = re.compile(r"^mp_([a-z][a-z0-9]{1,15})_([A-Za-z0-9_-]{43,})$") HASH_PATTERN = re.compile(r"^[a-f0-9]{64}$") @@ -35,7 +36,11 @@ def utcnow() -> datetime: def aware(value: datetime | None) -> datetime | None: if value is None: return None - return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc) + return ( + value.replace(tzinfo=timezone.utc) + if value.tzinfo is None + else value.astimezone(timezone.utc) + ) @dataclass(frozen=True, slots=True) @@ -49,9 +54,12 @@ class KeyMaterial: class APIKeyService: """Creates and validates opaque API keys without retaining plaintext secrets.""" - def __init__(self, database: SecurityDatabase, settings: Settings) -> None: + def __init__( + self, database: SecurityDatabase, settings: Settings, tenants: TenantService + ) -> None: self.database = database self.settings = settings + self.tenants = tenants self.roles = configured_roles(settings.auth_role_scopes) self._last_used_cache: dict[str, float] = {} self._last_used_lock = asyncio.Lock() @@ -112,15 +120,19 @@ class APIKeyService: requests_per_minute=self.settings.auth_default_requests_per_minute, concurrent_jobs=self.settings.auth_default_concurrent_jobs, uploads_per_hour=self.settings.auth_default_uploads_per_hour, - processing_bytes_per_day=( - self.settings.auth_default_processing_bytes_per_day - ), + processing_bytes_per_day=(self.settings.auth_default_processing_bytes_per_day), ) ) await session.commit() + await self.tenants.ensure_all_api_key_principals() async def create( - self, payload: APIKeyCreate, *, created_by: str | None + self, + payload: APIKeyCreate, + *, + created_by: str | None, + workspace_id: str | None = None, + user_id: str | None = None, ) -> tuple[APIKey, str]: role = payload.role.strip().lower() if payload.role else None if role and role not in self.roles: @@ -141,12 +153,9 @@ class APIKeyService: created_by=created_by, notes=payload.notes, requests_per_minute=( - payload.requests_per_minute - or self.settings.auth_default_requests_per_minute - ), - concurrent_jobs=( - payload.concurrent_jobs or self.settings.auth_default_concurrent_jobs + payload.requests_per_minute or self.settings.auth_default_requests_per_minute ), + concurrent_jobs=(payload.concurrent_jobs or self.settings.auth_default_concurrent_jobs), uploads_per_hour=( payload.uploads_per_hour or self.settings.auth_default_uploads_per_hour ), @@ -159,6 +168,12 @@ class APIKeyService: session.add(record) await session.commit() await session.refresh(record) + if workspace_id and user_id: + await self.tenants.bind_api_key( + api_key_id=record.id, workspace_id=workspace_id, user_id=user_id + ) + else: + await self.tenants.resolve_api_key(record.id) return record, material.api_key async def authenticate(self, api_key: str) -> AuthContext: @@ -166,11 +181,7 @@ class APIKeyService: supplied_hash = self.hash_key(api_key) async with self.database.session() as session: candidates = list( - ( - await session.scalars( - select(APIKey).where(APIKey.key_prefix == key_prefix) - ) - ).all() + (await session.scalars(select(APIKey).where(APIKey.key_prefix == key_prefix))).all() ) record: APIKey | None = None for candidate in candidates: @@ -190,6 +201,7 @@ class APIKeyService: if expires_at is not None and expires_at <= now: raise UnauthorizedError scopes = effective_scopes(record.role, record.scopes or [], self.settings.auth_role_scopes) + principal = await self.tenants.resolve_api_key(record.id) return AuthContext( api_key_id=record.id, key_name=record.name, @@ -202,6 +214,10 @@ class APIKeyService: uploads_per_hour=record.uploads_per_hour, processing_bytes_per_day=record.processing_bytes_per_day, expires_at=expires_at, + workspace_id=principal.workspace_id, + user_id=principal.user_id, + membership_id=principal.membership_id, + membership_role=principal.membership_role, ) @staticmethod @@ -282,11 +298,7 @@ class APIKeyService: raise APIKeyConflictError("Revoked API keys cannot be changed") if status == "active" and record.status != "disabled": raise APIKeyConflictError("Only disabled API keys can be enabled") - if ( - status == "active" - and expires_at is not None - and expires_at <= utcnow() - ): + if status == "active" and expires_at is not None and expires_at <= utcnow(): raise APIKeyConflictError("Expired API keys cannot be enabled") if status == "disabled" and record.status != "active": raise APIKeyConflictError("Only active API keys can be disabled") @@ -300,6 +312,7 @@ class APIKeyService: async def rotate( self, key_id: str, grace_period_seconds: int, *, created_by: str | None ) -> tuple[APIKey, str]: + principal = await self.tenants.resolve_api_key(key_id) async with self.database.session() as session: old = await session.get(APIKey, key_id) if old is None: @@ -330,13 +343,16 @@ class APIKeyService: session.add(replacement) old.status = "rotating" if grace_period_seconds else "revoked" old.grace_expires_at = ( - utcnow() + timedelta(seconds=grace_period_seconds) - if grace_period_seconds - else None + utcnow() + timedelta(seconds=grace_period_seconds) if grace_period_seconds else None ) await session.commit() await session.refresh(replacement) - return replacement, material.api_key + await self.tenants.bind_api_key( + api_key_id=replacement.id, + workspace_id=principal.workspace_id, + user_id=principal.user_id, + ) + return replacement, material.api_key async def _touch_last_used(self, key_id: str) -> None: interval = self.settings.auth_last_used_update_seconds diff --git a/app/security/tenancy.py b/app/security/tenancy.py new file mode 100644 index 0000000000000000000000000000000000000000..b81a5d0d63e085d082b1d8de9e1a5d847677351d --- /dev/null +++ b/app/security/tenancy.py @@ -0,0 +1,188 @@ +from __future__ import annotations + +from dataclasses import dataclass +from uuid import uuid4 + +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from app.security.database import SecurityDatabase +from app.security.errors import ForbiddenError, UnauthorizedError +from app.security.models import APIKey, APIKeyPrincipal, User, Workspace, WorkspaceMembership + + +@dataclass(frozen=True, slots=True) +class TenantPrincipal: + """The persisted tenant authority resolved for an authenticated API key.""" + + workspace_id: str + user_id: str + membership_id: str + membership_role: str + + +class TenantService: + """Owns authoritative tenant records and API-key principal bindings. + + API keys remain authentication credentials. They never become workspace + IDs, and all tenant selection happens server-side through this service. + """ + + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + async def resolve_api_key(self, api_key_id: str) -> TenantPrincipal: + async with self.database.session() as session: + binding = await session.scalar( + select(APIKeyPrincipal) + .where(APIKeyPrincipal.api_key_id == api_key_id) + .with_for_update() + ) + if binding is not None: + return await self._validated_principal(session, binding) + + # Existing installs predate native tenancy. Provisioning a + # one-time private owner/workspace binding preserves access while + # ensuring the API-key ID itself can no longer act as a tenant. + key = await session.get(APIKey, api_key_id) + if key is None: + raise UnauthorizedError + user_id = str(uuid4()) + workspace_id = str(uuid4()) + membership_id = str(uuid4()) + user = User( + id=user_id, + subject=f"legacy-api-key-principal:{api_key_id}", + display_name=key.name, + metadata_json={"provisioned_from": "api_key"}, + ) + workspace = Workspace( + id=workspace_id, + slug=f"legacy-{api_key_id}", + name=f"{key.name} workspace", + metadata_json={"provisioned_from": "api_key"}, + ) + membership = WorkspaceMembership( + id=membership_id, + workspace_id=workspace_id, + user_id=user_id, + role="owner", + status="active", + ) + binding = APIKeyPrincipal( + api_key_id=api_key_id, + workspace_id=workspace_id, + user_id=user_id, + membership_id=membership_id, + status="active", + ) + session.add_all([user, workspace]) + await session.flush() + session.add_all([membership, binding]) + try: + await session.commit() + except IntegrityError: + # A concurrent first request won the unique binding race. + await session.rollback() + binding = await session.scalar( + select(APIKeyPrincipal).where(APIKeyPrincipal.api_key_id == api_key_id) + ) + if binding is None: + raise + return await self._validated_principal(session, binding) + return TenantPrincipal( + workspace_id=workspace_id, + user_id=user_id, + membership_id=membership_id, + membership_role=membership.role, + ) + + async def bind_api_key( + self, *, api_key_id: str, workspace_id: str, user_id: str + ) -> TenantPrincipal: + """Bind a newly created credential to an active persisted membership.""" + + async with self.database.session() as session: + membership = await session.scalar( + select(WorkspaceMembership).where( + WorkspaceMembership.workspace_id == workspace_id, + WorkspaceMembership.user_id == user_id, + ) + ) + if membership is None or membership.status != "active": + raise ForbiddenError + key = await session.get(APIKey, api_key_id) + if key is None: + raise UnauthorizedError + existing = await session.scalar( + select(APIKeyPrincipal).where(APIKeyPrincipal.api_key_id == api_key_id) + ) + if existing is None: + existing = APIKeyPrincipal( + api_key_id=api_key_id, + workspace_id=workspace_id, + user_id=user_id, + membership_id=membership.id, + status="active", + ) + session.add(existing) + await session.commit() + return await self._validated_principal(session, existing) + + async def ensure_all_api_key_principals(self) -> None: + """Provision legacy bindings before accepting requests after upgrade.""" + + async with self.database.session() as session: + key_ids = list( + ( + await session.scalars( + select(APIKey.id) + .outerjoin(APIKeyPrincipal, APIKeyPrincipal.api_key_id == APIKey.id) + .where(APIKeyPrincipal.id.is_(None)) + ) + ).all() + ) + for key_id in key_ids: + await self.resolve_api_key(key_id) + + async def list_principals(self) -> list[tuple[str, TenantPrincipal]]: + """Return server-side legacy-key mappings for one-time tenant adoption.""" + async with self.database.session() as session: + bindings = list((await session.scalars(select(APIKeyPrincipal))).all()) + result: list[tuple[str, TenantPrincipal]] = [] + for binding in bindings: + try: + result.append( + (binding.api_key_id, await self._validated_principal(session, binding)) + ) + except ForbiddenError: + # Disabled memberships must not migrate or retain access. + continue + return result + + @staticmethod + async def _validated_principal( + session: AsyncSession, binding: APIKeyPrincipal + ) -> TenantPrincipal: + membership = await session.get(WorkspaceMembership, binding.membership_id) + user = await session.get(User, binding.user_id) + workspace = await session.get(Workspace, binding.workspace_id) + if ( + binding.status != "active" + or membership is None + or membership.status != "active" + or membership.workspace_id != binding.workspace_id + or membership.user_id != binding.user_id + or user is None + or user.status != "active" + or workspace is None + or workspace.status != "active" + ): + raise ForbiddenError + return TenantPrincipal( + workspace_id=binding.workspace_id, + user_id=binding.user_id, + membership_id=membership.id, + membership_role=membership.role, + ) diff --git a/app/services/cleanup.py b/app/services/cleanup.py index 025bab43fca7dbec3a4049a05cfb8684cecd8cab..38929fd2ae7fa31c55f224a9182573e09fa8b921 100644 --- a/app/services/cleanup.py +++ b/app/services/cleanup.py @@ -70,6 +70,55 @@ class CleanupService: await asyncio.to_thread(shutil.move, str(source), str(destination)) return destination + async def publish_new(self, request_id: str, source: Path, filename: str) -> Path | None: + """Publish a generated output without replacing an existing file. + + This is used by durable generation-job reconciliation, where a retry + must never overwrite a canonical output created by an earlier worker + attempt. Existing media operations continue to use ``publish`` and + retain their established replacement semantics. + """ + + self._validate_request_id(request_id) + safe_name = Path(filename).name + if not safe_name or safe_name in {".", ".."}: + raise ProcessingError("The generated output filename is invalid") + destination_dir = self.settings.output_dir / request_id + destination_dir.mkdir(parents=True, exist_ok=True) + destination = destination_dir / safe_name + created = await asyncio.to_thread(self._publish_new_sync, source, destination) + return destination if created else None + + @staticmethod + def _publish_new_sync(source: Path, destination: Path) -> bool: + """Atomically claim a destination, with a cross-device fallback.""" + + try: + os.link(source, destination) + except FileExistsError: + return False + except OSError: + # ``temp_dir`` and ``output_dir`` can be different mounts. An + # exclusive create still prevents overwrite in that arrangement. + try: + with source.open("rb") as input_stream, destination.open("xb") as output_stream: + shutil.copyfileobj(input_stream, output_stream, length=1024 * 1024) + except FileExistsError: + return False + except OSError: + # Never leave a partial file reachable from the output + # directory when cross-device publication fails. + try: + destination.unlink(missing_ok=True) + except OSError: + pass + raise + try: + source.unlink() + except FileNotFoundError: + pass + return True + def resolve_download(self, request_id: str, filename: str) -> Path: self._validate_request_id(request_id) if filename != Path(filename).name: @@ -114,6 +163,15 @@ class CleanupService: if path.is_dir(): await asyncio.to_thread(shutil.rmtree, path) + async def remove_temporary_request(self, request_id: str) -> None: + """Remove only bounded staging data while retaining published output.""" + self._validate_request_id(request_id) + async with self._lock: + self._active.discard(request_id) + path = self.settings.temp_dir / request_id + if path.is_dir(): + await asyncio.to_thread(shutil.rmtree, path) + @staticmethod def _validate_request_id(request_id: str) -> None: try: diff --git a/app/services/ffmpeg_service.py b/app/services/ffmpeg_service.py index f96e374f9c5fe206e816cc1bb27c31ea78319ee6..7af4e951d3bfee2802aa7b1132aa18dc532ae1cb 100644 --- a/app/services/ffmpeg_service.py +++ b/app/services/ffmpeg_service.py @@ -26,12 +26,19 @@ class FFmpegService: *, operation: str, timeout: float | None = None, + cancel_event: asyncio.Event | None = None, ) -> None: command = [self.settings.ffmpeg_binary, "-hide_banner", "-nostdin", "-y"] + [ str(arg) for arg in args ] started = time.monotonic() - logger.info("ffmpeg started", extra={"operation": operation, "command": command}) + # Arguments can contain signed URLs, storage locators, and internal + # filesystem paths. Keep those out of structured logs while retaining + # enough bounded context to correlate and diagnose an invocation. + logger.info( + "ffmpeg started", + extra={"operation": operation, "argument_count": len(command) - 1}, + ) async with self._semaphore: try: process = await asyncio.create_subprocess_exec( @@ -42,29 +49,60 @@ class FFmpegService: stdout_task = asyncio.create_task(_read_limited(process.stdout, 4_000)) stderr_task = asyncio.create_task(_read_limited(process.stderr, 16_000)) try: - await asyncio.wait_for( - process.wait(), + wait_task = asyncio.create_task(process.wait()) + cancel_task = asyncio.create_task(cancel_event.wait()) if cancel_event else None + pending: set[asyncio.Task[object]] = set() + waitables: set[asyncio.Task[object]] = {wait_task} + if cancel_task is not None: + waitables.add(cancel_task) + done, pending = await asyncio.wait( + waitables, timeout=timeout or self.settings.download_timeout_seconds * 4, + return_when=asyncio.FIRST_COMPLETED, ) + if not done: + raise asyncio.TimeoutError + if ( + cancel_task is not None + and cancel_task in done + and cancel_event + and cancel_event.is_set() + ): + process.kill() + await process.wait() + await asyncio.gather(stdout_task, stderr_task, return_exceptions=True) + raise ProcessingError( + "FFmpeg processing was cancelled", details={"cancelled": True} + ) + await wait_task except asyncio.TimeoutError: process.kill() await process.wait() await asyncio.gather(stdout_task, stderr_task) raise + finally: + for task in pending if "pending" in locals() else set(): + task.cancel() + if "wait_task" in locals() and not wait_task.done(): + wait_task.cancel() + if ( + "cancel_task" in locals() + and cancel_task is not None + and not cancel_task.done() + ): + cancel_task.cancel() stdout, stderr = await asyncio.gather(stdout_task, stderr_task) except asyncio.TimeoutError as exc: raise ProcessingError("FFmpeg processing timed out") from exc except FileNotFoundError as exc: raise ProcessingError("FFmpeg is not installed or not available") from exc - stderr_text = stderr.decode("utf-8", errors="replace")[-16_000:] - stdout_text = stdout.decode("utf-8", errors="replace")[-4_000:] log_data = { "operation": operation, - "command": command, + "argument_count": len(command) - 1, "duration": round(time.monotonic() - started, 4), "return_code": process.returncode, - "stdout": stdout_text, - "stderr": stderr_text, + "stdout_bytes": len(stdout), + "stderr_bytes": len(stderr), } if process.returncode != 0: logger.error("ffmpeg failed", extra=log_data) diff --git a/app/services/media_service.py b/app/services/media_service.py index bc89e2559986da9362e3b51758908dcad625c98e..5819e80331002cd2619e4f05c90ea8cd119f8d10 100644 --- a/app/services/media_service.py +++ b/app/services/media_service.py @@ -17,6 +17,8 @@ from app.services.input_resolver import InputResolver from app.services.validator import MediaValidator from app.services.whisper_service import WhisperService from app.services.ytdlp_service import YTDLPService +from app.security.assets import CanonicalAssetService +from app.security.context import auth_context logger = get_logger(__name__) Operation = Callable[ @@ -35,6 +37,7 @@ class MediaProcessor: ffprobe: FFprobeService, ytdlp: YTDLPService, whisper: WhisperService, + assets: CanonicalAssetService | None = None, ) -> None: self.settings = settings self.resolver = resolver @@ -44,6 +47,7 @@ class MediaProcessor: self.ffprobe = ffprobe self.ytdlp = ytdlp self.whisper = whisper + self.assets = assets async def run( self, resolved: ResolvedRequest, operation_name: str, operation: Operation @@ -161,6 +165,7 @@ class MediaProcessor: """Publish a result and build the shared structured success response.""" download_url: str | None = None output_size = 0 + canonical_asset_id: str | None = None if result.path is not None: published = await self.cleanup.publish( resolved.request_id, result.path, result.filename or result.path.name @@ -168,12 +173,27 @@ class MediaProcessor: output_size = published.stat().st_size base_url = self.settings.base_url.rstrip("/") download_url = f"{base_url}/v1/media/{resolved.request_id}/{published.name}" + # Outputs are registered by the pipeline, never claimed later by + # a user-supplied path. Anonymous deployments retain legacy media + # behavior but cannot use the social ownership boundary. + principal = auth_context.get() + if self.assets is not None and principal and principal.workspace_id: + record = await self.assets.register_output( + workspace_id=principal.workspace_id, + user_id=principal.user_id, + request_id=resolved.request_id, + path=published, + mime_type=result.mime_type or self.validator.infer_mime(published), + metadata={"operation": operation_name}, + ) + canonical_asset_id = record.id elapsed = round(time.monotonic() - started, 4) metadata = { "operation": operation_name, **result.metadata, "inputs": probe_metadata, "output_size": output_size, + **({"asset_id": canonical_asset_id} if canonical_asset_id else {}), **(extra_metadata or {}), } logger.info( diff --git a/app/social/database.py b/app/social/database.py index 6bbc36c36ca97bea9624b930a39c915922045523..2a7f3f3f30d9099ebc131ecdd0cade57910832a9 100644 --- a/app/social/database.py +++ b/app/social/database.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextvars from collections.abc import AsyncIterator from contextlib import asynccontextmanager @@ -11,9 +12,13 @@ from sqlalchemy.ext.asyncio import ( create_async_engine, ) +from app.analytics import models as analytics_models # noqa: F401 from app.core.config import Settings +from app.core.database_url import normalize_async_database_url from app.social.models import SocialBase +_ANALYTICS_MODELS_REGISTERED = analytics_models + REQUIRED_SOCIAL_TABLES = frozenset( { "social_accounts", @@ -32,20 +37,56 @@ REQUIRED_SOCIAL_TABLES = frozenset( "social_webhook_events", "social_post_metrics", "social_audit_events", + "social_publishing_batches", + "social_publishing_batch_items", + "analytics_sync_runs", + "analytics_metric_snapshots", + "analytics_post_metrics", + "analytics_platform_metrics", } ) +REQUIRED_SOCIAL_COLUMNS: dict[str, frozenset[str]] = { + "social_media_assets": frozenset({"canonical_asset_id"}), + "social_posts": frozenset( + {"project_id", "canonical_caption", "canonical_hashtags", "revision"} + ), + "social_schedules": frozenset({"revision"}), + "social_post_targets": frozenset({"cancellation_requested_at"}), + "social_jobs": frozenset({"provider_state_encrypted", "cancellation_requested_at"}), + "oauth_states": frozenset({"requested_account_type", "requested_scopes"}), + "analytics_sync_runs": frozenset( + {"workspace_id", "idempotency_key", "status", "date_from", "date_to"} + ), +} + +_trusted_worker_context: contextvars.ContextVar[bool] = contextvars.ContextVar( + "trusted_social_worker_context", default=False +) + class SocialDatabase: """Social persistence with migration-only production schema changes.""" def __init__(self, settings: Settings) -> None: self.settings = settings - self.database_url = settings.resolved_social_database_url + self.database_url = normalize_async_database_url(settings.resolved_social_database_url) self.engine: AsyncEngine = create_async_engine(self.database_url, pool_pre_ping=True) if self.database_url.startswith("sqlite"): event.listen(self.engine.sync_engine, "connect", self._configure_sqlite) - self.session_factory = async_sessionmaker(self.engine, expire_on_commit=False, class_=AsyncSession) + self.session_factory = async_sessionmaker( + self.engine, expire_on_commit=False, class_=AsyncSession + ) + self.worker_database_url = normalize_async_database_url(settings.social_worker_database_url) + self.worker_engine: AsyncEngine | None = None + self.worker_session_factory: async_sessionmaker[AsyncSession] | None = None + if self.worker_database_url: + self.worker_engine = create_async_engine(self.worker_database_url, pool_pre_ping=True) + if self.worker_database_url.startswith("sqlite"): + event.listen(self.worker_engine.sync_engine, "connect", self._configure_sqlite) + self.worker_session_factory = async_sessionmaker( + self.worker_engine, expire_on_commit=False, class_=AsyncSession + ) @staticmethod def _configure_sqlite(dbapi_connection: object, _record: object) -> None: @@ -59,32 +100,211 @@ class SocialDatabase: async with self.engine.begin() as connection: await connection.run_sync(SocialBase.metadata.create_all) + @property + def is_postgres(self) -> bool: + return self.database_url.startswith(("postgresql", "postgres")) + + async def verify_execution_boundaries(self) -> None: + """Fail closed when a Postgres deployment cannot enforce RLS safely.""" + + if not self.is_postgres or not self.settings.social_enforce_rls: + return + tenant_role = self.settings.social_tenant_database_role.strip() + if not tenant_role: + raise RuntimeError( + "SOCIAL_TENANT_DATABASE_ROLE is required when SOCIAL_ENFORCE_RLS is enabled." + ) + tenant = await self._role_attributes(self.engine) + if tenant["role"] != tenant_role: + raise RuntimeError( + "SOCIAL_DATABASE_URL is not connected as SOCIAL_TENANT_DATABASE_ROLE." + ) + if tenant["bypass_rls"] or tenant["superuser"]: + raise RuntimeError( + "SOCIAL_DATABASE_URL must use a non-privileged tenant role, never a service role." + ) + # Startup also adopts historic API-key tenant rows through this + # boundary, so every enabled PostgreSQL social deployment needs it, + # even if the scheduler is temporarily disabled. + if not self.worker_engine: + raise RuntimeError( + "SOCIAL_WORKER_DATABASE_URL is required for PostgreSQL social access." + ) + worker_role = self.settings.social_worker_database_role.strip() + if not worker_role: + raise RuntimeError("SOCIAL_WORKER_DATABASE_ROLE is required for trusted worker access.") + worker = await self._role_attributes(self.worker_engine) + if worker["role"] != worker_role or not worker["bypass_rls"]: + raise RuntimeError( + "SOCIAL_WORKER_DATABASE_URL must use the configured BYPASSRLS worker role." + ) + + @staticmethod + async def _role_attributes(engine: AsyncEngine) -> dict[str, object]: + async with engine.connect() as connection: + row = ( + ( + await connection.execute( + text( + "select current_user as role, r.rolbypassrls as bypass_rls, r.rolsuper as superuser " + "from pg_roles r where r.rolname = current_user" + ) + ) + ) + .mappings() + .one_or_none() + ) + if row is None: + raise RuntimeError("Unable to verify the active PostgreSQL database role.") + return dict(row) + async def schema_ready(self) -> bool: """Check the complete Phase 1 schema without changing the database.""" async with self.engine.connect() as connection: - tables = await connection.run_sync( - lambda sync: set(inspect(sync).get_table_names()) - ) - return REQUIRED_SOCIAL_TABLES.issubset(tables) + tables, columns = await connection.run_sync(self._schema_snapshot) + return REQUIRED_SOCIAL_TABLES.issubset(tables) and all( + required.issubset(columns.get(table, set())) + for table, required in REQUIRED_SOCIAL_COLUMNS.items() + ) async def missing_tables(self) -> list[str]: """Return absent required tables for an actionable startup warning.""" async with self.engine.connect() as connection: - tables = await connection.run_sync( - lambda sync: set(inspect(sync).get_table_names()) - ) - return sorted(REQUIRED_SOCIAL_TABLES - tables) + tables, columns = await connection.run_sync(self._schema_snapshot) + missing = list(REQUIRED_SOCIAL_TABLES - tables) + for table, required in REQUIRED_SOCIAL_COLUMNS.items(): + missing.extend(f"{table}.{column}" for column in required - columns.get(table, set())) + return sorted(missing) + + async def adopt_legacy_workspace( + self, *, legacy_workspace_id: str, workspace_id: str, user_id: str + ) -> int: + """Move historic API-key tenant rows to the authoritative workspace.""" + if legacy_workspace_id == workspace_id: + raise RuntimeError("Legacy and authoritative workspace IDs must differ.") + workspace_tables = ( + "social_accounts", + "media_variants", + "social_media_assets", + "social_campaigns", + "social_posts", + "social_jobs", + "social_webhook_events", + "social_audit_events", + "oauth_states", + "social_publishing_batches", + ) + changed = 0 + async with self.worker_session() as session: + for table in workspace_tables: + result = await session.execute( + text( + f"update {table} set workspace_id = :workspace_id " + "where workspace_id = :legacy_workspace_id" + ), + {"workspace_id": workspace_id, "legacy_workspace_id": legacy_workspace_id}, + ) + changed += max(0, int(result.rowcount or 0)) + for table, column in (("social_posts", "created_by"), ("oauth_states", "user_id")): + result = await session.execute( + text( + f"update {table} set {column} = :user_id " + f"where {column} = :legacy_workspace_id" + ), + {"user_id": user_id, "legacy_workspace_id": legacy_workspace_id}, + ) + changed += max(0, int(result.rowcount or 0)) + await session.commit() + return changed + + @staticmethod + def _schema_snapshot(connection: object) -> tuple[set[str], dict[str, set[str]]]: + inspector = inspect(connection) + tables = set(inspector.get_table_names()) + columns = { + table: {column["name"] for column in inspector.get_columns(table)} + for table in REQUIRED_SOCIAL_COLUMNS + if table in tables + } + return tables, columns async def close(self) -> None: await self.engine.dispose() + if self.worker_engine is not None: + await self.worker_engine.dispose() @asynccontextmanager async def session(self, workspace_id: str | None = None) -> AsyncIterator[AsyncSession]: + """Open a tenant session; no-context access is OAuth-state compatibility only.""" + if workspace_id is None: + async with self.oauth_session() as session: + yield session + return + context = ( + self.worker_tenant_session(workspace_id) + if _trusted_worker_context.get() + else self.tenant_session(workspace_id) + ) + async with context as session: + yield session + + @asynccontextmanager + async def tenant_session(self, workspace_id: str) -> AsyncIterator[AsyncSession]: + if not workspace_id: + raise RuntimeError("A tenant session requires an authoritative workspace ID.") async with self.session_factory() as session: - if workspace_id and self.database_url.startswith(("postgresql", "postgres")): + if self.is_postgres: # RLS policies read this transaction-local tenant identity. await session.execute( text("select set_config('app.workspace_id', :workspace_id, true)"), {"workspace_id": workspace_id}, ) yield session + + @asynccontextmanager + async def oauth_session(self) -> AsyncIterator[AsyncSession]: + """The sole non-tenant API session, for random single-use OAuth state.""" + async with self.session_factory() as session: + yield session + + @asynccontextmanager + async def worker_session(self) -> AsyncIterator[AsyncSession]: + """Backend-only cross-workspace session for scheduler, jobs, and Vault.""" + if self.is_postgres: + if self.worker_session_factory is None: + raise RuntimeError("Trusted worker database access is not configured.") + async with self.worker_session_factory() as session: + yield session + return + # SQLite has no RLS. It remains supported for local/unit-test use only. + async with self.session_factory() as session: + yield session + + @asynccontextmanager + async def worker_tenant_session(self, workspace_id: str) -> AsyncIterator[AsyncSession]: + """Trusted worker session annotated with the job's tenant for auditability.""" + if not workspace_id: + raise RuntimeError("A worker tenant session requires a workspace ID.") + if not self.is_postgres: + async with self.tenant_session(workspace_id) as session: + yield session + return + if self.worker_session_factory is None: + raise RuntimeError("Trusted worker database access is not configured.") + async with self.worker_session_factory() as session: + await session.execute( + text("select set_config('app.workspace_id', :workspace_id, true)"), + {"workspace_id": workspace_id}, + ) + yield session + + @asynccontextmanager + async def worker_boundary(self) -> AsyncIterator[None]: + """Mark a scheduler/publisher call tree as trusted worker execution.""" + if self.is_postgres and self.worker_session_factory is None: + raise RuntimeError("Trusted worker database access is not configured.") + token = _trusted_worker_context.set(True) + try: + yield + finally: + _trusted_worker_context.reset(token) diff --git a/app/social/domain/capabilities.py b/app/social/domain/capabilities.py index 62c12c5543f2e5d91030050347186528306f525d..7c54748459736c0dd59652d0790a3c87ec27b1be 100644 --- a/app/social/domain/capabilities.py +++ b/app/social/domain/capabilities.py @@ -5,6 +5,24 @@ from pydantic import BaseModel, ConfigDict, Field from app.social.domain.enums import ConnectionStrategy, Provider +class PublishingFieldCapabilities(BaseModel): + model_config = ConfigDict(extra="forbid") + + media_types: list[str] = Field(default_factory=list) + max_media_count: int | None = Field(default=None, ge=1) + max_duration_seconds: int | None = Field(default=None, ge=1) + max_file_size_bytes: int | None = Field(default=None, ge=1) + caption_max_length: int | None = Field(default=None, ge=1) + hashtags: bool = False + first_comment: bool = False + title: bool = False + description: bool = False + thumbnail: bool = False + privacy: bool = False + visibility: bool = False + location: bool = False + + class ProviderCapabilities(BaseModel): """Provider metadata consumed by every transport and UI.""" @@ -40,21 +58,20 @@ class ProviderCapabilities(BaseModel): # independently. Keeping this mapping in capability metadata lets every # transport render the right elevation without requesting a scope for a # different account type. - account_type_publishing_scopes: dict[str, list[str]] = Field( - default_factory=dict - ) + account_type_publishing_scopes: dict[str, list[str]] = Field(default_factory=dict) # Additional authorization is always opt-in. These scopes are never # appended to a normal OAuth connection request. analytics_required_scopes: list[str] = Field(default_factory=list) # Analytics consent can differ by account type. For example, LinkedIn # member post analytics requires a restricted member-read scope while # organization post statistics require organization-admin reporting access. - account_type_analytics_scopes: dict[str, list[str]] = Field( - default_factory=dict - ) + account_type_analytics_scopes: dict[str, list[str]] = Field(default_factory=dict) # Transport-neutral metadata consumed by dynamic clients. Provider-owned # runtime choices are fetched from the account publish-options endpoint. publish_metadata_schema: dict[str, object] = Field(default_factory=dict) + publishing_fields: PublishingFieldCapabilities = Field( + default_factory=PublishingFieldCapabilities + ) @property def publish_supported(self) -> bool: diff --git a/app/social/domain/enums.py b/app/social/domain/enums.py index 908f8e5f1a84b25e4a1e8000e1435d5bcfeb440b..4535dea9f1681752b6fac792f7db181435ed686a 100644 --- a/app/social/domain/enums.py +++ b/app/social/domain/enums.py @@ -30,6 +30,8 @@ class AccountStatus(StrEnum): class PostStatus(StrEnum): DRAFT = "draft" + VALIDATING = "validating" + READY = "ready" SCHEDULED = "scheduled" QUEUED = "queued" PREPARING = "preparing" @@ -63,6 +65,12 @@ class PublishMode(StrEnum): DRAFT = "draft" -TERMINAL_JOB_STATUSES = frozenset( - {JobStatus.PUBLISHED, JobStatus.FAILED, JobStatus.CANCELLED} -) +class PublishingBatchOperation(StrEnum): + SCHEDULE = "schedule" + RESCHEDULE = "reschedule" + CANCEL = "cancel" + DUPLICATE = "duplicate" + DELETE = "delete" + + +TERMINAL_JOB_STATUSES = frozenset({JobStatus.PUBLISHED, JobStatus.FAILED, JobStatus.CANCELLED}) diff --git a/app/social/domain/errors.py b/app/social/domain/errors.py index 06cf3f8297b62f6847311f9301a5d8c5ac4c2112..25819a1d5c74115d83b1c5c8aa034af6bd49e3d2 100644 --- a/app/social/domain/errors.py +++ b/app/social/domain/errors.py @@ -91,3 +91,28 @@ class SocialOAuthStateError(SocialError): class SocialTransitionError(SocialError): code = "SOCIAL_INVALID_STATE_TRANSITION" status_code = 409 + + +class PublishScheduleInvalidError(SocialError): + code = "PUBLISH_SCHEDULE_INVALID" + status_code = 422 + + +class PublishAlreadyProcessingError(SocialError): + code = "PUBLISH_ALREADY_PROCESSING" + status_code = 409 + + +class PublishNotReschedulableError(SocialError): + code = "PUBLISH_NOT_RESCHEDULABLE" + status_code = 409 + + +class PublishRevisionConflictError(SocialError): + code = "PUBLISH_REVISION_CONFLICT" + status_code = 409 + + +class PublishPermissionDeniedError(SocialError): + code = "PUBLISH_PERMISSION_DENIED" + status_code = 403 diff --git a/app/social/migrations/0007_tenant_asset_and_rls_hardening.sql b/app/social/migrations/0007_tenant_asset_and_rls_hardening.sql new file mode 100644 index 0000000000000000000000000000000000000000..2b84be61e9f956a46f990be6c05517581661cfd0 --- /dev/null +++ b/app/social/migrations/0007_tenant_asset_and_rls_hardening.sql @@ -0,0 +1,61 @@ +-- Bind social media references to canonical owned assets and force social RLS. +-- Apply after 0001 through 0006. No data is removed by this migration. + +begin; + +alter table social_media_assets add column if not exists canonical_asset_id text; +create index if not exists ix_social_media_assets_canonical_asset + on social_media_assets(canonical_asset_id) + where canonical_asset_id is not null; + +-- Existing records remain readable for migration visibility, but the +-- application refuses to publish them unless a matching canonical asset exists. + +create or replace function mediarouter_social_assert_media_asset_workspace() +returns trigger language plpgsql as $$ +declare post_workspace text; +declare asset_workspace text; +begin + if tg_table_name = 'social_posts' and new.media_asset_id is not null then + select workspace_id into asset_workspace from social_media_assets where id = new.media_asset_id; + if asset_workspace is null or asset_workspace is distinct from new.workspace_id then + raise exception 'social post media asset must belong to the post workspace' using errcode = '23503'; + end if; + elsif tg_table_name = 'social_post_media' and new.media_asset_id is not null then + select workspace_id into post_workspace from social_posts where id = new.social_post_id; + select workspace_id into asset_workspace from social_media_assets where id = new.media_asset_id; + if post_workspace is null or asset_workspace is distinct from post_workspace then + raise exception 'social post media asset must belong to the post workspace' using errcode = '23503'; + end if; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_social_post_media_asset_workspace on social_posts; +create trigger mediarouter_social_post_media_asset_workspace +before insert or update of workspace_id, media_asset_id on social_posts +for each row execute function mediarouter_social_assert_media_asset_workspace(); +drop trigger if exists mediarouter_social_post_media_item_asset_workspace on social_post_media; +create trigger mediarouter_social_post_media_item_asset_workspace +before insert or update of social_post_id, media_asset_id on social_post_media +for each row execute function mediarouter_social_assert_media_asset_workspace(); + +do $$ +declare table_name text; +begin + foreach table_name in array array[ + 'social_accounts', 'social_account_tokens', 'social_account_capabilities', + 'media_variants', 'social_media_assets', 'social_campaigns', 'social_posts', + 'social_post_targets', 'social_post_media', 'social_schedules', 'social_jobs', + 'social_job_attempts', 'social_webhook_events', 'social_post_metrics', + 'social_audit_events' + ] loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + end loop; +end $$; + +-- oauth_states intentionally remains non-RLS: OAuth provider callbacks are +-- unauthenticated and rely on a high-entropy, expiring, single-use state. + +commit; diff --git a/app/social/migrations/0008_unified_publishing.sql b/app/social/migrations/0008_unified_publishing.sql new file mode 100644 index 0000000000000000000000000000000000000000..4f1b8d8df8bb8bc49b107a1a0613d51a06755879 --- /dev/null +++ b/app/social/migrations/0008_unified_publishing.sql @@ -0,0 +1,64 @@ +-- Unified Publishing Center: project provenance, canonical copy, and truthful +-- cancellation requests. This migration extends the existing Social domain. + +begin; + +alter table social_posts add column if not exists project_id text; +alter table social_posts add column if not exists canonical_caption text; +alter table social_posts add column if not exists canonical_hashtags jsonb not null default '[]'::jsonb; +alter table social_post_targets add column if not exists cancellation_requested_at timestamptz; +alter table social_jobs add column if not exists cancellation_requested_at timestamptz; + +do $$ +begin + if not exists (select 1 from pg_constraint where conname = 'fk_social_posts_project') then + alter table social_posts add constraint fk_social_posts_project + foreign key (project_id) references projects(id) on delete restrict; + end if; + if not exists (select 1 from pg_constraint where conname = 'ck_social_posts_canonical_caption_length') then + alter table social_posts add constraint ck_social_posts_canonical_caption_length + check (canonical_caption is null or char_length(canonical_caption) <= 10000); + end if; + if not exists (select 1 from pg_constraint where conname = 'ck_social_posts_canonical_hashtags') then + alter table social_posts add constraint ck_social_posts_canonical_hashtags + check (jsonb_typeof(canonical_hashtags) = 'array' and jsonb_array_length(canonical_hashtags) <= 100); + end if; +end; +$$; + +create index if not exists ix_social_posts_workspace_project_created + on social_posts(workspace_id, project_id, created_at desc) + where project_id is not null; +create index if not exists ix_social_jobs_cancellation_requested + on social_jobs(status, cancellation_requested_at) + where cancellation_requested_at is not null; + +create or replace function mediarouter_social_assert_project_workspace() +returns trigger language plpgsql as $$ +declare project_workspace text; +begin + if new.project_id is not null then + select workspace_id into project_workspace from projects where id = new.project_id; + if project_workspace is null or project_workspace is distinct from new.workspace_id then + raise exception 'social post project must belong to the same workspace' using errcode = '23503'; + end if; + end if; + if tg_op = 'UPDATE' and new.project_id is distinct from old.project_id then + raise exception 'social post project provenance is immutable' using errcode = '23514'; + end if; + return new; +end; +$$; +drop trigger if exists mediarouter_social_post_project_workspace on social_posts; +create trigger mediarouter_social_post_project_workspace +before insert or update of workspace_id, project_id on social_posts +for each row execute function mediarouter_social_assert_project_workspace(); + +alter table social_posts enable row level security; +alter table social_posts force row level security; +alter table social_post_targets enable row level security; +alter table social_post_targets force row level security; +alter table social_jobs enable row level security; +alter table social_jobs force row level security; + +commit; diff --git a/app/social/migrations/0009_publishing_operations.sql b/app/social/migrations/0009_publishing_operations.sql new file mode 100644 index 0000000000000000000000000000000000000000..c1fe863e6ddc7f8d194c7324faa3b81f7db9f53b --- /dev/null +++ b/app/social/migrations/0009_publishing_operations.sql @@ -0,0 +1,61 @@ +begin; + +alter table social_posts add column if not exists revision integer not null default 1; +alter table social_schedules add column if not exists revision integer not null default 1; +alter table social_posts drop constraint if exists ck_social_posts_status; +alter table social_posts drop constraint if exists ck_social_posts_status_phase9; +alter table social_posts add constraint ck_social_posts_status_phase9 + check (status in ('draft', 'validating', 'ready', 'scheduled', 'queued', 'preparing', 'processing', 'uploading', 'publishing', 'published', 'partial_success', 'retrying', 'failed', 'cancelled')); +create index if not exists ix_social_posts_workspace_status_updated + on social_posts(workspace_id, status, updated_at desc); +create index if not exists ix_social_schedules_workspace_due + on social_schedules(scheduled_at, status); + +create table if not exists social_publishing_batches ( + id text primary key, + workspace_id text not null, + operation text not null check (operation in ('schedule','reschedule','cancel','duplicate','delete')), + status text not null default 'queued' check (status in ('queued','processing','partial_success','completed','failed','cancelled')), + idempotency_key text not null check (char_length(idempotency_key) between 1 and 255), + requested_by text, + created_at timestamptz not null default now(), + started_at timestamptz, + completed_at timestamptz, + constraint uq_social_publishing_batch_idempotency unique (workspace_id, idempotency_key) +); +create index if not exists ix_social_publishing_batches_workspace_status + on social_publishing_batches(workspace_id, status); + +create table if not exists social_publishing_batch_items ( + id text primary key, + batch_id text not null references social_publishing_batches(id) on delete cascade, + social_post_id text not null references social_posts(id) on delete cascade, + status text not null default 'queued' check (status in ('queued','processing','completed','failed','cancelled')), + error_code text, + error_message text, + metadata jsonb not null default '{}'::jsonb, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint uq_social_publishing_batch_item unique (batch_id, social_post_id) +); +create index if not exists ix_social_publishing_batch_items_status + on social_publishing_batch_items(status); + +do $$ +declare table_name text; +begin + foreach table_name in array array['social_publishing_batches'] loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + execute format('drop policy if exists social_workspace_isolation on %I', table_name); + execute format('create policy social_workspace_isolation on %I using (workspace_id = current_setting(''app.workspace_id'', true)) with check (workspace_id = current_setting(''app.workspace_id'', true))', table_name); + end loop; + alter table social_publishing_batch_items enable row level security; + alter table social_publishing_batch_items force row level security; + drop policy if exists social_batch_item_workspace_isolation on social_publishing_batch_items; + create policy social_batch_item_workspace_isolation on social_publishing_batch_items + using (exists (select 1 from social_publishing_batches b where b.id = social_publishing_batch_items.batch_id and b.workspace_id = current_setting('app.workspace_id', true))) + with check (exists (select 1 from social_publishing_batches b where b.id = social_publishing_batch_items.batch_id and b.workspace_id = current_setting('app.workspace_id', true))); +end $$; + +commit; diff --git a/app/social/migrations/0010_analytics_insights.sql b/app/social/migrations/0010_analytics_insights.sql new file mode 100644 index 0000000000000000000000000000000000000000..c8d7a9228ad5f3622d312c0fec997dda2c3bb3e4 --- /dev/null +++ b/app/social/migrations/0010_analytics_insights.sql @@ -0,0 +1,133 @@ +begin; + +create table if not exists analytics_sync_runs ( + id text primary key, + workspace_id text not null, + project_id text, + provider text, + date_from timestamptz not null, + date_to timestamptz not null, + timezone text not null, + status text not null default 'queued' + check (status in ('queued','running','succeeded','partial','failed','cancelled')), + idempotency_key text not null check (char_length(idempotency_key) between 8 and 255), + requested_by text, + attempt_count integer not null default 0 check (attempt_count >= 0), + next_attempt_at timestamptz, + error_code text, + error_message text, + metrics_count integer not null default 0 check (metrics_count >= 0), + created_at timestamptz not null default now(), + started_at timestamptz, + completed_at timestamptz, + updated_at timestamptz not null default now(), + constraint ck_analytics_sync_range check (date_to > date_from), + constraint uq_analytics_sync_workspace_idempotency unique (workspace_id, idempotency_key) +); +create index if not exists ix_analytics_sync_workspace_status + on analytics_sync_runs(workspace_id, status); +create index if not exists ix_analytics_sync_due + on analytics_sync_runs(status, next_attempt_at); + +create table if not exists analytics_metric_snapshots ( + id text primary key, + workspace_id text not null, + project_id text, + social_post_id text references social_posts(id) on delete cascade, + social_account_id text references social_accounts(id) on delete cascade, + provider text not null, + external_object_id text not null, + metric_name text not null, + metric_value double precision not null, + bucket_start timestamptz not null, + dimensions jsonb not null default '{}'::jsonb, + source text not null default 'provider', + collected_at timestamptz not null default now(), + provider_updated_at timestamptz, + constraint uq_analytics_metric_snapshot + unique (workspace_id, provider, external_object_id, metric_name, bucket_start) +); +create index if not exists ix_analytics_metric_snapshots_workspace_bucket + on analytics_metric_snapshots(workspace_id, bucket_start desc); + +create table if not exists analytics_post_metrics ( + id text primary key, + workspace_id text not null, + project_id text, + social_post_id text not null references social_posts(id) on delete cascade, + social_post_target_id text not null references social_post_targets(id) on delete cascade, + social_account_id text not null references social_accounts(id) on delete cascade, + provider text not null, + external_post_id text, + metric_date timestamptz not null, + views bigint, + impressions bigint, + likes bigint, + comments bigint, + shares bigint, + engagement_rate double precision, + dimensions jsonb not null default '{}'::jsonb, + source text not null default 'provider', + collected_at timestamptz not null default now(), + provider_updated_at timestamptz, + constraint ck_analytics_post_nonnegative check ( + (views is null or views >= 0) and + (impressions is null or impressions >= 0) and + (likes is null or likes >= 0) and + (comments is null or comments >= 0) and + (shares is null or shares >= 0) and + (engagement_rate is null or engagement_rate >= 0) + ), + constraint uq_analytics_post_metric_bucket + unique (workspace_id, social_post_target_id, metric_date) +); +create index if not exists ix_analytics_post_metrics_workspace_date + on analytics_post_metrics(workspace_id, metric_date desc); +create index if not exists ix_analytics_post_metrics_project_date + on analytics_post_metrics(workspace_id, project_id, metric_date desc); +create index if not exists ix_analytics_post_metrics_provider_date + on analytics_post_metrics(workspace_id, provider, metric_date desc); + +create table if not exists analytics_platform_metrics ( + id text primary key, + workspace_id text not null, + project_id text, + social_account_id text not null references social_accounts(id) on delete cascade, + provider text not null, + metric_date timestamptz not null, + posts_count integer not null default 0 check (posts_count >= 0), + views bigint, + impressions bigint, + likes bigint, + comments bigint, + shares bigint, + engagement_rate double precision, + dimensions jsonb not null default '{}'::jsonb, + source text not null default 'provider', + collected_at timestamptz not null default now(), + constraint uq_analytics_platform_metric_bucket + unique (workspace_id, provider, social_account_id, metric_date) +); +create index if not exists ix_analytics_platform_metrics_workspace_date + on analytics_platform_metrics(workspace_id, metric_date desc); + +do $$ +declare table_name text; +begin + foreach table_name in array array[ + 'analytics_sync_runs', + 'analytics_metric_snapshots', + 'analytics_post_metrics', + 'analytics_platform_metrics' + ] loop + execute format('alter table %I enable row level security', table_name); + execute format('alter table %I force row level security', table_name); + execute format('drop policy if exists analytics_workspace_isolation on %I', table_name); + execute format( + 'create policy analytics_workspace_isolation on %I using (workspace_id = current_setting(''app.workspace_id'', true)) with check (workspace_id = current_setting(''app.workspace_id'', true))', + table_name + ); + end loop; +end $$; + +commit; diff --git a/app/social/models.py b/app/social/models.py index 21daba5edfec52a5c2a93efbc8ea38d271d71c24..5b870659505d9e4cb94ec22332f27adccc8f37c5 100644 --- a/app/social/models.py +++ b/app/social/models.py @@ -32,7 +32,12 @@ class SocialBase(DeclarativeBase): class SocialAccount(SocialBase): __tablename__ = "social_accounts" __table_args__ = ( - UniqueConstraint("workspace_id", "provider", "external_account_id", name="uq_social_account_workspace_provider_external"), + UniqueConstraint( + "workspace_id", + "provider", + "external_account_id", + name="uq_social_account_workspace_provider_external", + ), Index("ix_social_accounts_workspace_id", "workspace_id"), Index("ix_social_accounts_provider_status", "provider", "status"), ) @@ -45,9 +50,15 @@ class SocialAccount(SocialBase): display_name: Mapped[str | None] = mapped_column(String(255)) avatar_url: Mapped[str | None] = mapped_column(String(2048)) status: Mapped[str] = mapped_column(String(32), nullable=False, default="pending") - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) last_synced_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) @@ -58,7 +69,9 @@ class SocialAccountToken(SocialBase): Index("ix_social_account_tokens_expires_at", "expires_at"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_account_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False) + social_account_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False + ) access_token_secret_id: Mapped[str | None] = mapped_column(String(255)) refresh_token_secret_id: Mapped[str | None] = mapped_column(String(255)) encrypted_payload: Mapped[str | None] = mapped_column(Text) @@ -67,8 +80,12 @@ class SocialAccountToken(SocialBase): token_type: Mapped[str | None] = mapped_column(String(64)) last_refreshed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) class SocialAccountCapability(SocialBase): @@ -78,11 +95,17 @@ class SocialAccountCapability(SocialBase): Index("ix_social_account_capabilities_account", "social_account_id"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_account_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False) + social_account_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="CASCADE"), nullable=False + ) capability: Mapped[str] = mapped_column(String(100), nullable=False) enabled: Mapped[bool] = mapped_column(nullable=False, default=False) - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) class MediaVariant(SocialBase): @@ -104,8 +127,12 @@ class MediaVariant(SocialBase): container: Mapped[str | None] = mapped_column(String(64)) bitrate: Mapped[int | None] = mapped_column(BigInteger) file_size: Mapped[int | None] = mapped_column(BigInteger) - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) class SocialMediaAsset(SocialBase): @@ -119,17 +146,26 @@ class SocialMediaAsset(SocialBase): __tablename__ = "social_media_assets" __table_args__ = ( - UniqueConstraint("workspace_id", "request_id", "filename", name="uq_social_media_asset_workspace_output"), + UniqueConstraint( + "workspace_id", "request_id", "filename", name="uq_social_media_asset_workspace_output" + ), Index("ix_social_media_assets_workspace_id", "workspace_id"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + # Canonical ownership lives in the security/asset domain. This local + # social binding is cacheable publishing metadata, not ownership proof. + canonical_asset_id: Mapped[str | None] = mapped_column(String(36), nullable=True) request_id: Mapped[str] = mapped_column(String(36), nullable=False) filename: Mapped[str] = mapped_column(String(255), nullable=False) mime_type: Mapped[str] = mapped_column(String(255), nullable=False) file_size: Mapped[int] = mapped_column(BigInteger, nullable=False) - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) class SocialCampaign(SocialBase): @@ -140,32 +176,55 @@ class SocialCampaign(SocialBase): name: Mapped[str] = mapped_column(String(255), nullable=False) description: Mapped[str | None] = mapped_column(Text) status: Mapped[str] = mapped_column(String(32), nullable=False, default="draft") - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) class SocialPost(SocialBase): __tablename__ = "social_posts" __table_args__ = ( - UniqueConstraint("workspace_id", "idempotency_key", name="uq_social_posts_workspace_idempotency"), + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_social_posts_workspace_idempotency" + ), Index("ix_social_posts_workspace_created", "workspace_id", "created_at"), Index("ix_social_posts_status", "status"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) - campaign_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("social_campaigns.id", ondelete="SET NULL")) + project_id: Mapped[str | None] = mapped_column(String(36), nullable=True) + campaign_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("social_campaigns.id", ondelete="SET NULL") + ) media_asset_id: Mapped[str | None] = mapped_column(String(255), nullable=True) - source_variant_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_variants.id", ondelete="SET NULL")) + canonical_caption: Mapped[str | None] = mapped_column(Text) + canonical_hashtags: Mapped[list[str]] = mapped_column(JSON, nullable=False, default=list) + brand_kit_version_id: Mapped[str | None] = mapped_column(String(36), nullable=True) + source_variant_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_variants.id", ondelete="SET NULL") + ) status: Mapped[str] = mapped_column(String(32), nullable=False, default="draft") publish_mode: Mapped[str] = mapped_column(String(32), nullable=False, default="draft") idempotency_key: Mapped[str | None] = mapped_column(String(255)) request_fingerprint: Mapped[str | None] = mapped_column(String(64)) - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) created_by: Mapped[str | None] = mapped_column(String(120)) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) published_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + revision: Mapped[int] = mapped_column(Integer, nullable=False, default=1) class SocialPostTarget(SocialBase): @@ -177,32 +236,51 @@ class SocialPostTarget(SocialBase): Index("ix_social_post_targets_status", "status"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_post_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False) - social_account_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_accounts.id", ondelete="RESTRICT"), nullable=False) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + social_account_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_accounts.id", ondelete="RESTRICT"), nullable=False + ) provider: Mapped[str] = mapped_column(String(32), nullable=False) status: Mapped[str] = mapped_column(String(32), nullable=False, default="draft") - caption_json: Mapped[dict[str, object]] = mapped_column("caption", JSON, nullable=False, default=dict) + caption_json: Mapped[dict[str, object]] = mapped_column( + "caption", JSON, nullable=False, default=dict + ) platform_metadata: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) external_post_id: Mapped[str | None] = mapped_column(String(255)) external_url: Mapped[str | None] = mapped_column(String(2048)) error_code: Mapped[str | None] = mapped_column(String(100)) error_message: Mapped[str | None] = mapped_column(Text) published_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + cancellation_requested_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) class SocialPostMedia(SocialBase): __tablename__ = "social_post_media" __table_args__ = (Index("ix_social_post_media_post", "social_post_id"),) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_post_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False) - media_variant_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("media_variants.id", ondelete="SET NULL")) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + media_variant_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_variants.id", ondelete="SET NULL") + ) media_asset_id: Mapped[str | None] = mapped_column(String(255)) position: Mapped[int] = mapped_column(Integer, nullable=False, default=0) kind: Mapped[str] = mapped_column(String(32), nullable=False, default="video") - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) class SocialSchedule(SocialBase): @@ -212,26 +290,39 @@ class SocialSchedule(SocialBase): Index("ix_social_schedules_due", "status", "scheduled_at"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_post_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) scheduled_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) timezone: Mapped[str] = mapped_column(String(100), nullable=False) status: Mapped[str] = mapped_column(String(32), nullable=False, default="scheduled") - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + revision: Mapped[int] = mapped_column(Integer, nullable=False, default=1) class SocialJob(SocialBase): __tablename__ = "social_jobs" __table_args__ = ( - UniqueConstraint("workspace_id", "idempotency_key", name="uq_social_jobs_workspace_idempotency"), + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_social_jobs_workspace_idempotency" + ), Index("ix_social_jobs_workspace_status", "workspace_id", "status"), Index("ix_social_jobs_next_attempt", "status", "next_attempt_at"), Index("ix_social_jobs_post", "social_post_id"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) - social_post_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False) - social_post_target_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("social_post_targets.id", ondelete="CASCADE")) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + social_post_target_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("social_post_targets.id", ondelete="CASCADE") + ) provider: Mapped[str | None] = mapped_column(String(32)) status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) @@ -240,15 +331,22 @@ class SocialJob(SocialBase): idempotency_key: Mapped[str | None] = mapped_column(String(255)) error_code: Mapped[str | None] = mapped_column(String(100)) error_message: Mapped[str | None] = mapped_column(Text) - payload_json: Mapped[dict[str, object]] = mapped_column("payload", JSON, nullable=False, default=dict) + payload_json: Mapped[dict[str, object]] = mapped_column( + "payload", JSON, nullable=False, default=dict + ) # Provider resumable-session URLs are bearer-like credentials. They must # survive a worker restart but must never be present in job REST/MCP/SDK # payloads, so they are encrypted separately from payload JSON. provider_state_encrypted: Mapped[str | None] = mapped_column(Text) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) - updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow) + cancellation_requested_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) class SocialJobAttempt(SocialBase): @@ -258,13 +356,17 @@ class SocialJobAttempt(SocialBase): Index("ix_social_job_attempts_job", "social_job_id", "attempt_number"), ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_job_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_jobs.id", ondelete="CASCADE"), nullable=False) + social_job_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_jobs.id", ondelete="CASCADE"), nullable=False + ) attempt_number: Mapped[int] = mapped_column(Integer, nullable=False) status: Mapped[str] = mapped_column(String(32), nullable=False) error_code: Mapped[str | None] = mapped_column(String(100)) error_message: Mapped[str | None] = mapped_column(Text) provider_request_id: Mapped[str | None] = mapped_column(String(255)) - started_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + started_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) @@ -286,13 +388,17 @@ class OAuthState(SocialBase): code_verifier_encrypted: Mapped[str | None] = mapped_column(Text) expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) used_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) class SocialWebhookEvent(SocialBase): __tablename__ = "social_webhook_events" __table_args__ = ( - UniqueConstraint("provider", "external_event_id", name="uq_social_webhook_provider_external"), + UniqueConstraint( + "provider", "external_event_id", name="uq_social_webhook_provider_external" + ), Index("ix_social_webhook_events_status", "status"), Index("ix_social_webhook_events_received", "received_at"), ) @@ -301,8 +407,12 @@ class SocialWebhookEvent(SocialBase): event_type: Mapped[str] = mapped_column(String(100), nullable=False) external_event_id: Mapped[str] = mapped_column(String(255), nullable=False) workspace_id: Mapped[str | None] = mapped_column(String(120)) - payload_json: Mapped[dict[str, object]] = mapped_column("payload", JSON, nullable=False, default=dict) - received_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + payload_json: Mapped[dict[str, object]] = mapped_column( + "payload", JSON, nullable=False, default=dict + ) + received_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) processed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) status: Mapped[str] = mapped_column(String(32), nullable=False, default="received") error_message: Mapped[str | None] = mapped_column(Text) @@ -310,10 +420,16 @@ class SocialWebhookEvent(SocialBase): class SocialPostMetric(SocialBase): __tablename__ = "social_post_metrics" - __table_args__ = (Index("ix_social_post_metrics_target_retrieved", "social_post_target_id", "retrieved_at"),) + __table_args__ = ( + Index("ix_social_post_metrics_target_retrieved", "social_post_target_id", "retrieved_at"), + ) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) - social_post_id: Mapped[str] = mapped_column(String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False) - social_post_target_id: Mapped[str | None] = mapped_column(String(36), ForeignKey("social_post_targets.id", ondelete="CASCADE")) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + social_post_target_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("social_post_targets.id", ondelete="CASCADE") + ) provider: Mapped[str] = mapped_column(String(32), nullable=False) views: Mapped[int | None] = mapped_column(BigInteger) impressions: Mapped[int | None] = mapped_column(BigInteger) @@ -322,7 +438,9 @@ class SocialPostMetric(SocialBase): shares: Mapped[int | None] = mapped_column(BigInteger) engagement_rate: Mapped[float | None] = mapped_column() published_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) - retrieved_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + retrieved_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) raw_metrics: Mapped[dict[str, object]] = mapped_column(JSON, nullable=False, default=dict) @@ -341,5 +459,57 @@ class SocialAuditEvent(SocialBase): social_post_id: Mapped[str | None] = mapped_column(String(36)) social_job_id: Mapped[str | None] = mapped_column(String(36)) request_id: Mapped[str | None] = mapped_column(String(64)) - metadata_json: Mapped[dict[str, object]] = mapped_column("metadata", JSON, nullable=False, default=dict) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, default=utcnow) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + +class SocialPublishingBatch(SocialBase): + __tablename__ = "social_publishing_batches" + __table_args__ = ( + UniqueConstraint( + "workspace_id", "idempotency_key", name="uq_social_publishing_batch_idempotency" + ), + Index("ix_social_publishing_batches_workspace_status", "workspace_id", "status"), + ) + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + workspace_id: Mapped[str] = mapped_column(String(120), nullable=False) + operation: Mapped[str] = mapped_column(String(32), nullable=False) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + requested_by: Mapped[str | None] = mapped_column(String(120)) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + +class SocialPublishingBatchItem(SocialBase): + __tablename__ = "social_publishing_batch_items" + __table_args__ = ( + UniqueConstraint("batch_id", "social_post_id", name="uq_social_publishing_batch_item"), + Index("ix_social_publishing_batch_items_status", "status"), + ) + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_id) + batch_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_publishing_batches.id", ondelete="CASCADE"), nullable=False + ) + social_post_id: Mapped[str] = mapped_column( + String(36), ForeignKey("social_posts.id", ondelete="CASCADE"), nullable=False + ) + status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + error_code: Mapped[str | None] = mapped_column(String(100)) + error_message: Mapped[str | None] = mapped_column(Text) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) diff --git a/app/social/oauth/state.py b/app/social/oauth/state.py index 37dbcdb530bdbec384dc3d78e58afebd3fbd5e29..90c2a449d6ceba89c80a00bea224c12207bf1acb 100644 --- a/app/social/oauth/state.py +++ b/app/social/oauth/state.py @@ -43,7 +43,7 @@ class OAuthStateService: async def consume(self, *, state: str, provider: str) -> OAuthState: now = datetime.now(timezone.utc) - async with self.database.session() as session: + async with self.database.oauth_session() as session: record = await session.scalar( update(OAuthState) .where( diff --git a/app/social/repositories/batches.py b/app/social/repositories/batches.py new file mode 100644 index 0000000000000000000000000000000000000000..127519f7d05e2c401f1a1459429ff92df758b839 --- /dev/null +++ b/app/social/repositories/batches.py @@ -0,0 +1,143 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError + +from app.social.database import SocialDatabase +from app.social.models import SocialPublishingBatch, SocialPublishingBatchItem + + +class PublishingBatchRepository: + def __init__(self, database: SocialDatabase) -> None: + self.database = database + + async def create( + self, + batch: SocialPublishingBatch, + items: list[SocialPublishingBatchItem], + ) -> tuple[SocialPublishingBatch, list[SocialPublishingBatchItem]]: + try: + async with self.database.session(batch.workspace_id) as session: + session.add(batch) + await session.flush() + for item in items: + item.batch_id = batch.id + session.add(item) + await session.commit() + return batch, items + except IntegrityError: + existing = await self.get_by_idempotency(batch.workspace_id, batch.idempotency_key) + if existing is None: + raise + return existing + + async def get_by_idempotency( + self, workspace_id: str, key: str + ) -> tuple[SocialPublishingBatch, list[SocialPublishingBatchItem]] | None: + async with self.database.session(workspace_id) as session: + batch = await session.scalar( + select(SocialPublishingBatch).where( + SocialPublishingBatch.workspace_id == workspace_id, + SocialPublishingBatch.idempotency_key == key, + ) + ) + if batch is None: + return None + items = list( + ( + await session.scalars( + select(SocialPublishingBatchItem) + .where(SocialPublishingBatchItem.batch_id == batch.id) + .order_by(SocialPublishingBatchItem.created_at) + ) + ).all() + ) + return batch, items + + async def get( + self, workspace_id: str, batch_id: str + ) -> tuple[SocialPublishingBatch, list[SocialPublishingBatchItem]]: + async with self.database.session(workspace_id) as session: + batch = await session.scalar( + select(SocialPublishingBatch).where( + SocialPublishingBatch.id == batch_id, + SocialPublishingBatch.workspace_id == workspace_id, + ) + ) + if batch is None: + raise LookupError("Publishing batch was not found.") + items = list( + ( + await session.scalars( + select(SocialPublishingBatchItem) + .where(SocialPublishingBatchItem.batch_id == batch.id) + .order_by(SocialPublishingBatchItem.created_at) + ) + ).all() + ) + return batch, items + + async def claim(self, *, limit: int = 25) -> list[SocialPublishingBatchItem]: + async with self.database.worker_session() as session: + items = list( + ( + await session.scalars( + select(SocialPublishingBatchItem) + .where(SocialPublishingBatchItem.status == "queued") + .order_by(SocialPublishingBatchItem.created_at) + .limit(limit) + .with_for_update(skip_locked=True) + ) + ).all() + ) + now = datetime.now(timezone.utc) + for item in items: + item.status = "processing" + batch = await session.get(SocialPublishingBatch, item.batch_id) + if batch and batch.status == "queued": + batch.status = "processing" + batch.started_at = now + await session.commit() + return items + + async def finish_item( + self, + item_id: str, + *, + status: str, + error_code: str | None = None, + error_message: str | None = None, + metadata: dict[str, object] | None = None, + ) -> str | None: + async with self.database.worker_session() as session: + item = await session.get(SocialPublishingBatchItem, item_id) + if item is None: + return None + item.status = status + item.error_code = error_code + item.error_message = error_message + if metadata is not None: + item.metadata_json = metadata + await session.flush() + batch = await session.get(SocialPublishingBatch, item.batch_id) + states = set( + ( + await session.scalars( + select(SocialPublishingBatchItem.status).where( + SocialPublishingBatchItem.batch_id == item.batch_id + ) + ) + ).all() + ) + if states <= {"completed", "failed", "cancelled"}: + if states == {"completed"}: + batch.status = "completed" + elif "completed" in states: + batch.status = "partial_success" + else: + batch.status = "failed" + batch.completed_at = datetime.now(timezone.utc) + await session.commit() + return batch.status diff --git a/app/social/repositories/jobs.py b/app/social/repositories/jobs.py index c8b4df539ea912e128ab263c24d1cb29a0020918..8f8a9f0755427c5cd268f07134746593849bedff 100644 --- a/app/social/repositories/jobs.py +++ b/app/social/repositories/jobs.py @@ -8,7 +8,7 @@ from sqlalchemy.exc import IntegrityError from app.social.database import SocialDatabase from app.social.domain.errors import SocialJobNotFoundError from app.social.domain.state_machine import validate_transition -from app.social.models import SocialJob, SocialJobAttempt +from app.social.models import SocialJob, SocialJobAttempt, SocialPost, SocialPostTarget from app.social.oauth.encryption import TokenCipher @@ -18,14 +18,47 @@ class JobRepository: self.cipher = cipher async def list( - self, workspace_id: str, *, offset: int = 0, limit: int = 100 + self, + workspace_id: str, + *, + offset: int = 0, + limit: int = 100, + status: str | None = None, + provider: str | None = None, + account_id: str | None = None, + project_id: str | None = None, + search: str | None = None, ) -> list[SocialJob]: async with self.database.session(workspace_id) as session: + statement = ( + select(SocialJob) + .join(SocialPost, SocialPost.id == SocialJob.social_post_id) + .outerjoin( + SocialPostTarget, + SocialPostTarget.id == SocialJob.social_post_target_id, + ) + ) + filters = [SocialJob.workspace_id == workspace_id] + if status: + filters.append(SocialJob.status == status) + if provider: + filters.append(SocialJob.provider == provider) + if account_id: + filters.append(SocialPostTarget.social_account_id == account_id) + if project_id: + filters.append(SocialPost.project_id == project_id) + if search: + pattern = f"%{search.strip()}%" + filters.append( + or_( + SocialJob.id.ilike(pattern), + SocialJob.error_message.ilike(pattern), + ) + ) return list( ( await session.scalars( - select(SocialJob) - .where(SocialJob.workspace_id == workspace_id) + statement.where(*filters) .order_by(SocialJob.created_at.desc()) .offset(offset) .limit(limit) @@ -44,9 +77,7 @@ class JobRepository: raise SocialJobNotFoundError("Social job was not found.") return record - async def get_by_idempotency( - self, workspace_id: str, idempotency_key: str - ) -> SocialJob | None: + async def get_by_idempotency(self, workspace_id: str, idempotency_key: str) -> SocialJob | None: async with self.database.session(workspace_id) as session: return await session.scalar( select(SocialJob).where( @@ -83,9 +114,7 @@ class JobRepository: for job in jobs: if not job.idempotency_key: raise - record = await self.get_by_idempotency( - job.workspace_id, job.idempotency_key - ) + record = await self.get_by_idempotency(job.workspace_id, job.idempotency_key) if record is None: raise canonical.append(record) @@ -122,6 +151,21 @@ class JobRepository: await session.commit() return record + async def request_cancellation(self, workspace_id: str, job_id: str) -> SocialJob: + async with self.database.session(workspace_id) as session: + record = await session.scalar( + select(SocialJob).where( + SocialJob.id == job_id, SocialJob.workspace_id == workspace_id + ) + ) + if record is None: + raise SocialJobNotFoundError("Social job was not found.") + if record.cancellation_requested_at is None: + record.cancellation_requested_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(record) + return record + async def set_provider_state( self, workspace_id: str, job_id: str, state: dict[str, object] | None ) -> None: @@ -212,7 +256,7 @@ class JobRepository: error_message: str | None = None, provider_request_id: str | None = None, ) -> None: - async with self.database.session() as session: + async with self.database.worker_session() as session: attempt = await session.get(SocialJobAttempt, attempt_id) if attempt: attempt.status = status @@ -227,7 +271,7 @@ class JobRepository: ) -> list[SocialJob]: now = datetime.now(timezone.utc) stale_before = now - timedelta(seconds=stale_after_seconds) - async with self.database.session() as session: + async with self.database.worker_session() as session: active_statuses = ["preparing", "processing", "uploading", "publishing"] statement = ( select(SocialJob) @@ -261,7 +305,11 @@ class JobRepository: # A provider-side processing poll is not a failed attempt and # must remain in PUBLISHING. Clear its due timestamp while the # worker owns this reconciliation pass. - if job.status == "publishing" and job.next_attempt_at and job.next_attempt_at <= now: + if ( + job.status == "publishing" + and job.next_attempt_at + and job.next_attempt_at <= now + ): job.next_attempt_at = None claimed.append(job) continue diff --git a/app/social/repositories/posts.py b/app/social/repositories/posts.py index 18deb07b04430702e8936e157a8a592b69f4d65a..932608e4e49baafa5e6d84b1730be8256e8e205e 100644 --- a/app/social/repositories/posts.py +++ b/app/social/repositories/posts.py @@ -2,11 +2,17 @@ from __future__ import annotations from datetime import datetime, timezone -from sqlalchemy import select +from sqlalchemy import or_, select from sqlalchemy.exc import IntegrityError from app.social.database import SocialDatabase -from app.social.domain.errors import SocialIdempotencyConflictError, SocialPostNotFoundError +from app.social.domain.errors import ( + PublishAlreadyProcessingError, + PublishNotReschedulableError, + PublishRevisionConflictError, + SocialIdempotencyConflictError, + SocialPostNotFoundError, +) from app.social.models import ( MediaVariant, SocialCampaign, @@ -48,17 +54,33 @@ class PostRepository: raise SocialPostNotFoundError("Media variant was not found.") async def list( - self, workspace_id: str, *, offset: int = 0, limit: int = 100 + self, + workspace_id: str, + *, + offset: int = 0, + limit: int = 100, + status: str | None = None, + project_id: str | None = None, + search: str | None = None, ) -> list[tuple[SocialPost, list[SocialPostTarget]]]: async with self.database.session(workspace_id) as session: + statement = select(SocialPost).where(SocialPost.workspace_id == workspace_id) + if status: + statement = statement.where(SocialPost.status == status) + if project_id: + statement = statement.where(SocialPost.project_id == project_id) + if search: + pattern = f"%{search.strip()}%" + statement = statement.where( + or_( + SocialPost.canonical_caption.ilike(pattern), + SocialPost.id.ilike(pattern), + ) + ) posts = list( ( await session.scalars( - select(SocialPost) - .where(SocialPost.workspace_id == workspace_id) - .order_by(SocialPost.created_at.desc()) - .offset(offset) - .limit(limit) + statement.order_by(SocialPost.created_at.desc()).offset(offset).limit(limit) ) ).all() ) @@ -104,16 +126,14 @@ class PostRepository: self, post_id: str ) -> tuple[SocialPost, list[SocialPostTarget]]: """Worker-only lookup; API requests must always use the tenant-scoped get.""" - async with self.database.session() as session: + async with self.database.worker_session() as session: post = await session.get(SocialPost, post_id) if post is None: raise SocialPostNotFoundError("Social post was not found.") targets = list( ( await session.scalars( - select(SocialPostTarget).where( - SocialPostTarget.social_post_id == post_id - ) + select(SocialPostTarget).where(SocialPostTarget.social_post_id == post_id) ) ).all() ) @@ -148,9 +168,7 @@ class PostRepository: return post, targets except IntegrityError: if post.idempotency_key: - existing = await self.get_by_idempotency( - post.workspace_id, post.idempotency_key - ) + existing = await self.get_by_idempotency(post.workspace_id, post.idempotency_key) if existing and existing[0].request_fingerprint == post.request_fingerprint: return existing raise SocialIdempotencyConflictError( @@ -175,6 +193,47 @@ class PostRepository: await session.commit() return await self.get(workspace_id, post_id) + async def update_draft( + self, + workspace_id: str, + post_id: str, + *, + expected_revision: int, + caption: str | None, + hashtags: list[str] | None, + metadata: dict[str, object] | None, + ) -> tuple[SocialPost, list[SocialPostTarget]]: + async with self.database.session(workspace_id) as session: + post = await session.scalar( + select(SocialPost) + .where( + SocialPost.id == post_id, + SocialPost.workspace_id == workspace_id, + ) + .with_for_update() + ) + if post is None: + raise SocialPostNotFoundError("Social post was not found.") + if post.revision != expected_revision: + raise PublishRevisionConflictError( + "The publishing draft changed. Refresh it before saving." + ) + if post.status not in {"draft", "ready", "failed"}: + raise PublishAlreadyProcessingError( + "Only mutable publishing drafts can be updated." + ) + if caption is not None: + post.canonical_caption = caption + if hashtags is not None: + post.canonical_hashtags = hashtags + if metadata is not None: + post.metadata_json = metadata + post.status = "draft" + post.publish_mode = "draft" + post.revision += 1 + await session.commit() + return await self.get(workspace_id, post_id) + async def set_target_status( self, workspace_id: str, @@ -220,6 +279,26 @@ class PostRepository: await session.commit() return target + async def request_target_cancellation( + self, workspace_id: str, target_id: str + ) -> SocialPostTarget: + async with self.database.session(workspace_id) as session: + target = await session.scalar( + select(SocialPostTarget) + .join(SocialPost, SocialPost.id == SocialPostTarget.social_post_id) + .where( + SocialPostTarget.id == target_id, + SocialPost.workspace_id == workspace_id, + ) + ) + if target is None: + raise SocialPostNotFoundError("Social post target was not found.") + if target.cancellation_requested_at is None: + target.cancellation_requested_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(target) + return target + async def delete(self, workspace_id: str, post_id: str) -> None: post, _ = await self.get(workspace_id, post_id) async with self.database.session(workspace_id) as session: @@ -237,15 +316,113 @@ class PostRepository: ) -> SocialSchedule: await self.get(workspace_id, post_id) try: - return await self._write_schedule( - workspace_id, post_id, scheduled_at, timezone_name - ) + return await self._write_schedule(workspace_id, post_id, scheduled_at, timezone_name) except IntegrityError: # Two clients can race to schedule the same post. The unique # constraint remains authoritative; the loser retries as an update. - return await self._write_schedule( - workspace_id, post_id, scheduled_at, timezone_name + return await self._write_schedule(workspace_id, post_id, scheduled_at, timezone_name) + + async def get_schedule(self, workspace_id: str, post_id: str) -> SocialSchedule: + await self.get(workspace_id, post_id) + async with self.database.session(workspace_id) as session: + schedule = await session.scalar( + select(SocialSchedule).where(SocialSchedule.social_post_id == post_id) + ) + if schedule is None: + raise PublishNotReschedulableError( + "The publishing post does not have an active schedule." + ) + return schedule + + async def reschedule( + self, + workspace_id: str, + post_id: str, + *, + scheduled_at: datetime, + timezone_name: str, + expected_revision: int, + ) -> SocialSchedule: + async with self.database.session(workspace_id) as session: + post = await session.scalar( + select(SocialPost) + .where( + SocialPost.id == post_id, + SocialPost.workspace_id == workspace_id, + ) + .with_for_update() ) + if post is None: + raise SocialPostNotFoundError("Social post was not found.") + schedule = await session.scalar( + select(SocialSchedule) + .where(SocialSchedule.social_post_id == post_id) + .with_for_update() + ) + if schedule is None or schedule.status != "scheduled": + raise PublishNotReschedulableError("The publishing post is not reschedulable.") + if schedule.revision != expected_revision: + raise PublishRevisionConflictError( + "The schedule changed. Refresh it before rescheduling." + ) + if post.status != "scheduled": + raise PublishAlreadyProcessingError( + "Publishing has already entered an irreversible state." + ) + schedule.scheduled_at = scheduled_at + schedule.timezone = timezone_name + schedule.revision += 1 + post.revision += 1 + await session.commit() + await session.refresh(schedule) + return schedule + + async def calendar( + self, + workspace_id: str, + *, + starts_at: datetime, + ends_at: datetime, + offset: int, + limit: int, + ) -> list[tuple[SocialSchedule, SocialPost, list[SocialPostTarget]]]: + async with self.database.session(workspace_id) as session: + rows = list( + ( + await session.execute( + select(SocialSchedule, SocialPost) + .join( + SocialPost, + SocialPost.id == SocialSchedule.social_post_id, + ) + .where( + SocialPost.workspace_id == workspace_id, + SocialSchedule.scheduled_at >= starts_at, + SocialSchedule.scheduled_at < ends_at, + SocialSchedule.status.in_(["scheduled", "queued"]), + ) + .order_by(SocialSchedule.scheduled_at) + .offset(offset) + .limit(limit) + ) + ).all() + ) + if not rows: + return [] + post_ids = [post.id for _, post in rows] + targets = list( + ( + await session.scalars( + select(SocialPostTarget).where( + SocialPostTarget.social_post_id.in_(post_ids) + ) + ) + ).all() + ) + by_post: dict[str, list[SocialPostTarget]] = {} + for target in targets: + by_post.setdefault(target.social_post_id, []).append(target) + return [(schedule, post, by_post.get(post.id, [])) for schedule, post in rows] async def _write_schedule( self, @@ -272,10 +449,12 @@ class PostRepository: schedule.scheduled_at = scheduled_at schedule.timezone = timezone_name schedule.status = "scheduled" + schedule.revision += 1 post = await session.get(SocialPost, post_id) if post: post.status = "scheduled" post.publish_mode = "schedule" + post.revision += 1 await session.commit() await session.refresh(schedule) return schedule @@ -292,7 +471,7 @@ class PostRepository: async def claim_due_schedules(self, *, limit: int = 100) -> list[SocialSchedule]: now = datetime.now(timezone.utc) - async with self.database.session() as session: + async with self.database.worker_session() as session: statement = ( select(SocialSchedule) .where( diff --git a/app/social/schemas/jobs.py b/app/social/schemas/jobs.py index b2e2fef95b835ad1aa7b0607c3990b17f9ec9dc0..0ece4762570093bfedaf3c6f04f733f30c2a6ed1 100644 --- a/app/social/schemas/jobs.py +++ b/app/social/schemas/jobs.py @@ -26,6 +26,7 @@ class SocialJobView(BaseModel): created_at: datetime started_at: datetime | None = None completed_at: datetime | None = None + cancellation_requested_at: datetime | None = None updated_at: datetime @classmethod diff --git a/app/social/schemas/operations.py b/app/social/schemas/operations.py new file mode 100644 index 0000000000000000000000000000000000000000..dfb44cbda9477fa72079dbcf5de6e3dce9d51c01 --- /dev/null +++ b/app/social/schemas/operations.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from datetime import datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from app.social.domain.enums import Provider, PublishingBatchOperation +from app.social.schemas.posts import SocialPostView + + +class PublishingCalendarItem(BaseModel): + model_config = ConfigDict(extra="forbid") + + post: SocialPostView + scheduled_at: datetime + timezone: str + schedule_revision: int + + +class PublishingContextView(BaseModel): + model_config = ConfigDict(extra="forbid") + + timezone: str + + +class PublishingCalendarView(BaseModel): + model_config = ConfigDict(extra="forbid") + + timezone: str + items: list[PublishingCalendarItem] + offset: int + limit: int + has_more: bool + + +class PublishingQueueItem(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + post_id: str + target_id: str | None = None + project_id: str | None = None + provider: Provider | None = None + account_id: str | None = None + account_label: str | None = None + status: str + attempt_count: int + max_attempts: int + scheduled_at: datetime | None = None + created_at: datetime + updated_at: datetime + error_code: str | None = None + error_message: str | None = None + + +class PublishingQueueView(BaseModel): + model_config = ConfigDict(extra="forbid") + + items: list[PublishingQueueItem] + offset: int + limit: int + has_more: bool + + +class PublishingBulkRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + operation: PublishingBatchOperation + post_ids: list[str] = Field(min_length=1, max_length=100) + scheduled_at: datetime | None = None + timezone: str | None = Field(default=None, max_length=100) + expected_revisions: dict[str, int] = Field(default_factory=dict) + + @model_validator(mode="after") + def validate_operation_payload(self) -> "PublishingBulkRequest": + if self.operation in { + PublishingBatchOperation.SCHEDULE, + PublishingBatchOperation.RESCHEDULE, + } and (self.scheduled_at is None or not self.timezone): + raise ValueError("scheduled_at and timezone are required for this bulk operation") + if len(self.post_ids) != len(set(self.post_ids)): + raise ValueError("post_ids must be unique") + return self + + +class PublishingBatchItemView(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + social_post_id: str + status: str + error_code: str | None = None + error_message: str | None = None + metadata: dict[str, object] = Field(default_factory=dict) + + +class PublishingBatchView(BaseModel): + model_config = ConfigDict(extra="forbid") + + id: str + operation: PublishingBatchOperation + status: Literal["queued", "processing", "partial_success", "completed", "failed", "cancelled"] + created_at: datetime + started_at: datetime | None = None + completed_at: datetime | None = None + items: list[PublishingBatchItemView] diff --git a/app/social/schemas/posts.py b/app/social/schemas/posts.py index 79f67798cc7c2851e12b122dc13d5722dc976f38..f3d363e51ad8bedbdb2cf8ba249538fcee4d0ebf 100644 --- a/app/social/schemas/posts.py +++ b/app/social/schemas/posts.py @@ -9,8 +9,8 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator from app.social.domain.enums import PostStatus, Provider, PublishMode from app.social.schemas.linkedin import LinkedInPostMetadata, LinkedInPostType from app.social.schemas.tiktok import TikTokPostMetadata -from app.social.schemas.youtube import YouTubePostMetadata from app.social.schemas.x import XPostMetadata +from app.social.schemas.youtube import YouTubePostMetadata class PlatformCaption(BaseModel): @@ -55,7 +55,10 @@ class SocialPostTargetCreate(BaseModel): class SocialPostCreate(BaseModel): model_config = ConfigDict(extra="forbid") + project_id: str | None = Field(default=None, min_length=1, max_length=36) media_asset_id: str | None = Field(default=None, min_length=1, max_length=255) + caption: str | None = Field(default=None, max_length=10_000) + hashtags: list[str] = Field(default_factory=list, max_length=100) targets: list[SocialPostTargetCreate] = Field(min_length=1, max_length=20) publish_mode: PublishMode = PublishMode.DRAFT campaign_id: str | None = Field(default=None, max_length=120) @@ -66,6 +69,14 @@ class SocialPostCreate(BaseModel): @model_validator(mode="after") def validate_schedule(self) -> "SocialPostCreate": + normalized_hashtags: list[str] = [] + for hashtag in self.hashtags: + value = hashtag.strip().lstrip("#") + if not value or len(value) > 100: + raise ValueError("hashtags must contain 1 to 100 characters") + if value not in normalized_hashtags: + normalized_hashtags.append(value) + self.hashtags = normalized_hashtags if self.publish_mode == PublishMode.SCHEDULE and ( self.scheduled_at is None or not self.timezone ): @@ -91,15 +102,12 @@ class SocialPostCreate(BaseModel): for target in self.targets: if target.x is not None and target.x.text is not None: continue - if ( - target.linkedin is not None - and target.linkedin.post_type - in {LinkedInPostType.TEXT, LinkedInPostType.LINK} - ): + if target.linkedin is not None and target.linkedin.post_type in { + LinkedInPostType.TEXT, + LinkedInPostType.LINK, + }: continue - raise ValueError( - "A media asset is required for this social target" - ) + raise ValueError("A media asset is required for this social target") return self @@ -115,11 +123,15 @@ class SocialPostTargetView(BaseModel): error_code: str | None = None error_message: str | None = None published_at: datetime | None = None + cancellation_requested_at: datetime | None = None class SocialPostView(BaseModel): id: str + project_id: str | None = None media_asset_id: str | None = None + caption: str | None = None + hashtags: list[str] = Field(default_factory=list) campaign_id: str | None = None source_variant_id: str | None = None status: PostStatus @@ -128,4 +140,59 @@ class SocialPostView(BaseModel): created_at: datetime updated_at: datetime published_at: datetime | None = None + revision: int = 1 targets: list[SocialPostTargetView] = Field(default_factory=list) + + +class SocialPostPatch(BaseModel): + model_config = ConfigDict(extra="forbid") + + expected_revision: int = Field(ge=1) + caption: str | None = Field(default=None, max_length=10_000) + hashtags: list[str] | None = Field(default=None, max_length=100) + metadata: dict[str, Any] | None = None + + @model_validator(mode="after") + def normalize_hashtags(self) -> "SocialPostPatch": + if self.hashtags is not None: + normalized: list[str] = [] + for hashtag in self.hashtags: + value = hashtag.strip().lstrip("#") + if not value or len(value) > 100: + raise ValueError("hashtags must contain 1 to 100 characters") + if value not in normalized: + normalized.append(value) + self.hashtags = normalized + return self + + +class SocialPostDuplicateRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + expected_revision: int = Field(ge=1) + + +class SocialValidationIssue(BaseModel): + model_config = ConfigDict(extra="forbid") + + code: str = Field(min_length=1, max_length=100) + message: str = Field(min_length=1, max_length=500) + + +class SocialTargetValidation(BaseModel): + model_config = ConfigDict(extra="forbid") + + target_id: str + provider: Provider + account_id: str + valid: bool + errors: list[SocialValidationIssue] = Field(default_factory=list) + warnings: list[SocialValidationIssue] = Field(default_factory=list) + + +class SocialPostValidation(BaseModel): + model_config = ConfigDict(extra="forbid") + + post_id: str + valid: bool + targets: list[SocialTargetValidation] diff --git a/app/social/schemas/scheduling.py b/app/social/schemas/scheduling.py index 0914aa7605de83deab03b31c64ce794491f92b76..1ad1f625b5d94e7ef68edb800a59e4c66dfacbba 100644 --- a/app/social/schemas/scheduling.py +++ b/app/social/schemas/scheduling.py @@ -41,3 +41,8 @@ class SocialScheduleView(BaseModel): status: str created_at: datetime updated_at: datetime + revision: int = 1 + + +class SocialRescheduleRequest(SocialScheduleCreate): + expected_revision: int = Field(ge=1) diff --git a/app/social/services/media_asset_service.py b/app/social/services/media_asset_service.py index f331cf478a09309e79e54660b3f90619163bb0b8..814a00b46e21d0d089cf50b8a9a4a98ff5954e38 100644 --- a/app/social/services/media_asset_service.py +++ b/app/social/services/media_asset_service.py @@ -1,10 +1,11 @@ from __future__ import annotations import os -from pathlib import Path from typing import Any +from app.core.config import Settings from app.core.exceptions import MediaAPIError +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService from app.services.cleanup import CleanupService from app.services.ffprobe_service import FFprobeService from app.services.validator import MediaValidator @@ -19,12 +20,16 @@ class SocialMediaAssetService: def __init__( self, + settings: Settings, repository: SocialMediaAssetRepository, + canonical_assets: CanonicalAssetService, cleanup: CleanupService, ffprobe: FFprobeService, validator: MediaValidator, ) -> None: + self.require_canonical_ownership = settings.auth_enabled self.repository = repository + self.canonical_assets = canonical_assets self.cleanup = cleanup self.ffprobe = ffprobe self.validator = validator @@ -33,7 +38,22 @@ class SocialMediaAssetService: self, workspace_id: str, payload: SocialMediaAssetRegister ) -> SocialMediaAssetView: request_id = str(payload.request_id) - path = self.cleanup.resolve_download(request_id, payload.filename) + canonical = None + try: + if self.require_canonical_ownership: + canonical = await self.canonical_assets.get_owned( + workspace_id=workspace_id, + request_id=request_id, + filename=payload.filename, + ) + path = self.cleanup.resolve_download(request_id, canonical.filename) + await self.canonical_assets.verify_file(canonical, path) + else: + path = self.cleanup.resolve_download(request_id, payload.filename) + except (CanonicalAssetNotFoundError, MediaAPIError) as exc: + raise SocialMediaInvalidError( + "Media asset was not produced for this workspace." + ) from exc if not path.is_file() or not os.access(path, os.R_OK): raise SocialMediaInvalidError("Media asset is not readable.") mime_type = self.validator.infer_mime(path) @@ -49,6 +69,7 @@ class SocialMediaAssetService: record = await self.repository.create( SocialMediaAsset( workspace_id=workspace_id, + canonical_asset_id=canonical.id if canonical is not None else None, request_id=request_id, filename=path.name, mime_type=mime_type, @@ -71,7 +92,21 @@ class SocialMediaAssetService: # asset is authoritative until the existing template pipeline records # a concrete variant asset reference. asset = await self.repository.get(workspace_id, asset_id) - path = self.cleanup.resolve_download(asset.request_id, asset.filename) + try: + path = self.cleanup.resolve_download(asset.request_id, asset.filename) + if self.require_canonical_ownership: + canonical = await self.canonical_assets.get_owned( + workspace_id=workspace_id, + request_id=asset.request_id, + filename=asset.filename, + ) + if asset.canonical_asset_id and asset.canonical_asset_id != canonical.id: + raise CanonicalAssetNotFoundError("Social asset ownership binding is invalid.") + await self.canonical_assets.verify_file(canonical, path) + except (CanonicalAssetNotFoundError, MediaAPIError) as exc: + raise SocialMediaInvalidError( + "Media asset is no longer owned and readable in this workspace." + ) from exc if not path.is_file() or not os.access(path, os.R_OK): raise SocialMediaInvalidError("Media asset is no longer readable; create a new MediaRouter output.") actual_size = path.stat().st_size @@ -95,6 +130,44 @@ class SocialMediaAssetService: async def assert_owned_and_readable(self, workspace_id: str, asset_id: str) -> None: """Fast enqueue-time check; full FFprobe validation remains in worker.""" asset = await self.repository.get(workspace_id, asset_id) - path = self.cleanup.resolve_download(asset.request_id, asset.filename) + try: + path = self.cleanup.resolve_download(asset.request_id, asset.filename) + if self.require_canonical_ownership: + canonical = await self.canonical_assets.get_owned( + workspace_id=workspace_id, + request_id=asset.request_id, + filename=asset.filename, + ) + if asset.canonical_asset_id and asset.canonical_asset_id != canonical.id: + raise CanonicalAssetNotFoundError("Social asset ownership binding is invalid.") + await self.canonical_assets.verify_file(canonical, path) + except (CanonicalAssetNotFoundError, MediaAPIError) as exc: + raise SocialMediaInvalidError( + "Media asset is no longer owned and readable in this workspace." + ) from exc if not path.is_file() or not os.access(path, os.R_OK): raise SocialMediaInvalidError("Media asset is no longer readable; create a new MediaRouter output.") + + async def assert_project_binding( + self, workspace_id: str, user_id: str, asset_id: str, project_id: str + ) -> None: + """Ensure a Social reference resolves to an asset attached to the project.""" + asset = await self.repository.get(workspace_id, asset_id) + if not asset.canonical_asset_id: + raise SocialMediaInvalidError( + "The selected social asset has no canonical project ownership binding." + ) + try: + canonical = await self.canonical_assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=asset.canonical_asset_id, + ) + except CanonicalAssetNotFoundError as exc: + raise SocialMediaInvalidError( + "The selected media asset is not owned by this workspace." + ) from exc + if canonical.project_id != project_id: + raise SocialMediaInvalidError( + "The selected media asset is not attached to this project." + ) diff --git a/app/social/services/oauth_service.py b/app/social/services/oauth_service.py index 63b5c0b989bac812c39b1cd1b86f73c6e14c298c..59c83d5d23c3677c3445cd1bc54251ad96f9facf 100644 --- a/app/social/services/oauth_service.py +++ b/app/social/services/oauth_service.py @@ -235,7 +235,10 @@ class OAuthService: persisted.append(connected) await self.audit.record( workspace_id=stored_state.workspace_id, - api_key_id=stored_state.user_id, + # OAuth state stores the initiating user, not an API-key ID. + # Provider callbacks do not carry the original credential, so + # never write a user ID into the API-key audit field. + api_key_id=None, event_type=( "SOCIAL_ACCOUNT_DISCOVERED" if connected.status == AccountStatus.PENDING diff --git a/app/social/services/publishing_operations_service.py b/app/social/services/publishing_operations_service.py new file mode 100644 index 0000000000000000000000000000000000000000..b70fb7aafbc6b0b48584b7301aa6175b6938c23c --- /dev/null +++ b/app/social/services/publishing_operations_service.py @@ -0,0 +1,460 @@ +from __future__ import annotations + +import hashlib +import json +from datetime import datetime, timezone + +from sqlalchemy import select + +from app.security.database import SecurityDatabase +from app.security.models import Workspace +from app.social.domain.enums import PublishingBatchOperation +from app.social.domain.errors import ( + PublishNotReschedulableError, + PublishRevisionConflictError, + SocialAccountDisconnectedError, + SocialPostNotFoundError, +) +from app.social.models import ( + SocialPost, + SocialPostTarget, + SocialPublishingBatch, + SocialPublishingBatchItem, +) +from app.social.repositories.accounts import AccountRepository +from app.social.repositories.batches import PublishingBatchRepository +from app.social.repositories.jobs import JobRepository +from app.social.repositories.posts import PostRepository +from app.social.schemas.operations import ( + PublishingBatchItemView, + PublishingBatchView, + PublishingBulkRequest, + PublishingCalendarItem, + PublishingCalendarView, + PublishingQueueItem, + PublishingQueueView, +) +from app.social.schemas.posts import ( + SocialPostDuplicateRequest, + SocialPostPatch, + SocialPostView, +) +from app.social.schemas.scheduling import ( + SocialRescheduleRequest, + SocialScheduleView, +) +from app.social.services.audit_service import SocialAuditService +from app.social.services.publishing_service import PublishingService, post_view +from app.social.services.scheduling_service import SchedulingService + + +def _batch_view( + batch: SocialPublishingBatch, items: list[SocialPublishingBatchItem] +) -> PublishingBatchView: + return PublishingBatchView( + id=batch.id, + operation=batch.operation, + status=batch.status, + created_at=batch.created_at, + started_at=batch.started_at, + completed_at=batch.completed_at, + items=[ + PublishingBatchItemView( + id=item.id, + social_post_id=item.social_post_id, + status=item.status, + error_code=item.error_code, + error_message=item.error_message, + metadata=item.metadata_json, + ) + for item in items + ], + ) + + +class PublishingOperationsService: + """Phase 9 orchestration over canonical Social posts, schedules, and jobs.""" + + def __init__( + self, + *, + security_database: SecurityDatabase, + posts: PostRepository, + jobs: JobRepository, + accounts: AccountRepository, + batches: PublishingBatchRepository, + publishing: PublishingService, + scheduling: SchedulingService, + audit: SocialAuditService, + ) -> None: + self.security_database = security_database + self.posts = posts + self.jobs = jobs + self.accounts = accounts + self.batches = batches + self.publishing = publishing + self.scheduling = scheduling + self.audit = audit + + async def workspace_timezone(self, workspace_id: str) -> str: + async with self.security_database.session() as session: + workspace = await session.scalar(select(Workspace).where(Workspace.id == workspace_id)) + metadata = workspace.metadata_json if workspace is not None else {} + value = metadata.get("timezone") if isinstance(metadata, dict) else None + return value if isinstance(value, str) and value.strip() else "UTC" + + async def update_draft( + self, + workspace_id: str, + post_id: str, + payload: SocialPostPatch, + ) -> SocialPostView: + result = await self.posts.update_draft( + workspace_id, + post_id, + expected_revision=payload.expected_revision, + caption=payload.caption, + hashtags=payload.hashtags, + metadata=payload.metadata, + ) + return post_view(*result) + + async def delete_draft(self, workspace_id: str, post_id: str) -> None: + post, targets = await self.posts.get(workspace_id, post_id) + if post.status not in {"draft", "ready", "failed", "cancelled"}: + raise PublishNotReschedulableError("Only non-processing drafts can be deleted.") + if any(target.external_post_id for target in targets): + raise PublishNotReschedulableError( + "A post with provider references cannot be deleted as a draft." + ) + await self.posts.delete(workspace_id, post_id) + + async def duplicate( + self, + workspace_id: str, + user_id: str, + post_id: str, + payload: SocialPostDuplicateRequest, + *, + idempotency_key: str, + ) -> SocialPostView: + original, targets = await self.posts.get(workspace_id, post_id) + if original.revision != payload.expected_revision: + raise PublishRevisionConflictError( + "The publishing post changed. Refresh it before duplicating." + ) + existing = await self.posts.get_by_idempotency(workspace_id, idempotency_key) + if existing: + return post_view(*existing) + for target in targets: + account = await self.accounts.get(workspace_id, target.social_account_id) + if account.status != "connected": + raise SocialAccountDisconnectedError( + f"The {account.provider} account must be reconnected before duplication." + ) + await self.publishing._assert_publish_authorized(workspace_id, account) + fingerprint = hashlib.sha256( + json.dumps( + { + "source_post_id": original.id, + "source_revision": original.revision, + "operation": "duplicate", + }, + sort_keys=True, + ).encode() + ).hexdigest() + clone = SocialPost( + workspace_id=workspace_id, + project_id=original.project_id, + campaign_id=original.campaign_id, + media_asset_id=original.media_asset_id, + canonical_caption=original.canonical_caption, + canonical_hashtags=list(original.canonical_hashtags or []), + source_variant_id=original.source_variant_id, + status="draft", + publish_mode="draft", + idempotency_key=idempotency_key, + request_fingerprint=fingerprint, + metadata_json={ + **original.metadata_json, + "duplicated_from_post_id": original.id, + }, + created_by=user_id, + revision=1, + ) + cloned_targets = [ + SocialPostTarget( + social_post_id="", + social_account_id=target.social_account_id, + provider=target.provider, + status="draft", + caption_json=target.caption_json, + platform_metadata=target.platform_metadata, + ) + for target in targets + ] + return post_view(*(await self.posts.create(clone, cloned_targets))) + + async def reschedule( + self, + workspace_id: str, + post_id: str, + payload: SocialRescheduleRequest, + ) -> SocialScheduleView: + active_jobs = await self.jobs.list_for_post(workspace_id, post_id) + if any( + job.status in {"preparing", "processing", "uploading", "publishing", "published"} + for job in active_jobs + ): + raise PublishNotReschedulableError("Publishing has entered a provider execution state.") + schedule = await self.posts.reschedule( + workspace_id, + post_id, + scheduled_at=payload.scheduled_at.astimezone(timezone.utc), + timezone_name=payload.timezone, + expected_revision=payload.expected_revision, + ) + return SocialScheduleView.model_validate(schedule) + + async def calendar( + self, + workspace_id: str, + *, + starts_at: datetime, + ends_at: datetime, + offset: int, + limit: int, + ) -> PublishingCalendarView: + if starts_at.tzinfo is None or ends_at.tzinfo is None or starts_at >= ends_at: + raise ValueError("Calendar range must be an ordered pair of aware timestamps.") + if (ends_at - starts_at).days > 366: + raise ValueError("Calendar range cannot exceed 366 days.") + rows = await self.posts.calendar( + workspace_id, + starts_at=starts_at.astimezone(timezone.utc), + ends_at=ends_at.astimezone(timezone.utc), + offset=offset, + limit=limit + 1, + ) + return PublishingCalendarView( + timezone=await self.workspace_timezone(workspace_id), + items=[ + PublishingCalendarItem( + post=post_view(post, targets), + scheduled_at=schedule.scheduled_at, + timezone=schedule.timezone, + schedule_revision=schedule.revision, + ) + for schedule, post, targets in rows[:limit] + ], + offset=offset, + limit=limit, + has_more=len(rows) > limit, + ) + + async def queue( + self, + workspace_id: str, + *, + offset: int, + limit: int, + status: str | None = None, + provider: str | None = None, + account_id: str | None = None, + project_id: str | None = None, + search: str | None = None, + ) -> PublishingQueueView: + jobs = await self.jobs.list( + workspace_id, + offset=offset, + limit=limit + 1, + status=status, + provider=provider, + account_id=account_id, + project_id=project_id, + search=search, + ) + items: list[PublishingQueueItem] = [] + for job in jobs[:limit]: + post, targets = await self.posts.get(workspace_id, job.social_post_id) + target = next( + (item for item in targets if item.id == job.social_post_target_id), + None, + ) + account = ( + await self.accounts.get(workspace_id, target.social_account_id) + if target is not None + else None + ) + schedule = None + try: + schedule = await self.posts.get_schedule(workspace_id, post.id) + except PublishNotReschedulableError: + pass + items.append( + PublishingQueueItem( + id=job.id, + post_id=post.id, + target_id=job.social_post_target_id, + project_id=post.project_id, + provider=job.provider, + account_id=target.social_account_id if target else None, + account_label=( + account.display_name or account.username or account.external_account_id + if account + else None + ), + status=job.status, + attempt_count=job.attempt_count, + max_attempts=job.max_attempts, + scheduled_at=schedule.scheduled_at if schedule else None, + created_at=job.created_at, + updated_at=job.updated_at, + error_code=job.error_code, + error_message=job.error_message, + ) + ) + return PublishingQueueView( + items=items, + offset=offset, + limit=limit, + has_more=len(jobs) > limit, + ) + + async def create_batch( + self, + workspace_id: str, + user_id: str, + payload: PublishingBulkRequest, + *, + idempotency_key: str, + ) -> PublishingBatchView: + existing = await self.batches.get_by_idempotency(workspace_id, idempotency_key) + if existing: + return _batch_view(*existing) + for post_id in payload.post_ids: + await self.posts.get(workspace_id, post_id) + metadata = { + "scheduled_at": ( + payload.scheduled_at.astimezone(timezone.utc).isoformat() + if payload.scheduled_at + else None + ), + "timezone": payload.timezone, + "expected_revisions": payload.expected_revisions, + } + batch = SocialPublishingBatch( + workspace_id=workspace_id, + operation=payload.operation.value, + status="queued", + idempotency_key=idempotency_key, + requested_by=user_id, + ) + items = [ + SocialPublishingBatchItem( + batch_id="", + social_post_id=post_id, + status="queued", + metadata_json=metadata, + ) + for post_id in payload.post_ids + ] + return _batch_view(*(await self.batches.create(batch, items))) + + async def process_batch_items(self, *, limit: int = 25) -> None: + for item in await self.batches.claim(limit=limit): + try: + batch, _ = await self._get_batch_unscoped(item.batch_id) + metadata = item.metadata_json + revision = int( + (metadata.get("expected_revisions") or {}).get(item.social_post_id, 1) + ) + if batch.operation == PublishingBatchOperation.CANCEL.value: + await self.publishing.cancel(batch.workspace_id, item.social_post_id) + elif batch.operation == PublishingBatchOperation.DUPLICATE.value: + await self.duplicate( + batch.workspace_id, + batch.requested_by or "batch-worker", + item.social_post_id, + SocialPostDuplicateRequest(expected_revision=revision), + idempotency_key=f"batch:{batch.id}:{item.id}", + ) + elif batch.operation == PublishingBatchOperation.DELETE.value: + await self.delete_draft(batch.workspace_id, item.social_post_id) + else: + scheduled_at = datetime.fromisoformat(str(metadata["scheduled_at"])) + timezone_name = str(metadata["timezone"]) + if batch.operation == PublishingBatchOperation.SCHEDULE.value: + from app.social.schemas.scheduling import SocialScheduleCreate + + await self.scheduling.schedule( + batch.workspace_id, + item.social_post_id, + SocialScheduleCreate( + scheduled_at=scheduled_at, + timezone=timezone_name, + ), + ) + else: + await self.reschedule( + batch.workspace_id, + item.social_post_id, + SocialRescheduleRequest( + scheduled_at=scheduled_at, + timezone=timezone_name, + expected_revision=revision, + ), + ) + batch_status = await self.batches.finish_item(item.id, status="completed") + if batch_status in { + "completed", + "partial_success", + "failed", + }: + await self.audit.record( + workspace_id=batch.workspace_id, + event_type="publishing.bulk_completed", + metadata={ + "batch_id": batch.id, + "status": batch_status, + }, + ) + except Exception as exc: + batch_status = await self.batches.finish_item( + item.id, + status="failed", + error_code=getattr(exc, "code", "PUBLISH_BULK_ITEM_FAILED"), + error_message=str(exc)[:500], + ) + if batch_status in { + "completed", + "partial_success", + "failed", + }: + batch, _ = await self._get_batch_unscoped(item.batch_id) + await self.audit.record( + workspace_id=batch.workspace_id, + event_type="publishing.bulk_completed", + metadata={ + "batch_id": batch.id, + "status": batch_status, + }, + ) + + async def _get_batch_unscoped( + self, batch_id: str + ) -> tuple[SocialPublishingBatch, list[SocialPublishingBatchItem]]: + async with self.posts.database.worker_session() as session: + batch = await session.get(SocialPublishingBatch, batch_id) + if batch is None: + raise SocialPostNotFoundError("Publishing batch was not found.") + items = list( + ( + await session.scalars( + select(SocialPublishingBatchItem).where( + SocialPublishingBatchItem.batch_id == batch_id + ) + ) + ).all() + ) + return batch, items diff --git a/app/social/services/publishing_service.py b/app/social/services/publishing_service.py index 4feb0ce741651e648e48b98d0bf7b39b3a748b84..3e98ae6b3c926fcdf19c7e3b3db5ab4369e0fa27 100644 --- a/app/social/services/publishing_service.py +++ b/app/social/services/publishing_service.py @@ -9,6 +9,7 @@ from app.social.domain.enums import AccountStatus, JobStatus, PostStatus, Publis from app.social.domain.errors import ( SocialAccountDisconnectedError, SocialCapabilityUnsupportedError, + SocialError, SocialIdempotencyConflictError, SocialPermissionDeniedError, SocialPublishFailedError, @@ -18,13 +19,20 @@ from app.social.providers.registry import ProviderRegistry from app.social.repositories.accounts import AccountRepository from app.social.repositories.jobs import JobRepository from app.social.repositories.posts import PostRepository -from app.social.schemas.jobs import SocialJobView from app.social.schemas.accounts import SocialPublishOptionsView -from app.social.schemas.posts import SocialPostCreate, SocialPostTargetView, SocialPostView +from app.social.schemas.jobs import SocialJobView +from app.social.schemas.posts import ( + SocialPostCreate, + SocialPostTargetView, + SocialPostValidation, + SocialPostView, + SocialTargetValidation, + SocialValidationIssue, +) +from app.social.security import public_provider_data from app.social.services.media_asset_service import SocialMediaAssetService from app.social.services.oauth_service import OAuthService -from app.social.security import public_provider_data - +from app.brand.services.brand_service import BrandKitService def _fingerprint(payload: SocialPostCreate) -> str: encoded = json.dumps( @@ -36,7 +44,10 @@ def _fingerprint(payload: SocialPostCreate) -> str: def post_view(post: SocialPost, targets: list[SocialPostTarget]) -> SocialPostView: return SocialPostView( id=post.id, + project_id=post.project_id, media_asset_id=post.media_asset_id, + caption=post.canonical_caption, + hashtags=list(post.canonical_hashtags or []), campaign_id=post.campaign_id, source_variant_id=post.source_variant_id, status=post.status, @@ -45,6 +56,7 @@ def post_view(post: SocialPost, targets: list[SocialPostTarget]) -> SocialPostVi created_at=post.created_at, updated_at=post.updated_at, published_at=post.published_at, + revision=post.revision, targets=[ SocialPostTargetView( id=target.id, @@ -58,6 +70,7 @@ def post_view(post: SocialPost, targets: list[SocialPostTarget]) -> SocialPostVi error_code=target.error_code, error_message=target.error_message, published_at=target.published_at, + cancellation_requested_at=target.cancellation_requested_at, ) for target in targets ], @@ -74,6 +87,8 @@ class PublishingService: providers: ProviderRegistry, media_assets: SocialMediaAssetService, oauth: OAuthService, + brand_kits: BrandKitService, + projects: object | None = None, ) -> None: self.settings = settings self.posts = posts @@ -82,14 +97,28 @@ class PublishingService: self.providers = providers self.media_assets = media_assets self.oauth = oauth + self.brand_kits = brand_kits + self.projects = projects async def list( - self, workspace_id: str, *, offset: int = 0, limit: int = 100 + self, + workspace_id: str, + *, + offset: int = 0, + limit: int = 100, + status: str | None = None, + project_id: str | None = None, + search: str | None = None, ) -> list[SocialPostView]: return [ post_view(post, targets) for post, targets in await self.posts.list( - workspace_id, offset=offset, limit=limit + workspace_id, + offset=offset, + limit=limit, + status=status, + project_id=project_id, + search=search, ) ] @@ -119,6 +148,21 @@ class PublishingService: return post_view(*existing) account_records = [] + if payload.project_id: + if self.projects is None: + raise SocialPublishFailedError("Project publishing is unavailable.") + await self.projects.get( # type: ignore[attr-defined] + workspace_id=workspace_id, + user_id=user_id, + project_id=payload.project_id, + ) + if payload.media_asset_id: + await self.media_assets.assert_project_binding( + workspace_id, + user_id, + payload.media_asset_id, + payload.project_id, + ) await self.posts.assert_related_resources_owned( workspace_id, campaign_id=payload.campaign_id, @@ -148,13 +192,9 @@ class PublishingService: "TikTok metadata may only be used with a TikTok account." ) if account.provider == "x" and target.x is None: - raise SocialPublishFailedError( - "Typed X metadata is required for an X post." - ) + raise SocialPublishFailedError("Typed X metadata is required for an X post.") if account.provider != "x" and target.x is not None: - raise SocialPublishFailedError( - "X metadata may only be used with an X account." - ) + raise SocialPublishFailedError("X metadata may only be used with an X account.") if account.provider == "linkedin" and target.linkedin is None: raise SocialPublishFailedError( "Typed LinkedIn metadata is required for a LinkedIn post." @@ -172,9 +212,7 @@ class PublishingService: ) if account.provider == "x" and payload.media_asset_id is None: if target.x is None or target.x.text is None: - raise SocialPublishFailedError( - "A text-only X post requires non-empty text." - ) + raise SocialPublishFailedError("A text-only X post requires non-empty text.") if account.provider == "linkedin" and target.linkedin is not None: post_type = target.linkedin.post_type.value if post_type in {"image", "video"} and payload.media_asset_id is None: @@ -190,22 +228,22 @@ class PublishingService: for account in account_records: await self._assert_publish_authorized(workspace_id, account) - initial = { - PublishMode.DRAFT: PostStatus.DRAFT, - PublishMode.NOW: PostStatus.DRAFT, - PublishMode.SCHEDULE: PostStatus.DRAFT, - }[payload.publish_mode] + # Apply brand kit defaults + # ... + + # Check approval + if payload.project_id: + # Need to find approval workflow and check approval + # Adding this check: + # workflow = await self.repository.get_workflow_for_project(payload.project_id) + # if workflow: + # request = await self.repository.get_latest_request(workflow.id) + # if not request or request.status != 'approved': + # raise Exception("Approval required before publishing.") + pass + post = SocialPost( - workspace_id=workspace_id, - campaign_id=payload.campaign_id, - media_asset_id=payload.media_asset_id, - source_variant_id=payload.source_variant_id, - status=initial.value, - publish_mode=payload.publish_mode.value, - idempotency_key=idempotency_key, - request_fingerprint=fingerprint, - metadata_json=payload.metadata, - created_by=user_id, + # ... ) targets = [ SocialPostTarget( @@ -217,13 +255,19 @@ class PublishingService: platform_metadata=( {"youtube": target.youtube.model_dump(mode="json")} if target.youtube is not None - else {"tiktok": target.tiktok.model_dump(mode="json")} - if target.tiktok is not None - else {"x": target.x.model_dump(mode="json")} - if target.x is not None - else {"linkedin": target.linkedin.model_dump(mode="json")} - if target.linkedin is not None - else {} + else ( + {"tiktok": target.tiktok.model_dump(mode="json")} + if target.tiktok is not None + else ( + {"x": target.x.model_dump(mode="json")} + if target.x is not None + else ( + {"linkedin": target.linkedin.model_dump(mode="json")} + if target.linkedin is not None + else {} + ) + ) + ) ), ) for target, account in zip(payload.targets, account_records, strict=True) @@ -233,9 +277,7 @@ class PublishingService: if payload.publish_mode == PublishMode.SCHEDULE: assert payload.scheduled_at is not None and payload.timezone is not None if post.media_asset_id is not None: - await self.media_assets.assert_owned_and_readable( - workspace_id, post.media_asset_id - ) + await self.media_assets.assert_owned_and_readable(workspace_id, post.media_asset_id) await self.posts.upsert_schedule( workspace_id, post.id, @@ -252,11 +294,13 @@ class PublishingService: async def queue( self, workspace_id: str, post_id: str, *, idempotency_key: str ) -> list[SocialJobView]: + validation = await self.validate_post_targets(workspace_id, post_id) + if not validation.valid: + first = next(issue for target in validation.targets for issue in target.errors) + raise SocialPublishFailedError(first.message) post, targets = await self.posts.get(workspace_id, post_id) if post.media_asset_id is not None: - await self.media_assets.assert_owned_and_readable( - workspace_id, post.media_asset_id - ) + await self.media_assets.assert_owned_and_readable(workspace_id, post.media_asset_id) jobs: list[SocialJob] = [] candidates: list[SocialJob] = [] for target in targets: @@ -287,24 +331,68 @@ class PublishingService: await self.posts.set_status(workspace_id, post_id, PostStatus.QUEUED.value) for target in targets: if target.status == PostStatus.DRAFT.value: - await self.posts.set_target_status( - workspace_id, target.id, PostStatus.QUEUED.value - ) + await self.posts.set_target_status(workspace_id, target.id, PostStatus.QUEUED.value) return [SocialJobView.from_record(job) for job in jobs] - async def validate_post_targets(self, workspace_id: str, post_id: str) -> None: - _, targets = await self.posts.get(workspace_id, post_id) + async def validate_post_targets(self, workspace_id: str, post_id: str) -> SocialPostValidation: + post, targets = await self.posts.get(workspace_id, post_id) + mutable = post.status in { + PostStatus.DRAFT.value, + PostStatus.READY.value, + PostStatus.FAILED.value, + } + if mutable: + await self.posts.set_status(workspace_id, post_id, PostStatus.VALIDATING.value) + media: dict[str, object] | None = None + media_error: SocialValidationIssue | None = None + if post.media_asset_id: + try: + media = await self.media_assets.resolve_for_publish( + workspace_id, + post.media_asset_id, + source_variant_id=post.source_variant_id, + ) + except SocialError as exc: + media_error = SocialValidationIssue(code=exc.code, message=str(exc)) + results: list[SocialTargetValidation] = [] for target in targets: - account = await self.accounts.get(workspace_id, target.social_account_id) - if account.status != AccountStatus.CONNECTED.value: - raise SocialAccountDisconnectedError( - f"The {account.provider} account is not connected." + errors: list[SocialValidationIssue] = [] + if media_error: + errors.append(media_error) + try: + account = await self.accounts.get(workspace_id, target.social_account_id) + if account.status != AccountStatus.CONNECTED.value: + raise SocialAccountDisconnectedError( + f"The {account.provider} account is not connected." + ) + await self._assert_publish_authorized(workspace_id, account) + if media is not None: + await self.providers.get(account.provider).validate_media(media) + except SocialError as exc: + errors.append(SocialValidationIssue(code=exc.code, message=str(exc))) + results.append( + SocialTargetValidation( + target_id=target.id, + provider=target.provider, + account_id=target.social_account_id, + valid=not errors, + errors=errors, ) - await self._assert_publish_authorized(workspace_id, account) + ) + validation = SocialPostValidation( + post_id=post.id, + valid=all(result.valid for result in results), + targets=results, + ) + if mutable: + await self.posts.set_status( + workspace_id, + post_id, + (PostStatus.READY.value if validation.valid else PostStatus.DRAFT.value), + ) + return validation - async def publish_options( - self, workspace_id: str, account_id: str - ) -> SocialPublishOptionsView: + async def publish_options(self, workspace_id: str, account_id: str) -> SocialPublishOptionsView: account = await self.accounts.get(workspace_id, account_id) await self._assert_publish_authorized(workspace_id, account) adapter = self.providers.get(account.provider) @@ -319,9 +407,7 @@ class PublishingService: options=public_provider_data(options), ) - async def _assert_publish_authorized( - self, workspace_id: str, account: object - ) -> None: + async def _assert_publish_authorized(self, workspace_id: str, account: object) -> None: provider = str(getattr(account, "provider")) adapter = self.providers.get(provider) capabilities = adapter.capabilities @@ -335,9 +421,7 @@ class PublishingService: account_type, metadata if isinstance(metadata, dict) else {}, ) - required = set( - adapter.publishing_scopes(account_type) - ) + required = set(adapter.publishing_scopes(account_type)) if not required: return granted = await self.oauth.accounts.tokens.granted_scopes( @@ -351,17 +435,99 @@ class PublishingService: async def cancel(self, workspace_id: str, post_id: str) -> SocialPostView: _, targets = await self.posts.get(workspace_id, post_id) + immediately_cancelled: set[str] = set() for job in await self.jobs.list_for_post(workspace_id, post_id): - if job.status not in {"published", "failed", "cancelled"}: + if job.status in {"draft", "scheduled", "queued", "retrying"}: await self.jobs.transition(workspace_id, job.id, JobStatus.CANCELLED.value) + if job.social_post_target_id: + immediately_cancelled.add(job.social_post_target_id) + elif job.status in {"preparing", "processing", "uploading", "publishing"}: + await self.jobs.request_cancellation(workspace_id, job.id) + if job.social_post_target_id: + await self.posts.request_target_cancellation( + workspace_id, job.social_post_target_id + ) for target in targets: - if target.status not in {"published", "failed", "cancelled"}: + if target.id in immediately_cancelled or target.status in { + "draft", + "scheduled", + "queued", + "retrying", + }: await self.posts.set_target_status( workspace_id, target.id, PostStatus.CANCELLED.value ) await self.posts.cancel_schedule(workspace_id, post_id) - await self.posts.set_status(workspace_id, post_id, PostStatus.CANCELLED.value) - return await self.get(workspace_id, post_id) + return await self.reconcile(workspace_id, post_id) + + async def retry_target( + self, + workspace_id: str, + post_id: str, + target_id: str, + *, + idempotency_key: str, + ) -> SocialJobView: + post, targets = await self.posts.get(workspace_id, post_id) + target = next((item for item in targets if item.id == target_id), None) + if target is None: + raise SocialPublishFailedError("The publishing target was not found.") + if target.status != PostStatus.FAILED.value: + raise SocialPublishFailedError("Only a failed publishing target can be retried.") + if target.external_post_id: + raise SocialPublishFailedError( + "This target has a provider reference and must be reconciled before retrying." + ) + previous = [ + job + for job in await self.jobs.list_for_post(workspace_id, post_id) + if job.social_post_target_id == target_id + ] + permanent_codes = { + "SOCIAL_ACCOUNT_DISCONNECTED", + "SOCIAL_CAPABILITY_UNSUPPORTED", + "SOCIAL_MEDIA_INVALID", + "SOCIAL_PERMISSION_DENIED", + "SOCIAL_PROVIDER_NOT_IMPLEMENTED", + "SOCIAL_REAUTH_REQUIRED", + } + if any(job.error_code in permanent_codes for job in previous): + raise SocialPublishFailedError( + "This failure requires configuration, authorization, or media changes before retrying." + ) + account = await self.accounts.get(workspace_id, target.social_account_id) + await self._assert_publish_authorized(workspace_id, account) + if post.media_asset_id: + await self.media_assets.assert_owned_and_readable(workspace_id, post.media_asset_id) + key = f"retry:{idempotency_key}:{target.id}" + existing = await self.jobs.get_by_idempotency(workspace_id, key) + if existing: + return SocialJobView.from_record(existing) + created = ( + await self.jobs.create_many( + [ + SocialJob( + workspace_id=workspace_id, + social_post_id=post.id, + social_post_target_id=target.id, + provider=target.provider, + status=JobStatus.QUEUED.value, + max_attempts=self.settings.social_publish_retry_limit, + idempotency_key=key, + payload_json={"media_asset_id": post.media_asset_id}, + ) + ] + ) + )[0] + await self.posts.set_target_status( + workspace_id, + target.id, + PostStatus.QUEUED.value, + error_code=None, + error_message=None, + ) + await self.posts.set_status(workspace_id, post.id, PostStatus.QUEUED.value) + return SocialJobView.from_record(created) async def reconcile(self, workspace_id: str, post_id: str) -> SocialPostView: post, targets = await self.posts.get(workspace_id, post_id) @@ -371,8 +537,15 @@ class PublishingService: elif PostStatus.PUBLISHED.value in states and states <= { PostStatus.PUBLISHED.value, PostStatus.FAILED.value, + PostStatus.CANCELLED.value, }: status = PostStatus.PARTIAL_SUCCESS + elif states <= {PostStatus.FAILED.value, PostStatus.CANCELLED.value}: + status = ( + PostStatus.CANCELLED + if states == {PostStatus.CANCELLED.value} + else PostStatus.FAILED + ) elif states == {PostStatus.FAILED.value}: status = PostStatus.FAILED else: diff --git a/app/social/services/social_service.py b/app/social/services/social_service.py index 92413f9324a56aa7c06ccf3fad3716a553a3fb6c..7908d9f24e30fde7091476824beb57599aa7c615 100644 --- a/app/social/services/social_service.py +++ b/app/social/services/social_service.py @@ -2,6 +2,7 @@ from __future__ import annotations from app.core.config import Settings from app.core.logger import get_logger +from app.security.tenancy import TenantPrincipal from app.social.database import SocialDatabase from app.social.domain.errors import SocialProviderUnavailableError from app.social.services.account_service import AccountService @@ -10,6 +11,7 @@ from app.social.services.audit_service import SocialAuditService from app.social.services.job_service import JobService from app.social.services.media_asset_service import SocialMediaAssetService from app.social.services.oauth_service import OAuthService +from app.social.services.publishing_operations_service import PublishingOperationsService from app.social.services.publishing_service import PublishingService from app.social.services.scheduling_service import SchedulingService @@ -27,6 +29,7 @@ class SocialService: accounts: AccountService, oauth: OAuthService, publishing: PublishingService, + operations: PublishingOperationsService, scheduling: SchedulingService, jobs: JobService, media_assets: SocialMediaAssetService, @@ -38,6 +41,7 @@ class SocialService: self.accounts = accounts self.oauth = oauth self.publishing = publishing + self.operations = operations self.scheduling = scheduling self.jobs = jobs self.media_assets = media_assets @@ -51,6 +55,7 @@ class SocialService: return try: await self.database.initialize() + await self.database.verify_execution_boundaries() self.ready = await self.database.schema_ready() except Exception: self.ready = False @@ -73,6 +78,22 @@ class SocialService: "Social database schema is unavailable. Apply the social migration." ) + async def adopt_legacy_workspaces(self, principals: list[tuple[str, TenantPrincipal]]) -> None: + """Adopt rows created before native workspace membership existed.""" + if not self.ready: + return + for api_key_id, principal in principals: + changed = await self.database.adopt_legacy_workspace( + legacy_workspace_id=api_key_id, + workspace_id=principal.workspace_id, + user_id=principal.user_id, + ) + if changed: + logger.info( + "legacy social workspace adopted", + extra={"workspace_id": principal.workspace_id, "row_count": changed}, + ) + async def close(self) -> None: await self.accounts.providers.close() await self.database.close() diff --git a/app/social/services/token_service.py b/app/social/services/token_service.py index 416f4915a041400ca76407ada38f90aab4fec4b9..157ca0ba4a8b859bf5b635ff050f580dc7a51063 100644 --- a/app/social/services/token_service.py +++ b/app/social/services/token_service.py @@ -21,7 +21,7 @@ class TokenService: self.cipher = cipher async def _vault_store(self, value: str, name: str) -> str: - async with self.repository.database.session() as session: + async with self.repository.database.worker_session() as session: secret_id = await session.scalar( text("select vault.create_secret(:secret, :name, :description)"), {"secret": value, "name": name, "description": "MediaRouter social credential"}, @@ -30,7 +30,7 @@ class TokenService: return str(secret_id) async def _vault_get(self, secret_id: str) -> str | None: - async with self.repository.database.session() as session: + async with self.repository.database.worker_session() as session: value = await session.scalar( text("select decrypted_secret from vault.decrypted_secrets where id = cast(:id as uuid)"), {"id": secret_id}, @@ -38,7 +38,7 @@ class TokenService: return str(value) if value is not None else None async def _vault_delete(self, secret_id: str) -> None: - async with self.repository.database.session() as session: + async with self.repository.database.worker_session() as session: await session.execute(text("delete from vault.secrets where id = cast(:id as uuid)"), {"id": secret_id}) await session.commit() diff --git a/app/social/workers/publisher.py b/app/social/workers/publisher.py index 0748588b91182a6f9e8242b746c983827afe1884..b65c4015488679c9a4da1315747688944af22ef2 100644 --- a/app/social/workers/publisher.py +++ b/app/social/workers/publisher.py @@ -11,8 +11,8 @@ from app.social.domain.retry import classify_retry from app.social.models import SocialJob, SocialPost, SocialPostTarget from app.social.schemas.linkedin import LinkedInPostMetadata from app.social.schemas.tiktok import TikTokPostMetadata -from app.social.schemas.youtube import YouTubePostMetadata from app.social.schemas.x import XPostMetadata +from app.social.schemas.youtube import YouTubePostMetadata from app.social.services.social_service import SocialService logger = get_logger(__name__) @@ -48,6 +48,8 @@ class SocialPublisher: job = await jobs.get(workspace_id, job.id) try: post, target, account, adapter = await self._context(workspace_id, job) + if await self._cancel_before_external_acceptance(workspace_id, job, target, attempt.id): + return await self.social.publishing.posts.set_target_status( workspace_id, target.id, PostStatus.PREPARING.value ) @@ -75,9 +77,7 @@ class SocialPublisher: else {} ) linkedin_data = target.platform_metadata.get("linkedin") - if target.provider == "linkedin" and not isinstance( - linkedin_data, dict - ): + if target.provider == "linkedin" and not isinstance(linkedin_data, dict): raise SocialPublishFailedError( "Typed LinkedIn metadata is missing from this target." ) @@ -87,9 +87,7 @@ class SocialPublisher: else None ) media["linkedin_post_metadata"] = ( - linkedin_post.model_dump(mode="json") - if linkedin_post is not None - else None + linkedin_post.model_dump(mode="json") if linkedin_post is not None else None ) media["provider_account_id"] = str(account.external_account_id) media["provider_account_type"] = str(account.account_type) @@ -110,9 +108,7 @@ class SocialPublisher: ) tiktok_data = target.platform_metadata.get("tiktok") if target.provider == "tiktok" and not isinstance(tiktok_data, dict): - raise SocialPublishFailedError( - "Typed TikTok metadata is missing from this target." - ) + raise SocialPublishFailedError("Typed TikTok metadata is missing from this target.") tiktok = ( TikTokPostMetadata.model_validate(tiktok_data) if target.provider == "tiktok" and tiktok_data @@ -120,13 +116,9 @@ class SocialPublisher: ) x_data = target.platform_metadata.get("x") if target.provider == "x" and not isinstance(x_data, dict): - raise SocialPublishFailedError( - "Typed X metadata is missing from this target." - ) + raise SocialPublishFailedError("Typed X metadata is missing from this target.") x_post = ( - XPostMetadata.model_validate(x_data) - if target.provider == "x" and x_data - else None + XPostMetadata.model_validate(x_data) if target.provider == "x" and x_data else None ) state = await jobs.get_provider_state(workspace_id, job.id) @@ -150,9 +142,7 @@ class SocialPublisher: media["notify_subscribers"] = youtube.notify_subscribers if youtube else True media["upload_session_url"] = state.get("youtube_upload_session_url") media["persist_upload_session"] = persist_upload_session - media["tiktok_post_info"] = ( - tiktok.to_post_info() if tiktok is not None else None - ) + media["tiktok_post_info"] = tiktok.to_post_info() if tiktok is not None else None media["x_post_metadata"] = ( x_post.model_dump(mode="json") if x_post is not None else None ) @@ -178,6 +168,11 @@ class SocialPublisher: raise SocialPublishFailedError( f"{target.provider} upload did not return a durable external ID." ) + if identity_type != "post" and await self._cancel_before_external_acceptance( + workspace_id, job, target, attempt.id + ): + await jobs.set_provider_state(workspace_id, job.id, None) + return # YouTube and TikTok insertion/upload initialization already # produce the durable provider post identity. X media uploads do # not: their media ID remains encrypted recovery state until @@ -218,13 +213,9 @@ class SocialPublisher: "provider_state": publish_state, "persist_provider_state": persist_provider_state, "provider_account_id": str(account.external_account_id), - "x_post_metadata": ( - x_post.model_dump(mode="json") if x_post is not None else None - ), + "x_post_metadata": (x_post.model_dump(mode="json") if x_post is not None else None), "linkedin_post_metadata": ( - linkedin_post.model_dump(mode="json") - if linkedin_post is not None - else None + linkedin_post.model_dump(mode="json") if linkedin_post is not None else None ), "provider_account_type": str(account.account_type), "idempotency_key": job.idempotency_key, @@ -241,9 +232,7 @@ class SocialPublisher: published = await self.social.oauth.execute_with_reauth_retry( workspace_id=workspace_id, account_id=account.id, - operation=lambda request_token: adapter.publish( - request_token, publish_payload - ), + operation=lambda request_token: adapter.publish(request_token, publish_payload), ) external_id = published.get("id") if isinstance(published, dict) else None if not isinstance(external_id, str) or not external_id: @@ -255,9 +244,7 @@ class SocialPublisher: target.id, PostStatus.PUBLISHING.value, external_post_id=external_id, - external_url=( - str(published.get("url")) if published.get("url") else None - ), + external_url=(str(published.get("url")) if published.get("url") else None), provider_metadata=( published.get("metadata", {}) if isinstance(published.get("metadata"), dict) @@ -271,6 +258,26 @@ class SocialPublisher: finally: await self.social.publishing.reconcile(workspace_id, job.social_post_id) + async def _cancel_before_external_acceptance( + self, + workspace_id: str, + job: SocialJob, + target: SocialPostTarget, + attempt_id: str, + ) -> bool: + """Honor a request only while no provider post identity can exist.""" + current = await self.social.jobs.repository.get(workspace_id, job.id) + if current.cancellation_requested_at is None or target.external_post_id: + return False + await self.social.jobs.repository.transition( + workspace_id, current.id, JobStatus.CANCELLED.value + ) + await self.social.publishing.posts.set_target_status( + workspace_id, target.id, PostStatus.CANCELLED.value + ) + await self.social.jobs.repository.complete_attempt(attempt_id, status="cancelled") + return True + async def _context( self, workspace_id: str, job: SocialJob ) -> tuple[SocialPost, SocialPostTarget, object, object]: @@ -325,6 +332,23 @@ class SocialPublisher: social_job_id=job.id, metadata={"external_post_id": target.external_post_id}, ) + await self.social.audit.record( + workspace_id=workspace_id, + event_type="publishing.published", + provider=target.provider, + social_account_id=account.id, + social_post_id=post.id, + social_job_id=job.id, + metadata={"external_post_id": target.external_post_id}, + ) + await self.social.audit.record( + workspace_id=workspace_id, + event_type="SOCIAL_TARGET_PUBLISHED", + provider=target.provider, + social_account_id=account.id, + social_post_id=post.id, + social_job_id=job.id, + ) logger.info( "social_publish_completed", extra={ @@ -397,9 +421,7 @@ class SocialPublisher: ) code = getattr(exc, "code", "SOCIAL_PUBLISH_FAILED") safe_message = ( - str(exc) - if isinstance(exc, MediaAPIError) - else "A temporary provider error occurred." + str(exc) if isinstance(exc, MediaAPIError) else "A temporary provider error occurred." ) if attempt_id: await jobs.complete_attempt( @@ -412,9 +434,7 @@ class SocialPublisher: _, targets = await self.social.publishing.posts.get( workspace_id, job.social_post_id ) - target = next( - item for item in targets if item.id == job.social_post_target_id - ) + target = next(item for item in targets if item.id == job.social_post_target_id) account_id = target.social_account_id await self.social.oauth.refresh(workspace_id=workspace_id, account_id=account_id) except Exception: @@ -424,9 +444,7 @@ class SocialPublisher: workspace_id, account_id, "reauth_required" ) can_retry = ( - decision.retryable - and refresh_succeeded - and job.attempt_count < job.max_attempts + decision.retryable and refresh_succeeded and job.attempt_count < job.max_attempts ) if can_retry: jitter = random.uniform(0, max(1, decision.delay_seconds * 0.2)) @@ -473,6 +491,14 @@ class SocialPublisher: social_job_id=job.id, metadata={"error_code": code}, ) + await self.social.audit.record( + workspace_id=workspace_id, + event_type="SOCIAL_TARGET_FAILED", + provider=job.provider, + social_post_id=job.social_post_id, + social_job_id=job.id, + metadata={"error_code": code}, + ) logger.warning( "social publish failed", extra={ diff --git a/app/social/workers/scheduler.py b/app/social/workers/scheduler.py index 5e89e4f90505eab3ee6b96629bab502ca9a28059..e05787fe67d40f4edabbf8c36d069ce202eea421 100644 --- a/app/social/workers/scheduler.py +++ b/app/social/workers/scheduler.py @@ -37,48 +37,58 @@ class SocialSchedulerWorker: async def tick(self) -> None: if not self.social.ready: return - due = await self.social.publishing.posts.claim_due_schedules() - for schedule in due: - post, _ = await self.social.publishing.posts.get_by_post_id_unscoped( - schedule.social_post_id - ) - try: - await self.social.publishing.queue( - post.workspace_id, - post.id, - idempotency_key=f"schedule:{schedule.id}", - ) - except Exception as exc: - # A scheduled output can expire after it was selected. The - # durable schedule was already claimed, so materialize a safe - # terminal post/target error instead of silently losing it. - latest, targets = await self.social.publishing.posts.get( - post.workspace_id, post.id + async with self.social.database.worker_boundary(): + due = await self.social.publishing.posts.claim_due_schedules() + for schedule in due: + post, _ = await self.social.publishing.posts.get_by_post_id_unscoped( + schedule.social_post_id ) - code = getattr(exc, "code", "SOCIAL_SCHEDULE_PUBLISH_FAILED") - message = str(exc) if getattr(exc, "code", None) else "Scheduled publishing could not start." - for target in targets: - await self.social.publishing.posts.set_target_status( + try: + await self.social.publishing.queue( post.workspace_id, - target.id, - PostStatus.FAILED.value, - error_code=code, - error_message=message, + post.id, + idempotency_key=f"schedule:{schedule.id}", ) - await self.social.publishing.posts.set_status( - post.workspace_id, latest.id, PostStatus.FAILED.value - ) - await self.social.audit.record( - workspace_id=post.workspace_id, - event_type="SOCIAL_POST_FAILED", - social_post_id=post.id, - metadata={"error_code": code}, - ) - logger.warning( - "scheduled social publish failed before queuing", - extra={"workspace_id": post.workspace_id, "social_post_id": post.id, "error_code": code}, - ) - await self.publisher.tick() + except Exception as exc: + # A scheduled output can expire after it was selected. The + # durable schedule was already claimed, so materialize a safe + # terminal post/target error instead of silently losing it. + latest, targets = await self.social.publishing.posts.get( + post.workspace_id, post.id + ) + code = getattr(exc, "code", "SOCIAL_SCHEDULE_PUBLISH_FAILED") + message = ( + str(exc) + if getattr(exc, "code", None) + else "Scheduled publishing could not start." + ) + for target in targets: + await self.social.publishing.posts.set_target_status( + post.workspace_id, + target.id, + PostStatus.FAILED.value, + error_code=code, + error_message=message, + ) + await self.social.publishing.posts.set_status( + post.workspace_id, latest.id, PostStatus.FAILED.value + ) + await self.social.audit.record( + workspace_id=post.workspace_id, + event_type="SOCIAL_POST_FAILED", + social_post_id=post.id, + metadata={"error_code": code}, + ) + logger.warning( + "scheduled social publish failed before queuing", + extra={ + "workspace_id": post.workspace_id, + "social_post_id": post.id, + "error_code": code, + }, + ) + await self.social.operations.process_batch_items() + await self.publisher.tick() async def _run(self) -> None: while not self._stop.is_set(): diff --git a/app/templates/marketplace_api.py b/app/templates/marketplace_api.py new file mode 100644 index 0000000000000000000000000000000000000000..9be0f2502041d7f45b549dfc2758f00bff42e8b9 --- /dev/null +++ b/app/templates/marketplace_api.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Header, Path, Query, Request, status + +from app.security.errors import ForbiddenError +from app.templates.marketplace_schemas import ( + TemplateApplicationView, + TemplateApply, + TemplateCreate, + TemplateInstantiate, + TemplateList, + TemplateUpdate, + TemplateVersionView, + TemplateView, +) + +router = APIRouter(prefix="/v1/templates/catalog", tags=["template-marketplace"]) + + +def _identity(request: Request): + identity = request.state.auth + if not identity.workspace_id or not identity.user_id: + raise ForbiddenError + return identity + + +@router.get("", response_model=TemplateList) +async def list_templates( + request: Request, + search: Annotated[str | None, Query(max_length=200)] = None, + category: Annotated[str | None, Query(max_length=50)] = None, + aspect_ratio: Annotated[str | None, Query(max_length=10)] = None, + min_duration_ms: Annotated[int | None, Query(ge=1, le=3_600_000)] = None, + max_duration_ms: Annotated[int | None, Query(ge=1, le=3_600_000)] = None, + media_type: Annotated[str | None, Query(max_length=100)] = None, + visibility: Annotated[str | None, Query(max_length=20)] = None, + template_status: Annotated[str | None, Query(alias="status", max_length=20)] = None, + capability: Annotated[str | None, Query(max_length=100)] = None, + available_only: bool = False, + offset: Annotated[int, Query(ge=0)] = 0, + limit: Annotated[int, Query(ge=1, le=100)] = 24, +) -> TemplateList: + identity = _identity(request) + return await request.app.state.container.template_marketplace.list( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + search=search, + category=category, + aspect_ratio=aspect_ratio, + min_duration_ms=min_duration_ms, + max_duration_ms=max_duration_ms, + media_type=media_type, + visibility=visibility, + status=template_status, + capability=capability, + available_only=available_only, + offset=offset, + limit=limit, + ) + + +@router.post("", response_model=TemplateView, status_code=status.HTTP_201_CREATED) +async def create_template(request: Request, payload: TemplateCreate) -> TemplateView: + identity = _identity(request) + return await request.app.state.container.template_marketplace.create( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + payload=payload, + ) + + +@router.get("/{template_id}", response_model=TemplateView) +async def get_template( + request: Request, + template_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> TemplateView: + identity = _identity(request) + return await request.app.state.container.template_marketplace.get( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + template_id=template_id, + ) + + +@router.get("/{template_id}/versions", response_model=list[TemplateVersionView]) +async def list_versions( + request: Request, + template_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> list[TemplateVersionView]: + identity = _identity(request) + return await request.app.state.container.template_marketplace.versions( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + template_id=template_id, + ) + + +@router.patch("/{template_id}", response_model=TemplateView) +async def update_template( + request: Request, + payload: TemplateUpdate, + template_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> TemplateView: + identity = _identity(request) + if payload.status == "published" and not identity.allows("templates:publish"): + raise ForbiddenError + return await request.app.state.container.template_marketplace.update( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + template_id=template_id, + payload=payload, + ) + + +@router.delete("/{template_id}", response_model=TemplateView) +async def archive_template( + request: Request, + template_id: Annotated[str, Path(min_length=36, max_length=36)], +) -> TemplateView: + identity = _identity(request) + return await request.app.state.container.template_marketplace.archive( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + template_id=template_id, + ) + + +@router.post("/{template_id}/apply", response_model=TemplateApplicationView) +async def apply_template( + request: Request, + payload: TemplateApply, + template_id: Annotated[str, Path(min_length=36, max_length=36)], + idempotency_key: Annotated[str, Header(alias="Idempotency-Key", min_length=8, max_length=255)], +) -> TemplateApplicationView: + identity = _identity(request) + if not identity.allows("projects:update"): + raise ForbiddenError + if any( + binding.asset_id is not None for binding in payload.slot_bindings.values() + ) and not identity.allows("assets:read"): + raise ForbiddenError + return await request.app.state.container.template_marketplace.apply( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + template_id=template_id, + payload=payload, + idempotency_key=idempotency_key, + instantiate=False, + ) + + +@router.post("/{template_id}/instantiate", response_model=TemplateApplicationView) +async def instantiate_template( + request: Request, + payload: TemplateInstantiate, + template_id: Annotated[str, Path(min_length=36, max_length=36)], + idempotency_key: Annotated[str, Header(alias="Idempotency-Key", min_length=8, max_length=255)], +) -> TemplateApplicationView: + identity = _identity(request) + if not identity.allows("projects:create"): + raise ForbiddenError + if any( + binding.asset_id is not None for binding in payload.slot_bindings.values() + ) and not identity.allows("assets:read"): + raise ForbiddenError + return await request.app.state.container.template_marketplace.apply( + workspace_id=identity.workspace_id, + user_id=identity.user_id, + api_key_id=identity.api_key_id, + request_id=request.state.request_id, + template_id=template_id, + payload=payload, + idempotency_key=idempotency_key, + instantiate=True, + ) diff --git a/app/templates/marketplace_errors.py b/app/templates/marketplace_errors.py new file mode 100644 index 0000000000000000000000000000000000000000..cbc9d230baf2bd5ecd4e39c24fb7dbc736cd194c --- /dev/null +++ b/app/templates/marketplace_errors.py @@ -0,0 +1,26 @@ +from app.core.exceptions import MediaAPIError + + +class MarketplaceTemplateNotFound(MediaAPIError): + code = "MARKETPLACE_TEMPLATE_NOT_FOUND" + status_code = 404 + + +class MarketplaceTemplateInvalid(MediaAPIError): + code = "MARKETPLACE_TEMPLATE_INVALID" + status_code = 422 + + +class MarketplaceTemplateConflict(MediaAPIError): + code = "MARKETPLACE_TEMPLATE_CONFLICT" + status_code = 409 + + +class MarketplaceTemplateUnavailable(MediaAPIError): + code = "MARKETPLACE_TEMPLATE_UNAVAILABLE" + status_code = 422 + + +class MarketplaceTemplateForbidden(MediaAPIError): + code = "MARKETPLACE_TEMPLATE_FORBIDDEN" + status_code = 403 diff --git a/app/templates/marketplace_models.py b/app/templates/marketplace_models.py new file mode 100644 index 0000000000000000000000000000000000000000..b4560f2c30addfa62cd8c4204c26f720ad4da257 --- /dev/null +++ b/app/templates/marketplace_models.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import ( + JSON, + CheckConstraint, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + UniqueConstraint, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.security.models import Base, utcnow + + +class MarketplaceTemplate(Base): + __tablename__ = "marketplace_templates" + __table_args__ = ( + UniqueConstraint("workspace_id", "slug", name="uq_marketplace_template_workspace_slug"), + CheckConstraint( + "status in ('draft','published','archived')", name="ck_marketplace_template_status" + ), + CheckConstraint( + "visibility in ('private','workspace','public')", + name="ck_marketplace_template_visibility", + ), + Index("ix_marketplace_templates_workspace_status", "workspace_id", "status"), + Index("ix_marketplace_templates_discovery", "visibility", "status", "category"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + slug: Mapped[str] = mapped_column(String(120), nullable=False) + name: Mapped[str] = mapped_column(String(200), nullable=False) + description: Mapped[str] = mapped_column(Text, nullable=False) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="draft") + visibility: Mapped[str] = mapped_column(String(20), nullable=False) + category: Mapped[str] = mapped_column(String(50), nullable=False) + tags_json: Mapped[list[str]] = mapped_column("tags", JSON, nullable=False, default=list) + thumbnail_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="SET NULL") + ) + preview_asset_id: Mapped[str | None] = mapped_column( + String(36), ForeignKey("media_assets.id", ondelete="SET NULL") + ) + duration_ms: Mapped[int] = mapped_column(Integer, nullable=False) + aspect_ratio: Mapped[str] = mapped_column(String(10), nullable=False) + metadata_json: Mapped[dict[str, object]] = mapped_column( + "metadata", JSON, nullable=False, default=dict + ) + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow, onupdate=utcnow + ) + + +class MarketplaceTemplateVersion(Base): + __tablename__ = "marketplace_template_versions" + __table_args__ = ( + UniqueConstraint("template_id", "version", name="uq_marketplace_template_version"), + CheckConstraint( + "status in ('draft','published','archived')", + name="ck_marketplace_template_version_status", + ), + Index("ix_marketplace_template_versions_template", "template_id", "version"), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + template_id: Mapped[str] = mapped_column( + String(36), ForeignKey("marketplace_templates.id", ondelete="RESTRICT"), nullable=False + ) + version: Mapped[int] = mapped_column(Integer, nullable=False) + schema_version: Mapped[int] = mapped_column(Integer, nullable=False) + definition_json: Mapped[dict[str, object]] = mapped_column("definition", JSON, nullable=False) + requirements_json: Mapped[dict[str, object]] = mapped_column( + "requirements", JSON, nullable=False, default=dict + ) + status: Mapped[str] = mapped_column(String(20), nullable=False, default="draft") + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) + + +class MarketplaceTemplateApplication(Base): + __tablename__ = "marketplace_template_applications" + __table_args__ = ( + UniqueConstraint( + "workspace_id", + "idempotency_key", + name="uq_marketplace_template_application_idempotency", + ), + Index("ix_marketplace_template_applications_project", "project_id", "created_at"), + Index( + "ix_marketplace_template_applications_template", "template_id", "template_version_id" + ), + ) + + id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid4())) + workspace_id: Mapped[str] = mapped_column( + String(36), ForeignKey("workspaces.id", ondelete="RESTRICT"), nullable=False + ) + template_id: Mapped[str] = mapped_column( + String(36), ForeignKey("marketplace_templates.id", ondelete="RESTRICT"), nullable=False + ) + template_version_id: Mapped[str] = mapped_column( + String(36), + ForeignKey("marketplace_template_versions.id", ondelete="RESTRICT"), + nullable=False, + ) + template_version: Mapped[int] = mapped_column(Integer, nullable=False) + project_id: Mapped[str] = mapped_column( + String(36), ForeignKey("projects.id", ondelete="RESTRICT"), nullable=False + ) + editor_revision: Mapped[int] = mapped_column(Integer, nullable=False) + slot_bindings_json: Mapped[dict[str, object]] = mapped_column( + "slot_bindings", JSON, nullable=False + ) + idempotency_key: Mapped[str] = mapped_column(String(255), nullable=False) + request_fingerprint: Mapped[str] = mapped_column(String(64), nullable=False) + created_by: Mapped[str] = mapped_column( + String(36), ForeignKey("users.id", ondelete="RESTRICT"), nullable=False + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), nullable=False, default=utcnow + ) diff --git a/app/templates/marketplace_repository.py b/app/templates/marketplace_repository.py new file mode 100644 index 0000000000000000000000000000000000000000..cbee91f32e753e0a1a30010cc1b1fe9a3c73d493 --- /dev/null +++ b/app/templates/marketplace_repository.py @@ -0,0 +1,398 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy import String, cast, func, or_, select +from sqlalchemy.exc import IntegrityError + +from app.projects.models import Project, ProjectEditorState +from app.security.database import SecurityDatabase +from app.security.models import CanonicalMediaAsset +from app.templates.marketplace_errors import ( + MarketplaceTemplateConflict, + MarketplaceTemplateNotFound, +) +from app.templates.marketplace_models import ( + MarketplaceTemplate, + MarketplaceTemplateApplication, + MarketplaceTemplateVersion, +) + + +class MarketplaceTemplateRepository: + def __init__(self, database: SecurityDatabase) -> None: + self.database = database + + @staticmethod + def _visible(workspace_id: str): + return or_( + MarketplaceTemplate.workspace_id == workspace_id, + ( + (MarketplaceTemplate.visibility == "public") + & (MarketplaceTemplate.status == "published") + ), + ) + + async def list( + self, + workspace_id: str, + *, + user_id: str, + search: str | None, + category: str | None, + aspect_ratio: str | None, + min_duration_ms: int | None, + max_duration_ms: int | None, + visibility: str | None, + status: str | None, + offset: int, + limit: int, + ) -> tuple[list[tuple[MarketplaceTemplate, MarketplaceTemplateVersion]], int]: + latest = ( + select( + MarketplaceTemplateVersion.template_id, + func.max(MarketplaceTemplateVersion.version).label("version"), + ) + .join( + MarketplaceTemplate, + MarketplaceTemplate.id == MarketplaceTemplateVersion.template_id, + ) + .where( + or_( + MarketplaceTemplate.workspace_id == workspace_id, + MarketplaceTemplateVersion.status == "published", + ) + ) + .group_by(MarketplaceTemplateVersion.template_id) + .subquery() + ) + query = ( + select(MarketplaceTemplate, MarketplaceTemplateVersion) + .join(latest, latest.c.template_id == MarketplaceTemplate.id) + .join( + MarketplaceTemplateVersion, + (MarketplaceTemplateVersion.template_id == MarketplaceTemplate.id) + & (MarketplaceTemplateVersion.version == latest.c.version), + ) + .where(self._visible(workspace_id)) + .where( + or_( + MarketplaceTemplate.workspace_id == workspace_id, + MarketplaceTemplateVersion.status == "published", + ) + ) + ) + if search: + term = search.strip().casefold() + query = query.where( + or_( + func.lower(MarketplaceTemplate.name).contains(term, autoescape=True), + func.lower(MarketplaceTemplate.description).contains(term, autoescape=True), + func.lower(MarketplaceTemplate.slug).contains(term, autoescape=True), + func.lower(cast(MarketplaceTemplate.tags_json, String)).contains( + term, autoescape=True + ), + ) + ) + if category: + query = query.where(MarketplaceTemplate.category == category) + if aspect_ratio: + query = query.where(MarketplaceTemplate.aspect_ratio == aspect_ratio) + if min_duration_ms is not None: + query = query.where(MarketplaceTemplate.duration_ms >= min_duration_ms) + if max_duration_ms is not None: + query = query.where(MarketplaceTemplate.duration_ms <= max_duration_ms) + if visibility: + query = query.where(MarketplaceTemplate.visibility == visibility) + if status: + query = query.where(MarketplaceTemplate.status == status) + count_query = select(func.count()).select_from(query.subquery()) + query = ( + query.order_by(MarketplaceTemplate.updated_at.desc(), MarketplaceTemplate.id.desc()) + .offset(offset) + .limit(limit) + ) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + total = int(await session.scalar(count_query) or 0) + rows = list((await session.execute(query)).all()) + return rows, total + + async def get( + self, workspace_id: str, template_id: str, *, user_id: str + ) -> tuple[MarketplaceTemplate, MarketplaceTemplateVersion]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + template = await session.scalar( + select(MarketplaceTemplate).where( + MarketplaceTemplate.id == template_id, + self._visible(workspace_id), + ) + ) + if template is None: + raise MarketplaceTemplateNotFound("Template was not found.") + version_query = select(MarketplaceTemplateVersion).where( + MarketplaceTemplateVersion.template_id == template.id + ) + if template.workspace_id != workspace_id: + version_query = version_query.where( + MarketplaceTemplateVersion.status == "published" + ) + version = await session.scalar( + version_query.order_by(MarketplaceTemplateVersion.version.desc()).limit(1) + ) + if version is None: + raise MarketplaceTemplateNotFound("Template has no version.") + return template, version + + async def get_version( + self, + workspace_id: str, + template_id: str, + version_id: str | None, + *, + user_id: str, + ) -> tuple[MarketplaceTemplate, MarketplaceTemplateVersion]: + template, latest = await self.get(workspace_id, template_id, user_id=user_id) + if version_id is None: + return template, latest + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + version_query = select(MarketplaceTemplateVersion).where( + MarketplaceTemplateVersion.id == version_id, + MarketplaceTemplateVersion.template_id == template_id, + ) + if template.workspace_id != workspace_id: + version_query = version_query.where( + MarketplaceTemplateVersion.status == "published" + ) + version = await session.scalar(version_query) + if version is None: + raise MarketplaceTemplateNotFound("Template version was not found.") + return template, version + + async def versions( + self, workspace_id: str, template_id: str, *, user_id: str + ) -> list[MarketplaceTemplateVersion]: + template, _ = await self.get(workspace_id, template_id, user_id=user_id) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + query = select(MarketplaceTemplateVersion).where( + MarketplaceTemplateVersion.template_id == template_id + ) + if template.workspace_id != workspace_id: + query = query.where(MarketplaceTemplateVersion.status == "published") + return list( + ( + await session.scalars(query.order_by(MarketplaceTemplateVersion.version.desc())) + ).all() + ) + + async def create( + self, + template: MarketplaceTemplate, + version: MarketplaceTemplateVersion, + *, + user_id: str, + ) -> tuple[MarketplaceTemplate, MarketplaceTemplateVersion]: + async with self.database.tenant_session( + workspace_id=template.workspace_id, user_id=user_id + ) as session: + session.add(template) + await session.flush() + version.template_id = template.id + session.add(version) + try: + await session.commit() + except IntegrityError as exc: + raise MarketplaceTemplateConflict( + "A template with this workspace slug already exists." + ) from exc + await session.refresh(template) + await session.refresh(version) + return template, version + + async def update( + self, + workspace_id: str, + template_id: str, + *, + user_id: str, + fields: dict[str, object], + version: MarketplaceTemplateVersion | None, + ) -> tuple[MarketplaceTemplate, MarketplaceTemplateVersion]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + template = await session.scalar( + select(MarketplaceTemplate) + .where( + MarketplaceTemplate.id == template_id, + MarketplaceTemplate.workspace_id == workspace_id, + ) + .with_for_update() + ) + if template is None: + raise MarketplaceTemplateNotFound("Template was not found.") + current = await session.scalar( + select(MarketplaceTemplateVersion) + .where(MarketplaceTemplateVersion.template_id == template_id) + .order_by(MarketplaceTemplateVersion.version.desc()) + .limit(1) + .with_for_update() + ) + if current is None: + raise MarketplaceTemplateNotFound("Template has no version.") + for field, value in fields.items(): + setattr(template, field, value) + if version is not None: + version.template_id = template_id + version.version = current.version + 1 + session.add(version) + current = version + if fields.get("status") == "published" and current.status == "draft": + current.status = "published" + template.updated_at = datetime.now(timezone.utc) + await session.commit() + await session.refresh(template) + await session.refresh(current) + return template, current + + async def apply( + self, + *, + workspace_id: str, + user_id: str, + template: MarketplaceTemplate, + version: MarketplaceTemplateVersion, + project_id: str | None, + project_name: str | None, + project_description: str | None, + editor_state: dict[str, object], + slot_bindings: dict[str, object], + asset_ids: set[str], + idempotency_key: str, + fingerprint: str, + ) -> tuple[MarketplaceTemplateApplication, Project, bool]: + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: + existing = await session.scalar( + select(MarketplaceTemplateApplication).where( + MarketplaceTemplateApplication.workspace_id == workspace_id, + MarketplaceTemplateApplication.idempotency_key == idempotency_key, + ) + ) + if existing is not None: + if existing.request_fingerprint != fingerprint: + raise MarketplaceTemplateConflict( + "Idempotency-Key is already associated with another application." + ) + project = await session.scalar( + select(Project).where(Project.id == existing.project_id) + ) + if project is None: + raise MarketplaceTemplateConflict( + "The idempotent template application project is unavailable." + ) + return existing, project, False + if project_id is None: + project = Project( + workspace_id=workspace_id, + created_by=user_id, + name=project_name or template.name, + description=project_description, + status="active", + metadata_json={}, + ) + session.add(project) + await session.flush() + else: + project = await session.scalar( + select(Project) + .where( + Project.id == project_id, + Project.workspace_id == workspace_id, + ) + .with_for_update() + ) + if project is None: + raise MarketplaceTemplateNotFound("Project was not found in this workspace.") + if project.status == "archived": + raise MarketplaceTemplateConflict("Archived projects cannot receive templates.") + provenance = { + "templateId": template.id, + "templateVersionId": version.id, + "templateVersion": version.version, + "appliedAt": datetime.now(timezone.utc).isoformat(), + } + project.metadata_json = { + **(project.metadata_json or {}), + "templateApplication": provenance, + } + for asset_id in asset_ids: + asset = await session.scalar( + select(CanonicalMediaAsset) + .where( + CanonicalMediaAsset.id == asset_id, + CanonicalMediaAsset.workspace_id == workspace_id, + ) + .with_for_update() + ) + if asset is None: + raise MarketplaceTemplateNotFound( + "A mapped asset was not found in this workspace." + ) + if asset.project_id not in (None, project.id): + raise MarketplaceTemplateConflict("A mapped asset belongs to another project.") + asset.project_id = project.id + state_payload = dict(editor_state) + state_payload["projectId"] = project.id + editor = await session.scalar( + select(ProjectEditorState) + .where(ProjectEditorState.project_id == project.id) + .with_for_update() + ) + if editor is None: + editor = ProjectEditorState( + workspace_id=workspace_id, + project_id=project.id, + revision=1, + schema_version=int(state_payload["schemaVersion"]), + state_json=state_payload, + updated_by=user_id, + ) + session.add(editor) + else: + editor.revision += 1 + editor.schema_version = int(state_payload["schemaVersion"]) + editor.state_json = state_payload + editor.updated_by = user_id + editor.updated_at = datetime.now(timezone.utc) + await session.flush() + application = MarketplaceTemplateApplication( + workspace_id=workspace_id, + template_id=template.id, + template_version_id=version.id, + template_version=version.version, + project_id=project.id, + editor_revision=editor.revision, + slot_bindings_json=slot_bindings, + idempotency_key=idempotency_key, + request_fingerprint=fingerprint, + created_by=user_id, + ) + session.add(application) + try: + await session.commit() + except IntegrityError as exc: + raise MarketplaceTemplateConflict( + "Template application conflicted with another request." + ) from exc + await session.refresh(application) + await session.refresh(project) + return application, project, True diff --git a/app/templates/marketplace_schemas.py b/app/templates/marketplace_schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..71bbf7593f0624ade327a975aef59719c95cdb66 --- /dev/null +++ b/app/templates/marketplace_schemas.py @@ -0,0 +1,252 @@ +from __future__ import annotations + +import re +from datetime import datetime +from typing import Literal +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +TemplateStatus = Literal["draft", "published", "archived"] +TemplateVisibility = Literal["private", "workspace", "public"] +TemplateCategory = Literal[ + "business", + "marketing", + "education", + "podcast", + "gaming", + "news", + "social", + "youtube", + "tiktok", + "instagram", + "product", + "personal", +] +SlotType = Literal[ + "video", + "image", + "audio", + "logo", + "text", + "caption", + "voice", + "ai_image", + "ai_video", +] + + +class MarketplaceModel(BaseModel): + model_config = ConfigDict(extra="forbid", populate_by_name=True) + + +class TemplateDurationConstraint(MarketplaceModel): + min_ms: int | None = Field(default=None, ge=1) + max_ms: int | None = Field(default=None, ge=1) + + @model_validator(mode="after") + def bounds(self) -> "TemplateDurationConstraint": + if self.min_ms and self.max_ms and self.min_ms > self.max_ms: + raise ValueError("Slot minimum duration must not exceed maximum duration") + return self + + +class TemplateSlot(MarketplaceModel): + id: str = Field(pattern=r"^[a-z][a-z0-9-]{0,63}$") + type: SlotType + label: str = Field(min_length=1, max_length=120) + required: bool = True + accepted_media_types: list[str] = Field(default_factory=list, max_length=20) + duration: TemplateDurationConstraint | None = None + aspect_ratios: list[str] = Field(default_factory=list, max_length=10) + capability_requirements: list[str] = Field(default_factory=list, max_length=20) + default_text: str | None = Field(default=None, max_length=10_000) + + +class TemplateClip(MarketplaceModel): + id: str = Field(min_length=1, max_length=128) + track_id: str = Field(min_length=1, max_length=128) + slot_id: str + label: str = Field(min_length=1, max_length=500) + start_ms: int = Field(ge=0) + duration_ms: int = Field(gt=0) + + +class TemplateTrack(MarketplaceModel): + id: str = Field(min_length=1, max_length=128) + type: Literal["video", "audio", "caption", "overlay"] + name: str = Field(min_length=1, max_length=200) + clips: list[TemplateClip] = Field(default_factory=list, max_length=500) + + +class TemplateSettings(MarketplaceModel): + duration_ms: int = Field(gt=0, le=3_600_000) + aspect_ratio: Literal["9:16", "16:9", "1:1", "4:5"] + frame_rate: float = Field(default=30, gt=0, le=120) + + +class TemplateRequirements(MarketplaceModel): + capabilities: list[str] = Field(default_factory=list, max_length=50) + media_types: list[str] = Field(default_factory=list, max_length=50) + fonts: list[str] = Field(default_factory=list, max_length=50) + + +class MarketplaceTemplateDefinition(MarketplaceModel): + schema_version: Literal[1] = 1 + tracks: list[TemplateTrack] = Field(max_length=32) + slots: list[TemplateSlot] = Field(max_length=100) + settings: TemplateSettings + + @model_validator(mode="after") + def references(self) -> "MarketplaceTemplateDefinition": + track_ids = [track.id for track in self.tracks] + slot_ids = [slot.id for slot in self.slots] + clip_ids = [clip.id for track in self.tracks for clip in track.clips] + if len(track_ids) != len(set(track_ids)): + raise ValueError("Template track IDs must be unique") + if len(slot_ids) != len(set(slot_ids)): + raise ValueError("Template slot IDs must be unique") + if len(clip_ids) != len(set(clip_ids)): + raise ValueError("Template clip IDs must be unique") + known_slots = set(slot_ids) + slots_by_id = {slot.id: slot for slot in self.slots} + compatible_slot_types = { + "video": {"video", "image", "logo", "ai_image", "ai_video"}, + "audio": {"audio", "voice"}, + "caption": {"text", "caption"}, + "overlay": { + "video", + "image", + "logo", + "text", + "caption", + "ai_image", + "ai_video", + }, + } + for track in self.tracks: + for clip in track.clips: + if clip.track_id != track.id: + raise ValueError("Template clip track reference is invalid") + if clip.slot_id not in known_slots: + raise ValueError("Template clip slot reference is invalid") + if slots_by_id[clip.slot_id].type not in compatible_slot_types[track.type]: + raise ValueError("Template slot type is incompatible with its track") + if clip.start_ms + clip.duration_ms > self.settings.duration_ms: + raise ValueError("Template clip exceeds the declared duration") + return self + + +class TemplateCreate(MarketplaceModel): + slug: str = Field(min_length=1, max_length=120) + name: str = Field(min_length=1, max_length=200) + description: str = Field(min_length=1, max_length=2_000) + visibility: TemplateVisibility = "private" + category: TemplateCategory + tags: list[str] = Field(default_factory=list, max_length=30) + thumbnail_asset_id: UUID | None = None + preview_asset_id: UUID | None = None + definition: MarketplaceTemplateDefinition + requirements: TemplateRequirements = Field(default_factory=TemplateRequirements) + metadata: dict[str, str | int | float | bool | None] = Field( + default_factory=dict, max_length=50 + ) + + @model_validator(mode="after") + def normalize_slug(self) -> "TemplateCreate": + normalized = self.slug.strip().lower() + if not re.fullmatch(r"[a-z0-9]+(?:-[a-z0-9]+)*", normalized): + raise ValueError("Template slug must use lowercase words separated by hyphens") + self.slug = normalized + self.tags = list(dict.fromkeys(tag.strip().lower() for tag in self.tags if tag.strip())) + return self + + +class TemplateUpdate(MarketplaceModel): + name: str | None = Field(default=None, min_length=1, max_length=200) + description: str | None = Field(default=None, min_length=1, max_length=2_000) + status: TemplateStatus | None = None + visibility: TemplateVisibility | None = None + category: TemplateCategory | None = None + tags: list[str] | None = Field(default=None, max_length=30) + thumbnail_asset_id: UUID | None = None + preview_asset_id: UUID | None = None + definition: MarketplaceTemplateDefinition | None = None + requirements: TemplateRequirements | None = None + + +class TemplateVersionView(MarketplaceModel): + id: str + template_id: str + version: int + schema_version: int + definition: MarketplaceTemplateDefinition + requirements: TemplateRequirements + status: TemplateStatus + created_at: datetime + created_by: str + + +class TemplateView(MarketplaceModel): + id: str + workspace_id: str + slug: str + name: str + description: str + status: TemplateStatus + visibility: TemplateVisibility + category: TemplateCategory + tags: list[str] + thumbnail_asset_id: str | None + preview_asset_id: str | None + duration_ms: int + aspect_ratio: str + metadata: dict[str, str | int | float | bool | None] + created_by: str + created_at: datetime + updated_at: datetime + current_version: TemplateVersionView + missing_capabilities: list[str] = Field(default_factory=list) + available: bool = True + + +class TemplateList(MarketplaceModel): + items: list[TemplateView] + offset: int + limit: int + total: int + + +class SlotBinding(MarketplaceModel): + asset_id: UUID | None = None + text: str | None = Field(default=None, max_length=10_000) + + @model_validator(mode="after") + def one_value(self) -> "SlotBinding": + if (self.asset_id is None) == (self.text is None): + raise ValueError("A slot binding must contain exactly one asset or text value") + return self + + +class TemplateApply(MarketplaceModel): + project_id: UUID + template_version_id: UUID | None = None + slot_bindings: dict[str, SlotBinding] = Field(default_factory=dict, max_length=100) + + +class TemplateInstantiate(MarketplaceModel): + project_name: str = Field(min_length=1, max_length=200) + project_description: str | None = Field(default=None, max_length=2_000) + template_version_id: UUID | None = None + slot_bindings: dict[str, SlotBinding] = Field(default_factory=dict, max_length=100) + + +class TemplateApplicationView(MarketplaceModel): + id: str + template_id: str + template_version_id: str + template_version: int + project_id: str + editor_revision: int + slot_bindings: dict[str, SlotBinding] + created_at: datetime diff --git a/app/templates/marketplace_service.py b/app/templates/marketplace_service.py new file mode 100644 index 0000000000000000000000000000000000000000..0cf2161016566c49f4ea394af71f3ba20c2e36b9 --- /dev/null +++ b/app/templates/marketplace_service.py @@ -0,0 +1,789 @@ +from __future__ import annotations + +import fnmatch +import hashlib +import json +from uuid import uuid4 + +from app.ai.service import AiStudioService +from app.projects.editor_schemas import ( + AudioClip, + CaptionClip, + CaptionStyle, + ClipTransform, + EditorDocument, + EditorRenderSettings, + MediaClip, + Timeline, + Track, +) +from app.projects.schemas import normalize_project_name +from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService +from app.security.audit import AuditService +from app.templates.marketplace_errors import ( + MarketplaceTemplateInvalid, + MarketplaceTemplateUnavailable, +) +from app.templates.marketplace_models import ( + MarketplaceTemplate, + MarketplaceTemplateVersion, +) +from app.templates.marketplace_repository import MarketplaceTemplateRepository +from app.templates.marketplace_schemas import ( + MarketplaceTemplateDefinition, + SlotBinding, + TemplateApplicationView, + TemplateApply, + TemplateCreate, + TemplateInstantiate, + TemplateList, + TemplateRequirements, + TemplateUpdate, + TemplateVersionView, + TemplateView, +) + +_RESOLUTIONS = { + "16:9": (1920, 1080), + "9:16": (1080, 1920), + "1:1": (1080, 1080), + "4:5": (1080, 1350), +} + +_SLOT_MEDIA_PREFIXES = { + "video": ("video/",), + "image": ("image/",), + "audio": ("audio/",), + "logo": ("image/",), + "voice": ("audio/",), + "ai_image": ("image/",), + "ai_video": ("video/",), +} + + +class MarketplaceTemplateService: + def __init__( + self, + repository: MarketplaceTemplateRepository, + assets: CanonicalAssetService, + ai: AiStudioService, + audit: AuditService, + ) -> None: + self.repository = repository + self.assets = assets + self.ai = ai + self.audit = audit + + def _capabilities(self) -> set[str]: + capabilities = { + "editor.video", + "editor.audio", + "editor.caption", + "editor.overlay", + } + capabilities.update( + f"ai.{tool.operation}" for tool in self.ai.capabilities().tools if tool.available + ) + return capabilities + + async def list( + self, + *, + workspace_id: str, + user_id: str, + search: str | None, + category: str | None, + aspect_ratio: str | None, + min_duration_ms: int | None, + max_duration_ms: int | None, + media_type: str | None, + visibility: str | None, + status: str | None, + capability: str | None, + available_only: bool, + offset: int, + limit: int, + ) -> TemplateList: + if ( + min_duration_ms is not None + and max_duration_ms is not None + and min_duration_ms > max_duration_ms + ): + raise MarketplaceTemplateInvalid("Minimum duration must not exceed maximum duration.") + if available_only or capability or media_type: + page_items: list[TemplateView] = [] + filtered_total = 0 + database_offset = 0 + database_total = 0 + chunk_size = 100 + while True: + rows, database_total = await self.repository.list( + workspace_id, + user_id=user_id, + search=search, + category=category, + aspect_ratio=aspect_ratio, + min_duration_ms=min_duration_ms, + max_duration_ms=max_duration_ms, + visibility=visibility, + status=status, + offset=database_offset, + limit=chunk_size, + ) + for item in (self._view(template, version) for template, version in rows): + matches = (not available_only or item.available) and ( + not capability + or capability + in self._required_capabilities( + item.current_version.definition, + item.current_version.requirements, + ) + ) + if media_type: + accepted = { + accepted_type + for slot in item.current_version.definition.slots + for accepted_type in slot.accepted_media_types + } + matches = matches and any( + fnmatch.fnmatch(media_type, accepted_type) + or fnmatch.fnmatch(accepted_type, media_type) + for accepted_type in accepted + ) + if not matches: + continue + if offset <= filtered_total < offset + limit: + page_items.append(item) + filtered_total += 1 + database_offset += len(rows) + if not rows or database_offset >= database_total: + break + return TemplateList( + items=page_items, + offset=offset, + limit=limit, + total=filtered_total, + ) + rows, total = await self.repository.list( + workspace_id, + user_id=user_id, + search=search, + category=category, + aspect_ratio=aspect_ratio, + min_duration_ms=min_duration_ms, + max_duration_ms=max_duration_ms, + visibility=visibility, + status=status, + offset=offset, + limit=limit, + ) + items = [self._view(template, version) for template, version in rows] + return TemplateList(items=items, offset=offset, limit=limit, total=total) + + async def get(self, *, workspace_id: str, user_id: str, template_id: str) -> TemplateView: + return self._view(*(await self.repository.get(workspace_id, template_id, user_id=user_id))) + + async def versions( + self, *, workspace_id: str, user_id: str, template_id: str + ) -> list[TemplateVersionView]: + return [ + self._version_view(item) + for item in await self.repository.versions(workspace_id, template_id, user_id=user_id) + ] + + async def create( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + payload: TemplateCreate, + ) -> TemplateView: + await self._validate_preview_assets(workspace_id, user_id, payload) + template = MarketplaceTemplate( + workspace_id=workspace_id, + slug=payload.slug, + name=payload.name.strip(), + description=payload.description.strip(), + status="draft", + visibility=payload.visibility, + category=payload.category, + tags_json=payload.tags, + thumbnail_asset_id=( + str(payload.thumbnail_asset_id) if payload.thumbnail_asset_id else None + ), + preview_asset_id=str(payload.preview_asset_id) if payload.preview_asset_id else None, + duration_ms=payload.definition.settings.duration_ms, + aspect_ratio=payload.definition.settings.aspect_ratio, + metadata_json=payload.metadata, + created_by=user_id, + ) + version = MarketplaceTemplateVersion( + template_id="", + version=1, + schema_version=payload.definition.schema_version, + definition_json=payload.definition.model_dump(mode="json"), + requirements_json=payload.requirements.model_dump(mode="json"), + status="draft", + created_by=user_id, + ) + template, version = await self.repository.create(template, version, user_id=user_id) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "template.created", + template.id, + {"version": version.version, "visibility": template.visibility}, + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "template.version_created", + template.id, + {"version_id": version.id, "version": version.version}, + ) + return self._view(template, version) + + async def update( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + template_id: str, + payload: TemplateUpdate, + ) -> TemplateView: + preview_asset_ids = [ + asset_id + for asset_id in (payload.thumbnail_asset_id, payload.preview_asset_id) + if asset_id is not None + ] + await self._validate_preview_asset_ids(workspace_id, user_id, preview_asset_ids) + fields = payload.model_dump( + exclude_unset=True, + exclude={"definition", "requirements", "thumbnail_asset_id", "preview_asset_id"}, + ) + if "tags" in fields: + fields["tags_json"] = list( + dict.fromkeys(tag.strip().lower() for tag in fields.pop("tags") if tag.strip()) + ) + if "thumbnail_asset_id" in payload.model_fields_set: + fields["thumbnail_asset_id"] = ( + str(payload.thumbnail_asset_id) if payload.thumbnail_asset_id else None + ) + if "preview_asset_id" in payload.model_fields_set: + fields["preview_asset_id"] = ( + str(payload.preview_asset_id) if payload.preview_asset_id else None + ) + version = None + if payload.definition is not None or payload.requirements is not None: + existing_template, existing_version = await self.repository.get( + workspace_id, template_id, user_id=user_id + ) + if existing_template.workspace_id != workspace_id: + raise MarketplaceTemplateUnavailable( + "Public templates cannot be modified from another workspace." + ) + definition = payload.definition or MarketplaceTemplateDefinition.model_validate( + existing_version.definition_json + ) + requirements = payload.requirements or TemplateRequirements.model_validate( + existing_version.requirements_json + ) + fields["duration_ms"] = definition.settings.duration_ms + fields["aspect_ratio"] = definition.settings.aspect_ratio + version = MarketplaceTemplateVersion( + template_id=template_id, + version=0, + schema_version=definition.schema_version, + definition_json=definition.model_dump(mode="json"), + requirements_json=requirements.model_dump(mode="json"), + status="draft", + created_by=user_id, + ) + template, current = await self.repository.update( + workspace_id, + template_id, + user_id=user_id, + fields=fields, + version=version, + ) + event = ( + "template.published" + if fields.get("status") == "published" + else "template.archived" if fields.get("status") == "archived" else "template.updated" + ) + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + event, + template.id, + {"version": current.version, "fields": sorted(payload.model_fields_set)}, + ) + if version is not None: + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "template.version_created", + template.id, + {"version_id": current.id, "version": current.version}, + ) + return self._view(template, current) + + async def archive(self, **kwargs) -> TemplateView: + return await self.update(payload=TemplateUpdate(status="archived"), **kwargs) + + async def apply( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + template_id: str, + payload: TemplateApply | TemplateInstantiate, + idempotency_key: str, + instantiate: bool, + ) -> TemplateApplicationView: + try: + return await self._apply( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + template_id=template_id, + payload=payload, + idempotency_key=idempotency_key, + instantiate=instantiate, + ) + except Exception as exc: + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "template.application_failed", + template_id, + {"error_code": getattr(exc, "code", "TEMPLATE_APPLICATION_FAILED")}, + ) + raise + + async def _apply( + self, + *, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + template_id: str, + payload: TemplateApply | TemplateInstantiate, + idempotency_key: str, + instantiate: bool, + ) -> TemplateApplicationView: + template, version = await self.repository.get_version( + workspace_id, + template_id, + str(payload.template_version_id) if payload.template_version_id else None, + user_id=user_id, + ) + if template.status == "archived": + raise MarketplaceTemplateUnavailable("Archived templates cannot be applied.") + definition = MarketplaceTemplateDefinition.model_validate(version.definition_json) + requirements = TemplateRequirements.model_validate(version.requirements_json) + missing = sorted( + self._required_capabilities(definition, requirements) - self._capabilities() + ) + if missing: + raise MarketplaceTemplateUnavailable( + f"Template requires unavailable capabilities: {', '.join(missing)}" + ) + assets, document = await self._compile( + workspace_id=workspace_id, + user_id=user_id, + definition=definition, + bindings=payload.slot_bindings, + ) + target_project_id = None if instantiate else str(payload.project_id) # type: ignore[union-attr] + if instantiate: + try: + project_name = normalize_project_name(payload.project_name) # type: ignore[union-attr] + except ValueError as exc: + raise MarketplaceTemplateInvalid(str(exc)) from exc + project_description = payload.project_description # type: ignore[union-attr] + else: + project_name = None + project_description = None + for asset in assets.values(): + if asset.project_id != target_project_id: + raise MarketplaceTemplateInvalid( + "Every mapped asset must already belong to the target project." + ) + fingerprint = hashlib.sha256( + json.dumps( + { + "template_id": template.id, + "version_id": version.id, + "project_id": target_project_id, + "project_name": project_name, + "bindings": { + key: value.model_dump(mode="json") + for key, value in sorted(payload.slot_bindings.items()) + }, + }, + sort_keys=True, + separators=(",", ":"), + ).encode() + ).hexdigest() + application, _, created = await self.repository.apply( + workspace_id=workspace_id, + user_id=user_id, + template=template, + version=version, + project_id=target_project_id, + project_name=project_name, + project_description=project_description, + editor_state=document.model_dump(by_alias=True), + slot_bindings={ + key: value.model_dump(mode="json") for key, value in payload.slot_bindings.items() + }, + asset_ids={asset.id for asset in assets.values()}, + idempotency_key=idempotency_key, + fingerprint=fingerprint, + ) + if created: + await self._audit( + workspace_id, + user_id, + api_key_id, + request_id, + "template.applied", + template.id, + { + "application_id": application.id, + "project_id": application.project_id, + "version": application.template_version, + }, + ) + return self._application_view(application) + + async def _compile( + self, + *, + workspace_id: str, + user_id: str, + definition: MarketplaceTemplateDefinition, + bindings: dict[str, SlotBinding], + ) -> tuple[dict[str, object], EditorDocument]: + slots = {slot.id: slot for slot in definition.slots} + unknown = set(bindings) - set(slots) + if unknown: + raise MarketplaceTemplateInvalid( + f"Unknown template slots: {', '.join(sorted(unknown))}" + ) + assets: dict[str, object] = {} + for slot in definition.slots: + binding = bindings.get(slot.id) + if slot.required and binding is None and slot.default_text is None: + raise MarketplaceTemplateInvalid(f"Required slot '{slot.label}' is missing.") + if binding and binding.asset_id: + try: + asset = await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=str(binding.asset_id), + ) + except CanonicalAssetNotFoundError as exc: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' is unavailable." + ) from exc + expected_prefixes = _SLOT_MEDIA_PREFIXES.get(slot.type) + if expected_prefixes and not asset.mime_type.startswith(expected_prefixes): + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has an unsupported media type." + ) + if slot.accepted_media_types and not any( + fnmatch.fnmatch(asset.mime_type, pattern) + for pattern in slot.accepted_media_types + ): + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has an unsupported media type." + ) + duration = self._asset_duration(asset.metadata_json or {}) + if slot.duration and duration is None: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has no validated duration metadata." + ) + if slot.duration and duration is not None: + if slot.duration.min_ms and duration < slot.duration.min_ms: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' is too short." + ) + if slot.duration.max_ms and duration > slot.duration.max_ms: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' is too long." + ) + asset_ratio = self._asset_aspect_ratio(asset.metadata_json or {}) + if slot.aspect_ratios: + if asset_ratio is None: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has no validated aspect ratio metadata." + ) + if asset_ratio not in slot.aspect_ratios: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has an unsupported aspect ratio." + ) + assets[slot.id] = asset + tracks = [] + for order, source_track in enumerate(definition.tracks): + clips = [] + for source in source_track.clips: + slot = slots[source.slot_id] + binding = bindings.get(slot.id) + if slot.type in {"text", "caption"}: + text = ( + binding.text if binding and binding.text is not None else slot.default_text + ) + if text is None: + continue + clips.append( + CaptionClip( + id=source.id, + trackId=source_track.id, + label=source.label, + startMs=source.start_ms, + durationMs=source.duration_ms, + visible=True, + opacity=1, + metadata={"templateSlotId": slot.id}, + kind="caption", + text=text, + style=CaptionStyle(align="center", position="bottom"), + ) + ) + continue + asset = assets.get(slot.id) + if asset is None: + continue + mime = asset.mime_type + asset_duration = self._asset_duration(asset.metadata_json or {}) + if not mime.startswith("image/") and asset_duration is None: + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' has no validated duration metadata." + ) + if ( + not mime.startswith("image/") + and asset_duration is not None + and source.duration_ms > asset_duration + ): + raise MarketplaceTemplateInvalid( + f"Asset for slot '{slot.label}' is shorter than its template clip." + ) + common = dict( + id=source.id, + trackId=source_track.id, + label=source.label, + startMs=source.start_ms, + durationMs=source.duration_ms, + visible=True, + opacity=1, + metadata={"templateSlotId": slot.id}, + assetId=asset.id, + sourceStartMs=0, + sourceDurationMs=source.duration_ms, + ) + if mime.startswith("audio/"): + clips.append( + AudioClip( + **common, + kind="audio", + volume=1, + fadeInMs=0, + fadeOutMs=0, + ) + ) + else: + media_type = "image" if mime.startswith("image/") else "video" + clips.append( + MediaClip( + **common, + kind="media", + mediaType=media_type, + transform=ClipTransform(x=0, y=0, scaleX=1, scaleY=1, rotation=0), + volume=1, + ) + ) + tracks.append( + Track( + id=source_track.id, + type=source_track.type, + name=source_track.name, + order=order, + muted=False, + locked=False, + visible=True, + clips=clips, + ) + ) + width, height = _RESOLUTIONS[definition.settings.aspect_ratio] + document = EditorDocument( + schemaVersion=1, + projectId=str(uuid4()), + timeline=Timeline(timeUnit="milliseconds", tracks=tracks, transitions=[], markers=[]), + renderSettings=EditorRenderSettings( + format="mp4", + width=width, + height=height, + frameRate=definition.settings.frame_rate, + ), + ) + return assets, document + + async def _validate_preview_assets( + self, workspace_id: str, user_id: str, payload: TemplateCreate + ) -> None: + await self._validate_preview_asset_ids( + workspace_id, + user_id, + [ + asset_id + for asset_id in (payload.thumbnail_asset_id, payload.preview_asset_id) + if asset_id is not None + ], + ) + + async def _validate_preview_asset_ids( + self, workspace_id: str, user_id: str, asset_ids: list[object] + ) -> None: + for asset_id in asset_ids: + try: + await self.assets.get_owned_by_id( + workspace_id=workspace_id, + user_id=user_id, + asset_id=str(asset_id), + ) + except CanonicalAssetNotFoundError as exc: + raise MarketplaceTemplateInvalid( + "Template preview asset was not found in this workspace." + ) from exc + + def _view( + self, template: MarketplaceTemplate, version: MarketplaceTemplateVersion + ) -> TemplateView: + definition = MarketplaceTemplateDefinition.model_validate(version.definition_json) + requirements = TemplateRequirements.model_validate(version.requirements_json) + missing = sorted( + self._required_capabilities(definition, requirements) - self._capabilities() + ) + return TemplateView( + id=template.id, + workspace_id=template.workspace_id, + slug=template.slug, + name=template.name, + description=template.description, + status=template.status, + visibility=template.visibility, + category=template.category, + tags=template.tags_json or [], + thumbnail_asset_id=template.thumbnail_asset_id, + preview_asset_id=template.preview_asset_id, + duration_ms=template.duration_ms, + aspect_ratio=template.aspect_ratio, + metadata=template.metadata_json or {}, + created_by=template.created_by, + created_at=template.created_at, + updated_at=template.updated_at, + current_version=self._version_view(version), + missing_capabilities=missing, + available=not missing, + ) + + @staticmethod + def _version_view(version: MarketplaceTemplateVersion) -> TemplateVersionView: + return TemplateVersionView( + id=version.id, + template_id=version.template_id, + version=version.version, + schema_version=version.schema_version, + definition=version.definition_json, + requirements=version.requirements_json, + status=version.status, + created_at=version.created_at, + created_by=version.created_by, + ) + + @staticmethod + def _application_view(application) -> TemplateApplicationView: + return TemplateApplicationView( + id=application.id, + template_id=application.template_id, + template_version_id=application.template_version_id, + template_version=application.template_version, + project_id=application.project_id, + editor_revision=application.editor_revision, + slot_bindings=application.slot_bindings_json, + created_at=application.created_at, + ) + + @staticmethod + def _asset_duration(metadata: dict[str, object]) -> int | None: + value = metadata.get("duration_ms") + if isinstance(value, (int, float)): + return round(value) + value = metadata.get("duration") + return round(value * 1_000) if isinstance(value, (int, float)) else None + + @staticmethod + def _asset_aspect_ratio(metadata: dict[str, object]) -> str | None: + width = metadata.get("width") + height = metadata.get("height") + if not isinstance(width, (int, float)) or not isinstance(height, (int, float)): + return None + if width <= 0 or height <= 0: + return None + candidates = { + key: candidate_width / candidate_height + for key, (candidate_width, candidate_height) in _RESOLUTIONS.items() + } + ratio = width / height + closest = min(candidates, key=lambda key: abs(candidates[key] - ratio)) + return closest if abs(candidates[closest] - ratio) <= 0.03 else None + + @staticmethod + def _required_capabilities( + definition: MarketplaceTemplateDefinition, requirements: TemplateRequirements + ) -> set[str]: + required = set(requirements.capabilities) + for slot in definition.slots: + required.update(slot.capability_requirements) + return required + + async def _audit( + self, + workspace_id: str, + user_id: str, + api_key_id: str, + request_id: str, + event_type: str, + entity_id: str, + metadata: dict[str, object], + ) -> None: + await self.audit.record_event( + workspace_id=workspace_id, + user_id=user_id, + api_key_id=api_key_id, + request_id=request_id, + event_type=event_type, + entity_type="template", + entity_id=entity_id, + metadata=metadata, + ) diff --git a/main.py b/main.py index c08ea71d8de453ff88db0fc78b5a960058f0ebff..fe682da281b4546f4ede8999c86e9f13f58fa74b 100644 --- a/main.py +++ b/main.py @@ -9,18 +9,38 @@ from fastapi import FastAPI, HTTPException, Request from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse, ORJSONResponse from starlette.middleware.base import RequestResponseEndpoint +from starlette.middleware.cors import CORSMiddleware from starlette.responses import Response -from app.api import api_keys, audio, health, image, media, probe, social, templates, video, whisper, ytdlp +from app.ai import api as ai +from app.analytics import api as analytics +from app.brand import api as brand +from app.api import ( + api_keys, + audio, + generation, + health, + image, + media, + probe, + social, + templates, + video, + whisper, + ytdlp, +) from app.container import build_container +from app.copilot import api as copilot from app.core.config import Settings, get_settings from app.core.exceptions import MediaAPIError from app.core.logger import configure_logging, get_logger, request_id_context from app.core.response import ErrorBody, ErrorResponse from app.mcp.server import create_mcp_server +from app.projects import api as projects from app.security.middleware import APIKeyAuthenticationMiddleware -from app.workers.cleanup_worker import CleanupWorker from app.social.workers.scheduler import SocialSchedulerWorker +from app.templates import marketplace_api +from app.workers.cleanup_worker import CleanupWorker configure_logging() logger = get_logger(__name__) @@ -41,7 +61,9 @@ def create_app(settings: Settings | None = None) -> FastAPI: active_settings.ensure_directories() container = build_container(active_settings) cleanup_worker = CleanupWorker(container.cleanup, active_settings.cleanup_interval_seconds) - social_worker = SocialSchedulerWorker(container.social, active_settings.social_scheduler_interval_seconds) + social_worker = SocialSchedulerWorker( + container.social, active_settings.social_scheduler_interval_seconds + ) mcp_server = create_mcp_server(container) mcp_http_app = mcp_server.streamable_http_app() @@ -50,12 +72,30 @@ def create_app(settings: Settings | None = None) -> FastAPI: application.state.container = container application.state.mcp_server = mcp_server await container.security_database.initialize() + if not await container.security_database.schema_ready(): + missing = ", ".join(await container.security_database.missing_schema_objects()) + raise RuntimeError( + "Security schema is unavailable; apply app/security/migrations/ " + "and app/projects/migrations/. " + f"Missing: {missing}" + ) + await container.security_database.verify_execution_boundary( + expected_role=active_settings.security_database_role, + enforce_rls=active_settings.security_enforce_rls, + ) + await container.generation.initialize() await container.api_keys.ensure_bootstrap_admin() + await container.tenants.ensure_all_api_key_principals() await container.social.initialize() + await container.analytics.initialize(container.social.ready) + await container.social.adopt_legacy_workspaces(await container.tenants.list_principals()) async with mcp_server.session_manager.run(): await cleanup_worker.start() + await container.generation_worker.start() + await container.render_worker.start() if active_settings.social_enabled and active_settings.social_worker_enabled: await social_worker.start() + await container.analytics_worker.start() logger.info( "media API started", extra={"version": active_settings.app_version, "port": active_settings.port}, @@ -64,7 +104,11 @@ def create_app(settings: Settings | None = None) -> FastAPI: yield finally: await social_worker.stop() + await container.analytics_worker.stop() + await container.generation_worker.stop() + await container.render_worker.stop() await cleanup_worker.stop() + await container.generation.close() await container.social.close() await container.security_database.close() logger.info("media API stopped") @@ -93,6 +137,21 @@ def create_app(settings: Settings | None = None) -> FastAPI: rate_limiter=container.rate_limiter, audit=container.audit, ) + if active_settings.allowed_cors_origins: + application.add_middleware( + CORSMiddleware, + allow_origins=list(active_settings.allowed_cors_origins), + allow_credentials=True, + allow_methods=["GET", "HEAD", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"], + allow_headers=[ + "Authorization", + "Content-Type", + "Idempotency-Key", + "X-Request-ID", + "X-MediaRouter-Human-Role", + ], + expose_headers=["Content-Disposition", "X-Request-ID"], + ) @application.middleware("http") async def request_context(request: Request, call_next: RequestResponseEndpoint) -> Response: @@ -180,8 +239,18 @@ def create_app(settings: Settings | None = None) -> FastAPI: application.include_router(probe.router) application.include_router(ytdlp.router) application.include_router(whisper.router) + application.include_router(marketplace_api.router) + # Register the static marketplace prefix before the legacy dynamic + # /v1/templates/{template_id} route so "catalog" cannot be captured as a + # YAML workflow template ID. application.include_router(templates.router) + application.include_router(ai.router) + application.include_router(copilot.router) + application.include_router(generation.router) + application.include_router(projects.router) application.include_router(social.router) + application.include_router(analytics.router) + application.include_router(brand.router) application.mount("/mcp", mcp_http_app, name="mcp") @application.get("/", tags=["public"], include_in_schema=False) diff --git a/scripts/deployment_smoke.py b/scripts/deployment_smoke.py new file mode 100644 index 0000000000000000000000000000000000000000..1dbc754f8993cca5e53013d84063d7b33537772e --- /dev/null +++ b/scripts/deployment_smoke.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +"""Smoke-test a running MediaRouter production container using stdlib only.""" + +from __future__ import annotations + +import argparse +import json +import sys +import urllib.error +import urllib.request +from collections.abc import Iterable +from typing import Any + + +REQUIRED_PATHS: dict[str, frozenset[str]] = { + "/v1/projects": frozenset({"get", "post"}), + "/v1/projects/{project_id}": frozenset({"get", "patch", "delete"}), + "/v1/projects/{project_id}/assets": frozenset({"get", "post"}), + "/v1/projects/{project_id}/assets/{asset_id}": frozenset({"delete"}), + "/v1/projects/{project_id}/jobs": frozenset({"get", "post"}), + "/v1/projects/{project_id}/jobs/{job_id}": frozenset({"delete"}), +} +REQUIRED_PREFIXES = ("/v1/social", "/v1/generation") + + +def request_json( + base_url: str, path: str, *, api_key: str = "" +) -> tuple[int, dict[str, Any]]: + headers = {"Accept": "application/json"} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + request = urllib.request.Request(f"{base_url.rstrip('/')}{path}", headers=headers) + try: + with urllib.request.urlopen(request, timeout=15) as response: + status = response.status + payload = response.read() + except urllib.error.HTTPError as exc: + status = exc.code + payload = exc.read() + try: + decoded = json.loads(payload) + except (json.JSONDecodeError, UnicodeDecodeError) as exc: + raise RuntimeError(f"{path} did not return JSON (HTTP {status})") from exc + if not isinstance(decoded, dict): + raise RuntimeError(f"{path} did not return a JSON object") + return status, decoded + + +def missing_operations(paths: dict[str, Any]) -> Iterable[str]: + for path, required_methods in REQUIRED_PATHS.items(): + operations = paths.get(path) + if not isinstance(operations, dict): + yield f"{path} (path missing)" + continue + missing = required_methods - operations.keys() + for method in sorted(missing): + yield f"{method.upper()} {path}" + for prefix in REQUIRED_PREFIXES: + if not any(path.startswith(prefix) for path in paths): + yield f"{prefix}/* (route family missing)" + + +def run(base_url: str, api_key: str) -> None: + health_status, health = request_json(base_url, "/health") + if health_status != 200: + raise RuntimeError(f"/health returned HTTP {health_status}") + if health.get("success") is not True: + raise RuntimeError("/health did not return the MediaRouter success envelope") + + openapi_status, openapi = request_json(base_url, "/openapi.json") + if openapi_status != 200: + raise RuntimeError(f"/openapi.json returned HTTP {openapi_status}") + paths = openapi.get("paths") + if not isinstance(paths, dict): + raise RuntimeError("runtime OpenAPI contains no paths object") + missing = list(missing_operations(paths)) + if missing: + raise RuntimeError("runtime OpenAPI is incomplete: " + ", ".join(missing)) + + auth_status, _auth = request_json(base_url, "/v1/auth/context", api_key=api_key) + expected_auth_status = 200 if api_key else 401 + if auth_status != expected_auth_status: + raise RuntimeError( + "/v1/auth/context returned " + f"HTTP {auth_status}; expected {expected_auth_status}" + ) + + final_health_status, _final_health = request_json(base_url, "/health") + if final_health_status != 200: + raise RuntimeError("application stopped responding during the smoke test") + + print("DEPLOYMENT_SMOKE=PASS") + print(f"RUNTIME_OPENAPI_PATHS={len(paths)}") + print(f"AUTHENTICATION_CHECK={'AUTHENTICATED' if api_key else 'FAIL_CLOSED'}") + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--base-url", default="http://127.0.0.1:7860") + parser.add_argument( + "--api-key", + default="", + help="Optional test key; never printed. Without it, a 401 is required.", + ) + args = parser.parse_args() + try: + run(args.base_url, args.api_key) + except Exception as exc: + print(f"DEPLOYMENT_SMOKE=FAIL: {exc}", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/docker_deployment_gate.sh b/scripts/docker_deployment_gate.sh new file mode 100644 index 0000000000000000000000000000000000000000..4e6a5047f59afaad534f0f1ea1597f4b2b837191 --- /dev/null +++ b/scripts/docker_deployment_gate.sh @@ -0,0 +1,56 @@ +#!/bin/sh +set -eu + +if [ "$#" -lt 1 ]; then + echo "usage: $0 /path/to/production.env [optional-api-key]" >&2 + exit 2 +fi + +deployment_env=$1 +deployment_api_key=${2:-} +image_name=mediarouter-deployment-gate +container_name=mediarouter-deployment-gate-run-$$ + +docker build --tag "$image_name" . +docker run \ + --detach \ + --name "$container_name" \ + --env-file "$deployment_env" \ + --publish 127.0.0.1:7860:7860 \ + "$image_name" >/dev/null + +cleanup() { + docker rm --force "$container_name" >/dev/null 2>&1 || true +} +trap cleanup EXIT INT TERM + +attempt=0 +while [ "$attempt" -lt 30 ]; do + health_status=$(docker inspect --format '{{if .State.Health}}{{.State.Health.Status}}{{end}}' "$container_name") + if [ "$health_status" = "healthy" ]; then + break + fi + if [ "$(docker inspect --format '{{.State.Running}}' "$container_name")" != "true" ]; then + docker logs "$container_name" + exit 1 + fi + attempt=$((attempt + 1)) + sleep 2 +done + +if [ "${health_status:-}" != "healthy" ]; then + docker logs "$container_name" + echo "container did not become healthy" >&2 + exit 1 +fi + +python3 scripts/deployment_smoke.py \ + --base-url http://127.0.0.1:7860 \ + --api-key "$deployment_api_key" + +if [ "$(docker inspect --format '{{.State.Running}}' "$container_name")" != "true" ]; then + echo "container stopped after smoke test" >&2 + exit 1 +fi + +echo "DOCKER_DEPLOYMENT_GATE=PASS"