diff --git a/README.md b/README.md index f675ab3282af70dd7b0c8414530478bfae07a31f..3a09f3855877c70a7ea9be939f94dbdd6d6e0356 100644 --- a/README.md +++ b/README.md @@ -84,6 +84,21 @@ 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. +## Current implementation surface + +The current source tree exposes the following product areas. Missing docs or backend modules are recorded as omitted from this working tree, not as unimplemented capabilities elsewhere. + +- Brand Kit: `/v1/brand`, frontend brand-kit module, Python/TypeScript SDK resources, MCP brand tools, n8n MediaBrandKit node. +- Projects and collaboration: `/v1/projects`, editor persistence/render jobs, teams/members/invitations, approvals/review comments, notification preferences, Python/TypeScript SDK project resources, MCP collaboration tools, n8n MediaCollaboration node. +- Publishing and social automation: `/v1/social`, unified publishing, scheduling, retries, reconciliation, analytics sync, Python/TypeScript SDK social resources, MCP social tools, n8n MediaSocial node. +- Analytics: `/v1/analytics`, overview/timeseries/platform/post analytics, sync runs/cancellation, frontend analytics module, MCP analytics tools, SDK analytics resource. See [`docs/analytics.md`](docs/analytics.md) for implementation and runtime status. +- AI Studio / Copilot: `/v1/ai`, `/v1/copilot`, durable jobs, deterministic planner, frontend AI/Copilot modules, SDK resources, MCP tools. See [`docs/ai-copilot.md`](docs/ai-copilot.md) for implementation and runtime status. +- Templates: `/v1/templates`, `/v1/templates/catalog`, template execution, marketplace API, frontend template/marketplace modules, SDK/n8n/MCP coverage. See [`docs/template-marketplace.md`](docs/template-marketplace.md) for implementation and runtime status. +- Content Studio: editor state, optimistic revisioning, autosave, render jobs, timeline UX, SDK/n8n integration points. See [`docs/content-studio-foundation.md`](docs/content-studio-foundation.md) and [`docs/content-studio-persistence-rendering.md`](docs/content-studio-persistence-rendering.md) for current implementation surface. +- Brand Kit: `/v1/brand`, frontend brand-kit module, Python/TypeScript SDK resources, MCP brand tools, and an n8n `MediaBrandKit` package entry. See [`docs/brand-kits.md`](docs/brand-kits.md) for implementation and runtime status. + +Provider certification, PostgreSQL/RLS runtime, Docker, and Hugging Face runtime verification remain deferred until the planned product phases are complete. + ## Generation foundation MediaRouter includes a tenant-scoped generation request/job foundation at diff --git a/api/__init__.py b/api/__init__.py index ed2cfe0e64914f6af144ece3753b2fd9692a0ab5..d12cfd93588e432a2fd691517a2452d3397c6b03 100644 --- a/api/__init__.py +++ b/api/__init__.py @@ -1 +1,14 @@ """Compatibility exports for the canonical :mod:`app.api` package.""" + +from app.api.api_keys import * # noqa +from app.api.audio import * # noqa +from app.api.generation import * # noqa +from app.api.health import * # noqa +from app.api.image import * # noqa +from app.api.media import * # noqa +from app.api.probe import * # noqa +from app.api.social import * # noqa +from app.api.templates import * # noqa +from app.api.video import * # noqa +from app.api.whisper import * # noqa +from app.api.ytdlp import * # noqa diff --git a/app/api/api_keys.py b/app/api/api_keys.py index f989528cf003733aa1609ce0dbe118bce45dbcf1..c0fd6eabc109346752763746c266efc709812541 100644 --- a/app/api/api_keys.py +++ b/app/api/api_keys.py @@ -40,7 +40,6 @@ async def current_auth_context(request: Request) -> AuthContextView: expires_at=context.expires_at, workspace_id=context.workspace_id, user_id=context.user_id, - membership_role=context.membership_role, ) diff --git a/app/mcp/server.py b/app/mcp/server.py index e41b0ce4e5f8076ce67b8e62c132d0c13352c75c..42fd4cf135b967e0638c4e15fdf8c4e3b7f79e55 100644 --- a/app/mcp/server.py +++ b/app/mcp/server.py @@ -21,6 +21,7 @@ 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.collaboration import register_collaboration_tools from app.mcp.tools.probe import register_probe_tools from app.mcp.tools.social import register_social_tools from app.mcp.tools.system import register_system_tools @@ -62,6 +63,7 @@ def create_mcp_server(container: Container) -> FastMCP[Any]: register_template_tools(server, registry) register_social_tools(server, registry) register_brand_tools(server, registry) + register_collaboration_tools(server, registry) register_ai_tools(server, registry) register_analytics_tools(server, registry) register_resources(server, registry) diff --git a/app/mcp/tools/collaboration.py b/app/mcp/tools/collaboration.py index cfc72db1e7010e34e250d8cf1389f574a0982a66..83a80ffe5ed67a210d5481d9ecc5929b073fa434 100644 --- a/app/mcp/tools/collaboration.py +++ b/app/mcp/tools/collaboration.py @@ -1,25 +1,35 @@ from __future__ import annotations from typing import Any -from app.mcp.registry import register_tool +from mcp.server.fastmcp import FastMCP + +from app.mcp.registry import MCPRegistry 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() + +def register_collaboration_tools(server: FastMCP[Any], registry: MCPRegistry) -> None: + @server.tool(description="List teams in the current workspace.") + async def list_teams() -> dict[str, Any]: + auth = auth_context.get() + if not auth or not auth.workspace_id: + raise ValueError("Unauthorized") + + async def action() -> list[dict[str, Any]]: + service: CollaborationService = registry.container.collaboration + return [item.model_dump(mode="json") for item in await service.list_teams(auth.workspace_id)] + + return await registry.run_metadata_tool("collaboration.list_teams", action) + + @server.tool(description="Create a team in the current workspace.") + async def create_team(name: str) -> dict[str, Any]: + auth = auth_context.get() + if not auth or not auth.workspace_id: + raise ValueError("Unauthorized") + + async def action() -> dict[str, Any]: + service: CollaborationService = registry.container.collaboration + return (await service.create_team(auth.workspace_id, name)).model_dump(mode="json") + + return await registry.run_metadata_tool("collaboration.create_team", action) + +register_collaboration_tools # re-export registration symbol for compatibility diff --git a/app/projects/api.py b/app/projects/api.py index 582e55df285e7e7a151a84ac17692371f8ac6f15..c799267b3130a0f1d5c06328de74d2848790488c 100644 --- a/app/projects/api.py +++ b/app/projects/api.py @@ -32,7 +32,7 @@ from app.projects.schemas.collaboration import ( TeamBase, MemberResponse ) -from app.projects.schemas.approval import ApprovalRequest, ReviewComment +from app.projects.schemas.approval import ApprovalRequest, ApprovalRequestCreate, ReviewComment from app.security.errors import ForbiddenError router = APIRouter(prefix="/v1/projects", tags=["projects"]) @@ -82,7 +82,24 @@ async def archive_team(request: Request, team_id: str) -> Response: @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) + workspace_id, _, _, _ = _identity(request) + return await request.app.state.container.approval.list_requests(workspace_id, workflow_id) + + + +@router.post("/workspace/workflows/{workflow_id}/requests", response_model=ApprovalRequest, status_code=status.HTTP_201_CREATED) +async def create_approval_request( + request: Request, + workflow_id: str, + payload: ApprovalRequestCreate, +) -> ApprovalRequest: + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.create_request( + workspace_id=workspace_id, + workflow_id=workflow_id, + project_id=payload.project_id, + user_id=user_id, + ) @router.post("/workspace/requests/{request_id}/approve", response_model=ApprovalRequest) @@ -90,16 +107,16 @@ 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) + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.approve_request(workspace_id, 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) + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.approval.reject_request(workspace_id, request_id, user_id) @router.post("/workspace/requests/{request_id}/comments", response_model=ReviewComment) async def add_review_comment( @@ -309,7 +326,10 @@ 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)) + workspace_id, user_id, _, _ = _identity(request) + return await request.app.state.container.collaboration.list_project_collaborators( + workspace_id, str(project_id), user_id=user_id + ) @router.post("/{project_id}/collaborators", response_model=MemberResponse) async def add_project_collaborator( @@ -329,7 +349,10 @@ async def remove_project_collaborator( project_id: UUID, user_id: str ) -> Response: - await request.app.state.container.collaboration.remove_project_collaborator(str(project_id), user_id) + workspace_id, _, _, _ = _identity(request) + await request.app.state.container.collaboration.remove_project_collaborator( + workspace_id, str(project_id), user_id + ) return Response(status_code=status.HTTP_204_NO_CONTENT) diff --git a/app/projects/migrations/0014_approval_workspace_integrity.sql b/app/projects/migrations/0014_approval_workspace_integrity.sql new file mode 100644 index 0000000000000000000000000000000000000000..4c5157bc56ac5a1a6438edf4aad8ae674625727f --- /dev/null +++ b/app/projects/migrations/0014_approval_workspace_integrity.sql @@ -0,0 +1,42 @@ +-- Additive approval workspace integrity migration. +-- +-- This migration brings the existing approval domain in line with the +-- authoritative ORM models and closes the confirmed cross-workspace +-- authorization gap by backfilling approval requests with their owning +-- workflow workspace, enforcing the relationship, and adding an index +-- used by authorization checks. +-- +-- Apply after 0013_notification_preferences.sql. + +begin; + +alter table approval_requests + add column if not exists workspace_id text; + +do $$ +begin + if exists ( + select 1 + from approval_requests + where workspace_id is null + ) then + update approval_requests + set workspace_id = approval_workflows.workspace_id + from approval_workflows + where approval_workflows.id = approval_requests.workflow_id + and approval_requests.workspace_id is null; + end if; +end $$; + +alter table approval_requests + alter column workspace_id set not null; + +alter table approval_requests + drop constraint if exists approval_requests_workflow_id_fkey, + add constraint approval_requests_workflow_id_fkey + foreign key (workflow_id) references approval_workflows(id) on delete cascade; + +create index if not exists ix_approval_requests_workspace + on approval_requests(workspace_id); + +commit; diff --git a/app/projects/repositories/approval_repository.py b/app/projects/repositories/approval_repository.py index 57a087072e442cabaf6c184bc3f76111b8554ac0..ac81c4774596a67c61d4b6fb08962fff9b237809 100644 --- a/app/projects/repositories/approval_repository.py +++ b/app/projects/repositories/approval_repository.py @@ -30,9 +30,14 @@ class ApprovalRepository: await session.refresh(request) return request - async def update_request_status(self, request_id: str, status: str) -> ApprovalRequest: + async def update_request_status(self, request_id: str, workspace_id: str, status: str) -> ApprovalRequest: async with self.database.session() as session: - request = await session.get(ApprovalRequest, request_id) + request = await session.scalar( + select(ApprovalRequest).where( + ApprovalRequest.id == request_id, + ApprovalRequest.workspace_id == workspace_id, + ) + ) if not request: raise Exception("Request not found") request.status = status @@ -40,9 +45,14 @@ class ApprovalRepository: await session.refresh(request) return request - async def get_request(self, request_id: str) -> ApprovalRequest: + async def get_request(self, request_id: str, workspace_id: str) -> ApprovalRequest: async with self.database.session() as session: - request = await session.get(ApprovalRequest, request_id) + request = await session.scalar( + select(ApprovalRequest).where( + ApprovalRequest.id == request_id, + ApprovalRequest.workspace_id == workspace_id, + ) + ) if not request: raise Exception("Request not found") return request @@ -60,10 +70,12 @@ class ApprovalRepository: await session.refresh(comment) return comment - async def list_requests(self, workflow_id: str) -> list[ApprovalRequest]: + async def list_requests(self, workspace_id: str, workflow_id: str) -> list[ApprovalRequest]: async with self.database.session() as session: result = await session.scalars( - select(ApprovalRequest).where(ApprovalRequest.workflow_id == workflow_id) + select(ApprovalRequest).where( + ApprovalRequest.workflow_id == workflow_id, + ApprovalRequest.workspace_id == workspace_id, + ) ) return list(result.all()) - diff --git a/app/projects/repositories/collaboration_repository.py b/app/projects/repositories/collaboration_repository.py index 153c627bf1e73e38ee2c072ede0dd19cda6fca0e..b8957f5e65a911c1c3bdcab66220d6e0e5582ff6 100644 --- a/app/projects/repositories/collaboration_repository.py +++ b/app/projects/repositories/collaboration_repository.py @@ -149,14 +149,23 @@ class CollaborationRepository: return list(result.all()) async def list_project_collaborators(self, project_id: str) -> list[ProjectCollaborator]: - async with self.database.session() as session: + project = await self._project_for_resources(project_id) + async with self.database.tenant_session( + workspace_id=project.workspace_id, user_id=None + ) as session: result = await session.scalars( - select(ProjectCollaborator).where(ProjectCollaborator.project_id == project_id) + select(ProjectCollaborator).where( + ProjectCollaborator.project_id == project_id, + ProjectCollaborator.workspace_id == project.workspace_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: + await self._project_for_resources(project_id) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: collaborator = ProjectCollaborator( workspace_id=workspace_id, project_id=project_id, @@ -168,14 +177,28 @@ class CollaborationRepository: 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: + async def remove_project_collaborator(self, workspace_id: str, project_id: str, user_id: str) -> None: + await self._project_for_resources(project_id) + async with self.database.tenant_session( + workspace_id=workspace_id, user_id=user_id + ) as session: collaborator = await session.scalar( select(ProjectCollaborator).where( ProjectCollaborator.project_id == project_id, + ProjectCollaborator.workspace_id == workspace_id, ProjectCollaborator.user_id == user_id ) ) if collaborator: await session.delete(collaborator) await session.commit() + + @staticmethod + async def _project_for_resources(project_id: str) -> Project: + async with CollaborationRepository(None).database.session() as session: + project = await session.scalar( + select(Project).where(Project.id == project_id) + ) + if project is None: + raise Exception("Project not found") + return project diff --git a/app/projects/schemas/approval.py b/app/projects/schemas/approval.py index cf590a65ee384a49b1095004154a0b58f2f1e3c2..af7eca21348c326122b03e93d4a86e84c12fa5f1 100644 --- a/app/projects/schemas/approval.py +++ b/app/projects/schemas/approval.py @@ -9,6 +9,11 @@ class ApprovalRequest(BaseModel): created_by: str created_at: str + +class ApprovalRequestCreate(BaseModel): + project_id: str + + class ReviewComment(BaseModel): id: str request_id: str diff --git a/app/projects/schemas/collaboration.py b/app/projects/schemas/collaboration.py index 3003944d449d86d786378200a33d8021db5668e1..a47150c55bdb28c63e42efe125c7ba420177ba7f 100644 --- a/app/projects/schemas/collaboration.py +++ b/app/projects/schemas/collaboration.py @@ -20,7 +20,6 @@ class InvitationResponse(InvitationCreate): status: str expires_at: str created_at: str - token: str | None = None class TeamBase(BaseModel): name: str diff --git a/app/projects/services/approval_service.py b/app/projects/services/approval_service.py index 16cec16292e8f3d2b3e4f162462d7fb862180953..b8385f030d8db5ce46c0571735f38838393672c4 100644 --- a/app/projects/services/approval_service.py +++ b/app/projects/services/approval_service.py @@ -33,5 +33,5 @@ class ApprovalService: 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) + async def list_requests(self, workspace_id: str, workflow_id: str) -> list[ApprovalRequest]: + return await self.repository.list_requests(workspace_id, workflow_id) diff --git a/app/projects/services/collaboration_service.py b/app/projects/services/collaboration_service.py index e5444171c6b56dcd3fd5aafaa5a8eeecba3832b2..5b55c302c095bcdb4eb6bd52e56207f608c12c35 100644 --- a/app/projects/services/collaboration_service.py +++ b/app/projects/services/collaboration_service.py @@ -50,7 +50,6 @@ class CollaborationService: 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]: @@ -107,14 +106,17 @@ class CollaborationService: 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) + async def list_project_collaborators(self, workspace_id: str, project_id: str, *, user_id: str) -> list[MemberResponse]: + collaborators = await self.repository.list_project_collaborators(workspace_id, 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 remove_project_collaborator(self, workspace_id: str, project_id: str, user_id: str) -> None: + await self.repository.remove_project_collaborator(workspace_id, project_id, user_id) + 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) diff --git a/app/security/schemas.py b/app/security/schemas.py index a68d709c59652ed37438003e344ddffc76b30c9a..cb67704ca8793acc5e536a5f90a387f206ca52f3 100644 --- a/app/security/schemas.py +++ b/app/security/schemas.py @@ -133,7 +133,6 @@ class AuthContextView(BaseModel): 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/core/__init__.py b/core/__init__.py index bd3863139fadb8241362b7c098686e159d45fd32..3b5668e81019a7de84cd3e00e121cbcbf5f4f0b9 100644 --- a/core/__init__.py +++ b/core/__init__.py @@ -1 +1,6 @@ """Compatibility exports for the canonical :mod:`app.core` package.""" + +from app.core.config import * # noqa +from app.core.exceptions import * # noqa +from app.core.logger import * # noqa +from app.core.response import * # noqa diff --git a/models/__init__.py b/models/__init__.py index 3e3320ee5b8bdd882002a45c5bd33b1036ecbc89..1350c3baf095082d37a144b0cde3d0ddd1b17088 100644 --- a/models/__init__.py +++ b/models/__init__.py @@ -1 +1,4 @@ """Compatibility exports for the canonical :mod:`app.models` package.""" + +from app.models.media import * # noqa +from app.models.requests import * # noqa diff --git a/operations/__init__.py b/operations/__init__.py index 88769466fca84a24de040f25600dfb5a43ab175a..819b36cd54833a0b3081c5a83335b8b386c6df31 100644 --- a/operations/__init__.py +++ b/operations/__init__.py @@ -1 +1,15 @@ """Compatibility exports for the canonical :mod:`app.operations` package.""" + +from app.operations.common import * # noqa +from app.operations.compress import * # noqa +from app.operations.concat import * # noqa +from app.operations.convert import * # noqa +from app.operations.crop import * # noqa +from app.operations.extract_audio import * # noqa +from app.operations.merge import * # noqa +from app.operations.resize import * # noqa +from app.operations.rotate import * # noqa +from app.operations.subtitles import * # noqa +from app.operations.thumbnails import * # noqa +from app.operations.trim import * # noqa +from app.operations.watermark import * # noqa diff --git a/services/__init__.py b/services/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..774f43c602072c37a4b81210d15921b77de2eba7 --- /dev/null +++ b/services/__init__.py @@ -0,0 +1,11 @@ +"""Compatibility exports for the canonical :mod:`app.services` package.""" + +from app.services.cleanup import * # noqa +from app.services.downloader import * # noqa +from app.services.ffmpeg_service import * # noqa +from app.services.ffprobe_service import * # noqa +from app.services.input_resolver import * # noqa +from app.services.media_service import * # noqa +from app.services.validator import * # noqa +from app.services.whisper_service import * # noqa +from app.services.ytdlp_service import * # noqa diff --git a/services/cleanup.py b/services/cleanup.py new file mode 100644 index 0000000000000000000000000000000000000000..1b885fedc1a6a7c457fb8ea4a7a23a9390618dac --- /dev/null +++ b/services/cleanup.py @@ -0,0 +1 @@ +from app.services.cleanup import * # noqa diff --git a/services/downloader.py b/services/downloader.py new file mode 100644 index 0000000000000000000000000000000000000000..dc73193a91ba456b3a411adeccc457b324d1a1ed --- /dev/null +++ b/services/downloader.py @@ -0,0 +1 @@ +from app.services.downloader import * # noqa diff --git a/services/ffmpeg_service.py b/services/ffmpeg_service.py new file mode 100644 index 0000000000000000000000000000000000000000..4145bb4ddc12b8373f722886b25453e6ac3821d6 --- /dev/null +++ b/services/ffmpeg_service.py @@ -0,0 +1 @@ +from app.services.ffmpeg_service import * # noqa diff --git a/services/ffprobe_service.py b/services/ffprobe_service.py new file mode 100644 index 0000000000000000000000000000000000000000..16577f2b4dbd1f3ee306cf4c5570a816076798d8 --- /dev/null +++ b/services/ffprobe_service.py @@ -0,0 +1 @@ +from app.services.ffprobe_service import * # noqa diff --git a/services/input_resolver.py b/services/input_resolver.py new file mode 100644 index 0000000000000000000000000000000000000000..97fa308554b717ddb4f9343fe459552f52122713 --- /dev/null +++ b/services/input_resolver.py @@ -0,0 +1 @@ +from app.services.input_resolver import * # noqa diff --git a/services/media_service.py b/services/media_service.py new file mode 100644 index 0000000000000000000000000000000000000000..411220290375601acdad0a505b747b91fb893605 --- /dev/null +++ b/services/media_service.py @@ -0,0 +1 @@ +from app.services.media_service import * # noqa diff --git a/services/validator.py b/services/validator.py new file mode 100644 index 0000000000000000000000000000000000000000..cb4640948b4fc268f475f3e0a2c3f4be9730df1d --- /dev/null +++ b/services/validator.py @@ -0,0 +1 @@ +from app.services.validator import * # noqa diff --git a/services/whisper_service.py b/services/whisper_service.py new file mode 100644 index 0000000000000000000000000000000000000000..8030ace437db40062578422adefd7b3d743628fb --- /dev/null +++ b/services/whisper_service.py @@ -0,0 +1 @@ +from app.services.whisper_service import * # noqa diff --git a/services/ytdlp_service.py b/services/ytdlp_service.py new file mode 100644 index 0000000000000000000000000000000000000000..e56d27450b12ef75c98d77acc6ce832b2b0d1bb4 --- /dev/null +++ b/services/ytdlp_service.py @@ -0,0 +1 @@ +from app.services.ytdlp_service import * # noqa diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000000000000000000000000000000000000..0bc237625c444eac957bb273d9bd574c0867b13c --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from app.core.config import Settings + + +@pytest.fixture +def settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + max_upload_size=10 * 1024 * 1024, + cleanup_minutes=1, + cleanup_interval_seconds=3600, + whisper_model="tiny", + max_workers=1, + allow_private_urls=True, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + ) diff --git a/tests/test_ai_copilot.py b/tests/test_ai_copilot.py new file mode 100644 index 0000000000000000000000000000000000000000..87a96bfcac73cd60727d7935334cee1d6f5428a0 --- /dev/null +++ b/tests/test_ai_copilot.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +from pathlib import Path +from uuid import uuid4 + +import pytest +from pydantic import ValidationError + +from app.copilot.actions import CopilotActionRegistry +from app.copilot.errors import CopilotInvalidRequestError +from app.copilot.planner import CopilotPlanner +from app.copilot.schemas import ( + CopilotContext, + CopilotEditorSummary, + CopilotPlan, +) + + +def context(*, capabilities: list[str], asset: bool = False, clip: bool = False): + project_id = uuid4() + return CopilotContext( + workspace_id=str(uuid4()), + project_id=project_id, + selected_asset_ids=[uuid4()] if asset else [], + selected_clip_ids=["clip-1"] if clip else [], + editor_summary=CopilotEditorSummary( + revision=4, duration_ms=30_000, track_count=1, clip_count=1 + ), + available_capabilities=capabilities, + ) + + +def test_planner_fails_closed_for_unavailable_transcription() -> None: + plan = CopilotPlanner().plan( + "Turn this podcast into a TikTok", + context(capabilities=["editor.render"], asset=True), + ) + assert not plan.executable + assert plan.unsupported_capabilities == ["ai.transcribe"] + assert plan.actions == [] + + +def test_planner_requires_confirmation_for_render_and_generation() -> None: + render = CopilotPlanner().plan("Render this project", context(capabilities=["editor.render"])) + assert render.executable and render.requires_confirmation + assert render.actions[0].type == "editor.render" + image = CopilotPlanner().plan( + "Generate an image of a lighthouse", + context(capabilities=["ai.generate_image"]), + ) + assert image.executable and image.requires_confirmation + assert image.actions[0].type == "ai.generate_image" + + +def test_action_plan_rejects_unknown_model_generated_structures() -> None: + with pytest.raises(ValidationError): + CopilotPlan.model_validate( + { + "intent": "unsafe", + "explanation": "unsafe", + "actions": [ + { + "id": "a", + "type": "shell.execute", + "arguments": {"command": "rm -rf /"}, + "reason": "unsafe", + "requires_confirmation": False, + "destructive": False, + "external_side_effect": False, + "required_permission": "admin", + "required_capability": "shell", + } + ], + "missing_information": [], + "unsupported_capabilities": [], + "executable": True, + "requires_confirmation": False, + } + ) + + +def test_action_registry_rejects_policy_metadata_tampering() -> None: + registry = CopilotActionRegistry( + projects=None, # type: ignore[arg-type] + assets=None, # type: ignore[arg-type] + editor=None, # type: ignore[arg-type] + renders=None, # type: ignore[arg-type] + ai=None, # type: ignore[arg-type] + templates=None, # type: ignore[arg-type] + ) + plan = CopilotPlanner().plan( + "Generate an image of a lighthouse", + context(capabilities=["ai.generate_image"]), + ) + tampered = plan.actions[0].model_copy(update={"requires_confirmation": False}) + with pytest.raises(CopilotInvalidRequestError): + registry.validate(tampered) + + +def test_copilot_migration_is_additive_and_tenant_isolated() -> None: + migration = ( + (Path(__file__).resolve().parents[1] / "app/projects/migrations/0005_ai_copilot.sql") + .read_text(encoding="utf-8") + .lower() + ) + for expected in ( + "create table if not exists copilot_runs", + "unique (workspace_id, idempotency_key)", + "enable row level security", + "force row level security", + "create policy copilot_runs_select", + "create policy copilot_runs_insert", + "create policy copilot_runs_update", + "copilot run identity fields are immutable", + ): + assert expected in migration + assert "drop table" not in migration diff --git a/tests/test_analytics_phase10_static.py b/tests/test_analytics_phase10_static.py new file mode 100644 index 0000000000000000000000000000000000000000..0269f170139e49e9d9f258e4781b97ade040268c --- /dev/null +++ b/tests/test_analytics_phase10_static.py @@ -0,0 +1,42 @@ +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] + + +def test_analytics_migration_is_additive_and_forces_rls() -> None: + text = (ROOT / "app/social/migrations/0010_analytics_insights.sql").read_text() + normalized = text.lower() + assert "drop table" not in normalized + for table in ( + "analytics_sync_runs", + "analytics_metric_snapshots", + "analytics_post_metrics", + "analytics_platform_metrics", + ): + assert f"create table if not exists {table}" in normalized + assert normalized.count("force row level security") >= 1 + assert "current_setting(''app.workspace_id''" in normalized + + +def test_analytics_routes_and_transports_are_narrow() -> None: + api = (ROOT / "app/analytics/api.py").read_text() + mcp = (ROOT / "app/mcp/tools/analytics.py").read_text() + sdk = (ROOT / "sdk/typescript/src/resources/analytics.ts").read_text() + for route in ( + '"/overview"', + '"/timeseries"', + '"/platforms"', + '"/posts"', + '"/sync"', + ): + assert route in api + assert "execute analytics query" not in mcp.lower() + assert "class AnalyticsResource" in sdk + + +def test_analytics_never_fabricates_provider_metrics() -> None: + service = (ROOT / "app/analytics/service.py").read_text() + provider = (ROOT / "app/social/providers/base.py").read_text() + assert "self.social_analytics.post" in service + assert "get_metrics" in provider + assert "random.randint" not in service diff --git a/tests/test_api_contract_regression.py b/tests/test_api_contract_regression.py new file mode 100644 index 0000000000000000000000000000000000000000..f7dca4b5c7c83600cf9dc09cc59204e07c1eb2cd --- /dev/null +++ b/tests/test_api_contract_regression.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +import ast +import re +from pathlib import Path +import unittest + + +ROUTES = { + "brand_api": Path("app/brand/api.py"), + "projects_api": Path("app/projects/api.py"), +} + +FRONTEND_CALLS = { + "brand_api": Path("frontend/features/brand-kits/api/index.ts"), + "collaboration_api": Path("frontend/features/workspace/collaboration/api/collaboration.ts"), +} + +BRAND_ROUTES = { + "router.post('', response_model=BrandKitResponse, status_code=status.HTTP_201_CREATED)": "/v1/brand POST", + "router.get('', response_model=list[BrandKitResponse])": "/v1/brand GET", + "router.patch('/{brand_kit_id}', response_model=BrandKitResponse)": "/v1/brand PATCH", + "router.delete('/{brand_kit_id}', status_code=status.HTTP_204_NO_CONTENT)": "/v1/brand DELETE", +} + +EXPECTED_BRAND_FRONTEND_CALLS = [ + "await apiClient.get('/v1/brand');", + "await apiClient.post('/v1/brand', payload);", +] + +EXPECTED_COLLABORATION_FRONTEND_CALLS = [ + "await apiClient.get('/v1/projects/workspace/teams');", + "await apiClient.post('/v1/projects/workspace/teams', payload);", + "await apiClient.post('/v1/projects/workspace/invitations', payload);", + "await apiClient.get('/v1/projects/workspace/members');", + "await apiClient.delete(`/v1/projects/workspace/members/${userId}`);", + "await apiClient.patch(`/v1/projects/workspace/members/${userId}/role?new_role=${newRole}`);", + "await apiClient.get(`/v1/projects/workspace/workflows/${workflowId}/requests`);", + "await apiClient.post(`/v1/projects/workspace/workflows/${workflowId}/requests`, { project_id: projectId });", + "await apiClient.post(`/v1/projects/workspace/requests/${requestId}/approve`);", + "await apiClient.post(`/v1/projects/workspace/requests/${requestId}/reject`);", + "await apiClient.post(`/v1/projects/workspace/requests/${requestId}/comments?content=${encodeURIComponent(content)}`);", + "await apiClient.get(`/v1/projects/${encodeURIComponent(projectId)}/collaborators`);", + "await apiClient.post(`/v1/projects/${encodeURIComponent(projectId)}/collaborators?user_id=${encodeURIComponent(userId)}&role=${encodeURIComponent(role)}`);", + "await apiClient.delete(`/v1/projects/${encodeURIComponent(projectId)}/collaborators/${encodeURIComponent(userId)}`);", +] + + +def _route_decorators(path: Path) -> list[str]: + tree = ast.parse(path.read_text()) + calls = [] + for node in tree.body: + if not isinstance(node, ast.AsyncFunctionDef): + continue + for decorator in node.decorator_list: + if isinstance(decorator, ast.Call) and isinstance(decorator.func, ast.Attribute) and decorator.func.attr in {"get", "post", "patch", "delete"}: + calls.append(ast.unparse(decorator)) + return calls + + +def _frontend_calls(path: Path) -> list[str]: + return [re.sub(r"^\s*const\s+\{[^}]*\}\s+=\s+", "", line.strip()) for line in path.read_text().splitlines() if "apiClient." in line] + + +def test_brand_kit_routes_match_expected_contract() -> None: + assert _route_decorators(ROUTES["brand_api"]) == list(BRAND_ROUTES.keys()) + + +def test_approval_request_create_route_is_exposed() -> None: + decorators = _route_decorators(ROUTES["projects_api"]) + assert any( + decorator == "router.post('/workspace/workflows/{workflow_id}/requests', response_model=ApprovalRequest, status_code=status.HTTP_201_CREATED)" + for decorator in decorators + ) + + +def test_brand_kit_frontend_uses_expected_backend_routes() -> None: + assert _frontend_calls(FRONTEND_CALLS["brand_api"]) == EXPECTED_BRAND_FRONTEND_CALLS + + +def test_collaboration_frontend_uses_expected_backend_routes() -> None: + assert _frontend_calls(FRONTEND_CALLS["collaboration_api"]) == EXPECTED_COLLABORATION_FRONTEND_CALLS + + +class ApiContractRegressionTests(unittest.TestCase): + def test_brand_kit_routes_match_expected_contract(self) -> None: + test_brand_kit_routes_match_expected_contract() + + def test_brand_kit_frontend_uses_expected_backend_routes(self) -> None: + test_brand_kit_frontend_uses_expected_backend_routes() + + def test_collaboration_frontend_uses_expected_backend_routes(self) -> None: + test_collaboration_frontend_uses_expected_backend_routes() + + def test_approval_request_create_route_is_exposed(self) -> None: + test_approval_request_create_route_is_exposed() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_authentication.py b/tests/test_authentication.py new file mode 100644 index 0000000000000000000000000000000000000000..f4db9b10a841ef6545a2f02b5b0f97cdd5ab3fda --- /dev/null +++ b/tests/test_authentication.py @@ -0,0 +1,307 @@ +from __future__ import annotations + +import base64 +import hashlib +from datetime import datetime, timedelta, timezone +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.mcp.registry import MCPRegistry +from app.security.context import auth_context +from app.security.errors import APIKeyConflictError, ForbiddenError, RateLimitError, UnauthorizedError +from app.security.models import APIKey, AuditLog +from app.security.schemas import APIKeyCreate +from app.security.service import APIKeyService +from main import create_app + + +def security_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "auth_enabled": True, + "auth_last_used_update_seconds": 0, + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + "max_workers": 1, + } + values.update(overrides) + return Settings(**values) + + +@pytest.fixture +async def security_container(tmp_path: Path): + container = build_container(security_settings(tmp_path)) + await container.security_database.initialize() + try: + yield container + finally: + await container.security_database.close() + + +async def create_key(container, **overrides: object) -> tuple[APIKey, str]: + values: dict[str, object] = { + "name": "Automation", + "environment": "test", + "role": None, + "scopes": ["templates:read"], + } + values.update(overrides) + return await container.api_keys.create(APIKeyCreate(**values), created_by="tests") + + +async def test_key_generation_has_256_bits_and_database_never_stores_secret( + security_container, +) -> None: + record, secret = await create_key(security_container) + + environment, encoded_secret = secret.split("_", 2)[1:] + raw_secret = base64.urlsafe_b64decode(encoded_secret + "=") + assert environment == "test" + assert len(raw_secret) == 32 + assert record.key_prefix == f"mp_test_{encoded_secret[:8]}" + assert record.key_hash == hashlib.sha256(secret.encode()).hexdigest() + + async with security_container.security_database.session() as session: + stored = await session.get(APIKey, record.id) + assert stored is not None + assert secret not in vars(stored).values() + assert not hasattr(stored, "api_key") + + +async def test_authentication_rejects_invalid_expired_disabled_and_revoked_keys( + security_container, +) -> None: + active, active_secret = await create_key(security_container) + assert (await security_container.api_keys.authenticate(active_secret)).api_key_id == active.id + + replacement = "A" if active_secret[-1] != "A" else "B" + with pytest.raises(UnauthorizedError): + await security_container.api_keys.authenticate(active_secret[:-1] + replacement) + + _, expired_secret = await create_key( + security_container, + name="Expired", + expires_at=datetime.now(timezone.utc) - timedelta(seconds=1), + ) + with pytest.raises(UnauthorizedError): + await security_container.api_keys.authenticate(expired_secret) + + await security_container.api_keys.set_status(active.id, "disabled") + with pytest.raises(UnauthorizedError): + await security_container.api_keys.authenticate(active_secret) + await security_container.api_keys.set_status(active.id, "active") + assert (await security_container.api_keys.authenticate(active_secret)).api_key_id == active.id + + await security_container.api_keys.set_status(active.id, "revoked") + with pytest.raises(UnauthorizedError): + await security_container.api_keys.authenticate(active_secret) + with pytest.raises(APIKeyConflictError): + await security_container.api_keys.set_status(active.id, "disabled") + with pytest.raises(APIKeyConflictError): + await security_container.api_keys.set_status(active.id, "active") + + +async def test_scope_enforcement_and_rotation_grace_period(security_container) -> None: + old, old_secret = await create_key(security_container) + context = await security_container.api_keys.authenticate(old_secret) + security_container.api_keys.authorize(context, "templates:read") + with pytest.raises(ForbiddenError): + security_container.api_keys.authorize(context, "operations:execute") + + replacement, replacement_secret = await security_container.api_keys.rotate( + old.id, 60, created_by="tests" + ) + assert replacement.rotated_from_id == old.id + assert (await security_container.api_keys.authenticate(old_secret)).api_key_id == old.id + assert ( + await security_container.api_keys.authenticate(replacement_secret) + ).api_key_id == replacement.id + with pytest.raises(APIKeyConflictError): + await security_container.api_keys.set_status(old.id, "disabled") + + async with security_container.security_database.session() as session: + rotating = await session.get(APIKey, old.id) + assert rotating is not None + rotating.grace_expires_at = datetime.now(timezone.utc) - timedelta(seconds=1) + await session.commit() + with pytest.raises(UnauthorizedError): + await security_container.api_keys.authenticate(old_secret) + assert (await security_container.api_keys.get(old.id)).status == "revoked" + + +async def test_per_key_request_and_concurrent_job_limits(security_container) -> None: + _, request_secret = await create_key( + security_container, name="RPM", requests_per_minute=1 + ) + request_context = await security_container.api_keys.authenticate(request_secret) + lease = await security_container.rate_limiter.acquire( + request_context, is_job=False, is_upload=False, uploaded_bytes=0 + ) + await lease.release() + with pytest.raises(RateLimitError) as rate_error: + await security_container.rate_limiter.acquire( + request_context, is_job=False, is_upload=False, uploaded_bytes=0 + ) + assert rate_error.value.retry_after >= 1 + + _, job_secret = await create_key( + security_container, name="Concurrency", concurrent_jobs=1 + ) + job_context = await security_container.api_keys.authenticate(job_secret) + running = await security_container.rate_limiter.acquire( + job_context, is_job=True, is_upload=False, uploaded_bytes=0 + ) + with pytest.raises(RateLimitError): + await security_container.rate_limiter.acquire( + job_context, is_job=True, is_upload=False, uploaded_bytes=0 + ) + await running.release() + next_job = await security_container.rate_limiter.acquire( + job_context, is_job=True, is_upload=False, uploaded_bytes=0 + ) + await next_job.release() + + +async def test_stdio_mcp_uses_shared_context_scopes_rate_limits_and_audit( + security_container, +) -> None: + _, secret = await create_key( + security_container, name="MCP Reader", scopes=["mcp:read"] + ) + context = await security_container.api_keys.authenticate(secret) + registry = MCPRegistry(security_container) + unauthorized = await registry.run_metadata_tool("system_info", registry.system_info_data) + token = auth_context.set(context) + try: + resource = await registry.safe_resource("version", registry.version_data) + forbidden = await registry.run_metadata_tool("system_info", registry.system_info_data) + finally: + auth_context.reset(token) + + assert unauthorized["success"] is False + assert unauthorized["error"]["code"] == "UNAUTHORIZED" + assert resource["success"] is True + assert forbidden["success"] is False + assert forbidden["error"]["code"] == "FORBIDDEN" + async with security_container.security_database.session() as session: + logs = list((await session.scalars(select(AuditLog))).all()) + assert {log.endpoint for log in logs} >= { + "mcp://tools/resource.version", + "mcp://tools/system_info", + } + + +def test_http_middleware_public_and_authentication_contracts(tmp_path: Path) -> None: + material = APIKeyService.generate_material("test") + settings = security_settings( + tmp_path, + auth_bootstrap_key_hash=material.key_hash, + auth_bootstrap_key_prefix=material.key_prefix, + auth_bootstrap_environment="test", + auth_default_requests_per_minute=1000, + ) + application = create_app(settings) + authorization = {"Authorization": f"Bearer {material.api_key}"} + + with TestClient(application) as client: + for path in ("/", "/health", "/version", "/docs", "/openapi.json", "/redoc"): + assert client.get(path).status_code == 200 + + missing = client.get("/v1/auth/context") + malformed = client.get( + "/v1/auth/context", headers={"Authorization": "Basic not-a-mediarouter-key"} + ) + invalid = client.get( + "/v1/auth/context", headers={"Authorization": "Bearer mp_test_invalid"} + ) + for response in (missing, malformed, invalid): + assert response.status_code == 401 + assert response.json() == { + "error": "Unauthorized", + "message": "Invalid or expired API key.", + } + assert response.headers["www-authenticate"] == "Bearer" + + mcp_missing = client.post("/mcp/", json={"jsonrpc": "2.0", "id": 1}) + assert mcp_missing.status_code == 401 + + identity = client.get("/v1/auth/context", headers=authorization) + assert identity.status_code == 200 + assert identity.json()["key_prefix"] == material.key_prefix + assert "admin" in identity.json()["scopes"] + + created = client.post( + "/v1/api-keys", + headers=authorization, + json={ + "name": "Template Reader", + "environment": "test", + "role": None, + "scopes": ["templates:read"], + }, + ) + assert created.status_code == 201 + limited_authorization = { + "Authorization": f"Bearer {created.json()['api_key']}" + } + assert client.get("/v1/auth/context", headers=limited_authorization).status_code == 200 + forbidden = client.get("/v1/health", headers=limited_authorization) + assert forbidden.status_code == 403 + assert forbidden.json() == { + "error": "Forbidden", + "message": "Missing required scope.", + } + + mcp_forbidden = client.post( + "/mcp/", + headers=limited_authorization, + json={ + "jsonrpc": "2.0", + "id": 1, + "method": "tools/call", + "params": {"name": "health", "arguments": {}}, + }, + ) + assert mcp_forbidden.status_code == 403 + + audit_logs = client.get("/v1/audit-logs", headers=authorization) + assert audit_logs.status_code == 200 + entries = audit_logs.json() + assert any( + entry["endpoint"] == "/v1/auth/context" + and entry["api_key_id"] == identity.json()["id"] + and entry["response_code"] == 200 + for entry in entries + ) + + +def test_http_rate_limit_returns_retry_after(tmp_path: Path) -> None: + material = APIKeyService.generate_material("test") + application = create_app( + security_settings( + tmp_path, + auth_bootstrap_key_hash=material.key_hash, + auth_bootstrap_key_prefix=material.key_prefix, + auth_bootstrap_environment="test", + auth_default_requests_per_minute=1, + ) + ) + headers = {"Authorization": f"Bearer {material.api_key}"} + with TestClient(application) as client: + assert client.get("/v1/auth/context", headers=headers).status_code == 200 + limited = client.get("/v1/auth/context", headers=headers) + + assert limited.status_code == 429 + assert limited.json() == { + "error": "Rate limit exceeded", + "message": "Retry later.", + } + assert int(limited.headers["retry-after"]) >= 1 diff --git a/tests/test_brand_kits.py b/tests/test_brand_kits.py new file mode 100644 index 0000000000000000000000000000000000000000..9c43927727e09968486c19bdab33df247b35088b --- /dev/null +++ b/tests/test_brand_kits.py @@ -0,0 +1,42 @@ +import pytest +from unittest.mock import AsyncMock, MagicMock +from app.brand.services.brand_service import BrandKitService +from app.brand.services.validation_service import BrandKitValidationService +from app.brand.models.brand import BrandKitVersion + +@pytest.fixture +def validation_service(): + return BrandKitValidationService() + +def test_brand_kit_validation_missing_logo(validation_service): + version = BrandKitVersion(version_number=1, created_by="user1") + result = validation_service.validate(version) + assert not result['valid'] + assert any(issue['field'] == 'logo_asset_id' for issue in result['issues']) + +def test_brand_kit_validation_valid(validation_service): + version = BrandKitVersion(version_number=1, created_by="user1", logo_asset_id="asset123") + result = validation_service.validate(version) + assert result['valid'] + +@pytest.mark.asyncio +async def test_brand_kit_service_create(): + mock_repo = AsyncMock() + mock_assets = AsyncMock() + mock_audit = AsyncMock() + + service = BrandKitService(mock_repo, mock_assets, mock_audit) + + workspace_id = "ws1" + name = "Test Kit" + data = {"logo_asset_id": "asset123"} + user_id = "user1" + + mock_assets.get_asset.return_value = {"id": "asset123"} + mock_repo.create.return_value = (MagicMock(id="kit1"), MagicMock(id="ver1")) + + await service.create_brand_kit(workspace_id, name, data, user_id=user_id) + + mock_assets.get_asset.assert_called_once() + mock_repo.create.assert_called_once() + mock_audit.log_event.assert_called_once() diff --git a/tests/test_cleanup_worker.py b/tests/test_cleanup_worker.py new file mode 100644 index 0000000000000000000000000000000000000000..2e150123f88c2496792c76bb644de1881316e0d0 --- /dev/null +++ b/tests/test_cleanup_worker.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +import asyncio +import os +import time +from uuid import uuid4 + +from app.services.cleanup import CleanupService +from app.workers.cleanup_worker import CleanupWorker + + +async def test_cleanup_removes_expired_workspace(settings) -> None: + service = CleanupService(settings) + request_id = str(uuid4()) + workspace = await service.create_workspace(request_id) + await service.complete(request_id) + old = time.time() - 120 + os.utime(workspace.root, (old, old)) + removed = await service.cleanup_expired() + assert removed == 1 + assert not workspace.root.exists() + + +async def test_cleanup_keeps_active_workspace(settings) -> None: + service = CleanupService(settings) + workspace = await service.create_workspace(str(uuid4())) + old = time.time() - 120 + os.utime(workspace.root, (old, old)) + assert await service.cleanup_expired() == 0 + assert workspace.root.exists() + + +async def test_cleanup_worker_runs_and_stops() -> None: + class FakeCleanup: + def __init__(self) -> None: + self.called = asyncio.Event() + + async def cleanup_expired(self) -> int: + self.called.set() + return 0 + + service = FakeCleanup() + worker = CleanupWorker(service, interval_seconds=60) # type: ignore[arg-type] + await worker.start() + await asyncio.wait_for(service.called.wait(), timeout=1) + await worker.stop() diff --git a/tests/test_collaboration.py b/tests/test_collaboration.py new file mode 100644 index 0000000000000000000000000000000000000000..01b5f7709e66d7020ce6a18611a3837704c6f85e --- /dev/null +++ b/tests/test_collaboration.py @@ -0,0 +1,33 @@ +import pytest +from app.projects.repositories.collaboration_repository import CollaborationRepository +from app.projects.services.collaboration_service import CollaborationService +from app.projects.errors import CollaborationUnauthorizedError + +@pytest.mark.asyncio +async def test_collaboration_logic_admin_removal_constraint(db_session): + # Setup test workspace and admin users + repo = CollaborationRepository(db_session) + service = CollaborationService(repo) + + workspace_id = "test_workspace" + admin_user_id = "admin_user" + target_user_id = "member_user" + + # 1. Mock memberships: 2 admins + # Use actual DB insert here if needed for true integration test + # ... setup DB state ... + + # 2. Test prevention of last admin removal + with pytest.raises(CollaborationUnauthorizedError, match="Cannot remove the last administrator."): + await service.remove_member(workspace_id, admin_user_id, target_user_id) + +@pytest.mark.asyncio +async def test_collaboration_logic_self_elevation_prevention(db_session): + repo = CollaborationRepository(db_session) + service = CollaborationService(repo) + + workspace_id = "test_workspace" + actor_user_id = "user_1" + + with pytest.raises(CollaborationUnauthorizedError, match="Cannot elevate your own privileges."): + await service.update_member_role(workspace_id, actor_user_id, actor_user_id, "admin") diff --git a/tests/test_collaboration_full.py b/tests/test_collaboration_full.py new file mode 100644 index 0000000000000000000000000000000000000000..dd67269230a07a3beb733e3da2ebbbb5f2a0156a --- /dev/null +++ b/tests/test_collaboration_full.py @@ -0,0 +1,40 @@ +import pytest +from app.projects.repositories.collaboration_repository import CollaborationRepository +from app.projects.services.collaboration_service import CollaborationService + +@pytest.mark.asyncio +async def test_collaboration_team_lifecycle(db_session): + repo = CollaborationRepository(db_session) + service = CollaborationService(repo) + + workspace_id = "test_workspace" + + # 1. Create Team + team = await service.create_team(workspace_id, "Engineering") + assert team.name == "Engineering" + + # 2. List Teams + teams = await service.list_teams(workspace_id) + assert len(teams) >= 1 + + # 3. Update Team + updated = await service.update_team(workspace_id, team.id, "Product") + assert updated.name == "Product" + + # 4. Archive + await service.archive_team(workspace_id, team.id) + teams = await service.list_teams(workspace_id) + assert not any(t.id == team.id for t in teams) + +@pytest.mark.asyncio +async def test_collaboration_invitation_lifecycle(db_session): + repo = CollaborationRepository(db_session) + service = CollaborationService(repo) + + workspace_id = "test_workspace" + email = "test@example.com" + + # Test invitation + invitation = await service.invite_member(workspace_id, email, "member") + assert invitation.email == email + assert hasattr(invitation, "token") # Check if token is returned diff --git a/tests/test_content_studio_phase2.py b/tests/test_content_studio_phase2.py new file mode 100644 index 0000000000000000000000000000000000000000..0912073bf6266276b4e66f38beed6795d00ad9e7 --- /dev/null +++ b/tests/test_content_studio_phase2.py @@ -0,0 +1,295 @@ +from __future__ import annotations + +from copy import deepcopy +from pathlib import Path +from uuid import uuid4 + +import pytest +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.projects.editor_schemas import EditorDocument, EditorSaveRequest, ProjectRenderCreate +from app.projects.errors import ( + ProjectEditorConflictError, + ProjectNotFoundError, + ProjectRenderLimitError, +) +from app.projects.schemas import ProjectCreate +from app.projects.services.render_compiler import compile_render +from app.security.models import AuditEvent +from app.security.schemas import APIKeyCreate + + +def settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=True, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + generation_worker_enabled=False, + render_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + whisper_model="tiny", + ) + + +def document(project_id: str, asset_id: str) -> EditorDocument: + return EditorDocument.model_validate( + { + "schemaVersion": 1, + "projectId": project_id, + "timeline": { + "timeUnit": "milliseconds", + "tracks": [ + { + "id": "video-1", + "type": "video", + "name": "Video 1", + "order": 0, + "muted": False, + "locked": False, + "visible": True, + "clips": [ + { + "id": "clip-1", + "kind": "media", + "trackId": "video-1", + "assetId": asset_id, + "label": "source.mp4", + "startMs": 0, + "durationMs": 1000, + "sourceStartMs": 0, + "sourceDurationMs": 1000, + "mediaType": "video", + "transform": { + "x": 0, + "y": 0, + "scaleX": 1, + "scaleY": 1, + "rotation": 0, + }, + "volume": 1, + "opacity": 1, + "visible": True, + "metadata": {}, + } + ], + } + ], + "transitions": [], + "markers": [], + }, + "renderSettings": {"format": "mp4", "width": 1280, "height": 720, "frameRate": 30}, + } + ) + + +async def actor(container, name: str): + key, secret = await container.api_keys.create( + APIKeyCreate( + name=name, + environment="test", + role=None, + scopes=[ + "projects:read", + "projects:create", + "projects:update", + "jobs:create", + "jobs:cancel", + ], + ), + created_by="tests", + ) + return key, await container.api_keys.authenticate(secret) + + +@pytest.mark.asyncio +async def test_editor_revision_isolation_render_idempotency_and_cancellation( + tmp_path: Path, +) -> None: + container = build_container(settings(tmp_path)) + await container.security_database.initialize() + try: + key_a, actor_a = await actor(container, "A") + _, actor_b = await actor(container, "B") + project = await container.projects.create( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Studio"), + ) + request_id = str(uuid4()) + output = container.settings.output_dir / request_id + output.mkdir(parents=True) + source = output / "source.mp4" + source.write_bytes(b"test media") + asset = await container.assets.register_output( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + request_id=request_id, + path=source, + mime_type="video/mp4", + project_id=project.id, + ) + editor_document = document(project.id, asset.id) + saved = await container.editor.save( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + payload=EditorSaveRequest(expected_revision=0, schema_version=1, state=editor_document), + ) + assert saved.revision == 1 + with pytest.raises(ProjectEditorConflictError): + await container.editor.save( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + payload=EditorSaveRequest( + expected_revision=0, schema_version=1, state=editor_document + ), + ) + with pytest.raises(ProjectNotFoundError): + await container.editor.get( + workspace_id=actor_b.workspace_id, + user_id=actor_b.user_id, + project_id=project.id, + ) + render_payload = ProjectRenderCreate( + editor_revision=1, output_format="mp4", width=1280, height=720 + ) + first = await container.renders.create( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + payload=render_payload, + idempotency_key="render-1", + ) + second = await container.renders.create( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + payload=render_payload, + idempotency_key="render-1", + ) + assert first.id == second.id and first.status == "queued" + with pytest.raises(ProjectRenderLimitError): + await container.renders.create( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + payload=render_payload, + idempotency_key="render-2", + ) + cancelled = await container.renders.cancel( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + render_id=first.id, + ) + assert cancelled.status == "cancelled" + repeated = await container.renders.cancel( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project.id, + render_id=first.id, + ) + assert repeated.status == "cancelled" + async with container.security_database.tenant_session( + workspace_id=actor_a.workspace_id, + user_id=actor_a.user_id, + ) as session: + cancellation_events = list( + ( + await session.scalars( + select(AuditEvent).where( + AuditEvent.entity_id == first.id, + AuditEvent.event_type == "project.render_cancelled", + ) + ) + ).all() + ) + assert len(cancellation_events) == 1 + finally: + await container.security_database.close() + + +def test_render_compiler_is_deterministic_and_uses_server_paths(tmp_path: Path) -> None: + source = tmp_path / "source.mp4" + source.write_bytes(b"media") + state = document(str(uuid4()), str(uuid4())) + asset_id = next(iter(state.asset_ids())) + first = compile_render( + state, + asset_paths={asset_id: (source, "video/mp4")}, + width=1280, + height=720, + frame_rate=30, + output_format="mp4", + quality="standard", + preset="balanced", + ) + second = compile_render( + state, + asset_paths={asset_id: (source, "video/mp4")}, + width=1280, + height=720, + frame_rate=30, + output_format="mp4", + quality="standard", + preset="balanced", + ) + assert first == second + assert source in first.args + assert first.duration_ms == 1000 + assert "yuv420p" in first.args + + +def test_render_compiler_ignores_hidden_timeline_tail(tmp_path: Path) -> None: + source = tmp_path / "source.mp4" + source.write_bytes(b"media") + state = document(str(uuid4()), str(uuid4())) + payload = state.model_dump(by_alias=True) + hidden_track = deepcopy(payload["timeline"]["tracks"][0]) + hidden_track.update({"id": "video-hidden", "name": "Hidden", "order": 1, "visible": False}) + hidden_track["clips"][0].update( + {"id": "clip-hidden", "trackId": "video-hidden", "startMs": 120_000} + ) + payload["timeline"]["tracks"].append(hidden_track) + state_with_hidden_tail = EditorDocument.model_validate(payload) + asset_id = next(iter(state_with_hidden_tail.asset_ids())) + + plan = compile_render( + state_with_hidden_tail, + asset_paths={asset_id: (source, "video/mp4")}, + width=1280, + height=720, + frame_rate=30, + output_format="webm", + quality="high", + preset="quality", + ) + + assert plan.duration_ms == 1000 + assert plan.args.count(source) == 1 + assert "18" in plan.args + assert "0" in plan.args diff --git a/tests/test_cors.py b/tests/test_cors.py new file mode 100644 index 0000000000000000000000000000000000000000..32b76493e76fc2262eb3ca40a2e1b21d78bd0d11 --- /dev/null +++ b/tests/test_cors.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +from pathlib import Path + +import httpx + +from app.core.config import Settings +from main import create_app + + +async def test_configured_frontend_origin_receives_cors_headers(tmp_path: Path) -> None: + settings = Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cors_allowed_origins="https://workspace.example.vercel.app", + ) + app = create_app(settings) + transport = httpx.ASGITransport(app=app) + + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + response = await client.options( + "/v1/projects", + headers={ + "Origin": "https://workspace.example.vercel.app", + "Access-Control-Request-Method": "GET", + "Access-Control-Request-Headers": "Authorization", + }, + ) + + assert response.status_code == 200 + assert response.headers["access-control-allow-origin"] == ( + "https://workspace.example.vercel.app" + ) + assert "authorization" in response.headers["access-control-allow-headers"].lower() + + +async def test_unconfigured_origin_receives_no_cors_authorization(tmp_path: Path) -> None: + settings = Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cors_allowed_origins="https://workspace.example.vercel.app", + ) + app = create_app(settings) + transport = httpx.ASGITransport(app=app) + + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + response = await client.options( + "/v1/projects", + headers={ + "Origin": "https://attacker.example", + "Access-Control-Request-Method": "GET", + }, + ) + + assert response.status_code == 400 + assert "access-control-allow-origin" not in response.headers diff --git a/tests/test_database_migration_contracts.py b/tests/test_database_migration_contracts.py new file mode 100644 index 0000000000000000000000000000000000000000..9283a81cd091083683fd448de01fc5f38ea0616e --- /dev/null +++ b/tests/test_database_migration_contracts.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import re +import unittest +from pathlib import Path + + +MIGRATION_FILES = sorted(Path("app/projects/migrations").glob("*.sql")) +ALLOWED_TABLES = { + "teams", + "team_members", + "invitations", + "project_collaborators", + "approval_workflows", + "approval_requests", + "review_comments", + "collaboration_activity", + "notification_preferences", +} + + +class MigrationContractTests(unittest.TestCase): + def test_latest_migration_adds_approval_request_workspace_integrity(self) -> None: + latest = MIGRATION_FILES[-1].read_text() + + self.assertIn("alter table approval_requests", latest) + self.assertIn("add column if not exists workspace_id text", latest) + self.assertIn("alter column workspace_id set not null", latest) + self.assertIn("create index if not exists ix_approval_requests_workspace", latest) + + def test_migrations_do_not_recreate_shared_tables(self) -> None: + counts: dict[str, int] = {} + for path in MIGRATION_FILES: + body = path.read_text() + for table in ALLOWED_TABLES: + counts[table] = counts.get(table, 0) + body.count(f"create table if not exists {table}") + + for table, count in counts.items(): + self.assertEqual( + count, + 1, + f"Duplicate table creation detected for {table}: {count} migrations recreate it", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_database_url.py b/tests/test_database_url.py new file mode 100644 index 0000000000000000000000000000000000000000..4d943c723ad4f373ca18103a4b74fa2828cbed53 --- /dev/null +++ b/tests/test_database_url.py @@ -0,0 +1,16 @@ +from app.core.database_url import normalize_async_database_url + + +def test_bare_postgres_urls_use_the_installed_async_driver() -> None: + assert ( + normalize_async_database_url("postgresql://user:secret@db.example/app") + == "postgresql+asyncpg://user:secret@db.example/app" + ) + assert ( + normalize_async_database_url("postgres://user:secret@db.example/app") + == "postgresql+asyncpg://user:secret@db.example/app" + ) + assert ( + normalize_async_database_url("postgresql+asyncpg://user:secret@db.example/app") + == "postgresql+asyncpg://user:secret@db.example/app" + ) diff --git a/tests/test_downloader.py b/tests/test_downloader.py new file mode 100644 index 0000000000000000000000000000000000000000..2f5292f2dd457101b4cf65a3273a44f259bfbefb --- /dev/null +++ b/tests/test_downloader.py @@ -0,0 +1,20 @@ +from unittest.mock import AsyncMock + +import respx +from httpx import Response + +from app.services.downloader import Downloader +from app.services.validator import MediaValidator + + +@respx.mock +async def test_url_download_streams_to_disk(settings, tmp_path) -> None: + url = "https://media.example.test/sample.mp3" + respx.get(url).mock( + return_value=Response(200, content=b"ID3data", headers={"content-type": "audio/mpeg"}) + ) + downloader = Downloader(settings, MediaValidator(settings)) + downloader.validate_url = AsyncMock() # type: ignore[method-assign] + path, mime_type = await downloader.download(url, tmp_path) + assert path.read_bytes() == b"ID3data" + assert mime_type == "audio/mpeg" diff --git a/tests/test_error_handling.py b/tests/test_error_handling.py new file mode 100644 index 0000000000000000000000000000000000000000..2fbbf7e2cd60e6817f0b308545e25028e78a8b49 --- /dev/null +++ b/tests/test_error_handling.py @@ -0,0 +1,27 @@ +import pytest +from fastapi.testclient import TestClient + +from app.core.exceptions import NotFoundError +from main import create_app + + +def test_errors_use_safe_standard_envelope(settings) -> None: + with TestClient(create_app(settings), raise_server_exceptions=False) as client: + response = client.post( + "/v1/probe", + json={"base64": "not-valid-base64!", "filename": "sample.mp3"}, + ) + assert response.status_code == 422 + payload = response.json() + assert payload["success"] is False + assert payload["request_id"] + assert payload["error"]["code"] == "INVALID_INPUT" + assert "traceback" not in response.text.lower() + + +def test_download_path_traversal_is_rejected(settings) -> None: + app = create_app(settings) + with pytest.raises(NotFoundError): + app.state.container.cleanup.resolve_download( + "00000000-0000-0000-0000-000000000000", "../secret.mp4" + ) diff --git a/tests/test_ffmpeg_operations.py b/tests/test_ffmpeg_operations.py new file mode 100644 index 0000000000000000000000000000000000000000..f06e8d3ecd9d8b10811f9f6d8a6aca840cdd0612 --- /dev/null +++ b/tests/test_ffmpeg_operations.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import shutil +import subprocess + +import pytest + +from app.models.media import InputMedia, MediaSource +from app.operations.convert import convert_audio +from app.services.ffmpeg_service import FFmpegService + + +@pytest.mark.skipif(shutil.which("ffmpeg") is None, reason="ffmpeg is not installed") +async def test_ffmpeg_audio_conversion(settings, tmp_path) -> None: + source = tmp_path / "tone.wav" + subprocess.run( + [ + "ffmpeg", + "-hide_banner", + "-loglevel", + "error", + "-f", + "lavfi", + "-i", + "sine=frequency=440:duration=0.2", + "-y", + str(source), + ], + check=True, + ) + media = InputMedia( + source=MediaSource.MULTIPART, + filename=source.name, + mime_type="audio/wav", + temp_path=source, + size=source.stat().st_size, + ) + result = await convert_audio( + FFmpegService(settings), [media], {"format": "mp3"}, tmp_path / "out" + ) + assert result.path is not None + assert result.path.is_file() + assert result.path.stat().st_size > 0 + + +async def test_ffmpeg_codec_listing_is_structured(settings, monkeypatch) -> None: + service = FFmpegService(settings) + + async def fake_capture(*args, **kwargs) -> str: + return """Codecs: + D..... = Decoding supported + .E.... = Encoding supported + ------- + DEV.LS h264 H.264 / AVC / MPEG-4 AVC + DEA.L. aac AAC (Advanced Audio Coding) +""" + + monkeypatch.setattr(service, "_capture", fake_capture) + + codecs = await service.codecs() + + assert [codec["name"] for codec in codecs] == ["h264", "aac"] + assert codecs[0]["decode"] is True + assert codecs[0]["encode"] is True + assert codecs[0]["type"] == "video" diff --git a/tests/test_ffprobe.py b/tests/test_ffprobe.py new file mode 100644 index 0000000000000000000000000000000000000000..172d2f5eb355505319dbf084404decfd3d11783c --- /dev/null +++ b/tests/test_ffprobe.py @@ -0,0 +1,21 @@ +from __future__ import annotations + +import shutil +import wave + +import pytest + +from app.services.ffprobe_service import FFprobeService + + +@pytest.mark.skipif(shutil.which("ffprobe") is None, reason="ffprobe is not installed") +async def test_ffprobe_returns_audio_metadata(settings, tmp_path) -> None: + audio = tmp_path / "tone.wav" + with wave.open(str(audio), "wb") as stream: + stream.setnchannels(1) + stream.setsampwidth(2) + stream.setframerate(8000) + stream.writeframes(b"\x00\x00" * 8000) + metadata = await FFprobeService(settings).probe(audio) + assert metadata["duration"] == pytest.approx(1.0, abs=0.01) + assert metadata["audio_streams"][0]["codec"] == "pcm_s16le" diff --git a/tests/test_generation_flux.py b/tests/test_generation_flux.py new file mode 100644 index 0000000000000000000000000000000000000000..7e0a3b1872668a04e2821cd7846d08310e20af26 --- /dev/null +++ b/tests/test_generation_flux.py @@ -0,0 +1,324 @@ +"""Mocked protocol tests for the audited FLUX.2 Klein worker integration.""" + +from __future__ import annotations + +from collections.abc import Callable +from pathlib import Path +from types import SimpleNamespace + +import httpx +import pytest + +from app.core.config import Settings +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerJobStatus, +) +from app.generation.domain.errors import ( + GenerationCapabilityUnsupportedError, + GenerationValidationError, + GenerationWorkerError, +) +from app.generation.domain.retry import GenerationRetryPolicy +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_MODEL_ID, + FLUX_PROVIDER_ID, + FluxProviderAdapter, +) +from app.generation.providers.worker_client import RemoteWorkerClient +from app.generation.schemas.requests import GenerationRequestCreate + + +async def _no_sleep() -> None: + return None + + +def _client( + handler: Callable[[httpx.Request], httpx.Response], *, retries: int = 2 +) -> RemoteWorkerClient: + return RemoteWorkerClient( + base_url="https://flux-worker.example", + bearer_token="x" * 32, + connect_timeout_seconds=1, + request_timeout_seconds=1, + read_timeout_seconds=1, + retry_policy=GenerationRetryPolicy(max_retries=retries, backoff_seconds=0), + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + sleep=lambda _: _no_sleep(), + ) + + +def _info() -> dict[str, object]: + return { + "id": FLUX_MODEL_ID, + "name": "FLUX.2 Klein 4B", + "type": "image", + "license": "Apache-2.0", + "status": "ready", + "models": {"distilled": FLUX_DISTILLED_MODEL_ID, "base": FLUX_BASE_MODEL_ID}, + } + + +def _payload(**overrides: object) -> GenerationRequestCreate: + value: dict[str, object] = { + "provider": FLUX_PROVIDER_ID, + "model_id": FLUX_MODEL_ID, + "modality": "image", + "prompt": "A cinematic coastal city at sunrise", + "flux": { + "mode_choice": "Distilled (4 steps)", + "seed": 42, + "randomize_seed": False, + "width": 1024, + "height": 1024, + "num_inference_steps": 4, + "guidance_scale": 1.0, + "prompt_upsampling": False, + }, + } + value.update(overrides) + return GenerationRequestCreate.model_validate(value) + + +@pytest.mark.asyncio +async def test_flux_exact_model_discovery_and_readiness() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/health": + return httpx.Response(200, json={"status": "ok"}) + if request.url.path == "/ready": + return httpx.Response( + 200, + json={ + "status": "ready", + "model_loaded": True, + "model": FLUX_MODEL_ID, + "accepting_jobs": True, + }, + ) + return httpx.Response(200, json=_info()) + + adapter = FluxProviderAdapter(client=_client(handler)) + registry = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id=FLUX_PROVIDER_ID, + model=FLUX_MODEL_CAPABILITY, + configuration_reference="flux-space", + ) + ] + ) + assert (await adapter.health()).status.value == "healthy" + models = registry.verify_readiness( + provider_id=FLUX_PROVIDER_ID, + worker_info=await adapter.info(), + readiness=await adapter.ready(), + provider_configured=adapter.available, + ) + assert models[0].model.id == FLUX_MODEL_ID + assert models[0].model.modality is GenerationModality.IMAGE + assert models[0].available + + +@pytest.mark.asyncio +async def test_flux_identity_mismatch_and_not_ready_are_not_advertised() -> None: + wrong = {**_info(), "models": {"distilled": "untrusted/model", "base": FLUX_BASE_MODEL_ID}} + + def identity_handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=wrong) + + with pytest.raises(GenerationWorkerError) as raised: + await FluxProviderAdapter(client=_client(identity_handler)).info() + assert raised.value.category is WorkerErrorCategory.PROVIDER_ERROR + + def not_ready_handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/ready": + return httpx.Response( + 503, + json={ + "status": "not_ready", + "model_loaded": False, + "model": FLUX_MODEL_ID, + "accepting_jobs": False, + }, + ) + return httpx.Response(200, json=_info()) + + with pytest.raises(GenerationWorkerError) as raised: + await FluxProviderAdapter(client=_client(not_ready_handler)).ready() + assert raised.value.category is WorkerErrorCategory.WORKER_NOT_READY + + +@pytest.mark.asyncio +async def test_flux_text_submission_uses_strict_form_and_has_no_automatic_retry() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(202, json={"job_id": "flux_" + "a" * 32, "status": "queued"}) + + job = await FluxProviderAdapter(client=_client(handler)).submit( + payload={"prompt": "A city at sunrise", "flux": {"width": 1024, "height": 1024}}, + idempotency_key="generation-request-id", + ) + assert job.status is WorkerJobStatus.QUEUED + assert requests[0].headers["authorization"] == "Bearer " + "x" * 32 + assert requests[0].headers["content-type"].startswith("application/x-www-form-urlencoded") + assert b"width=1024" in requests[0].content + + +@pytest.mark.asyncio +async def test_flux_optional_canonical_image_uses_multipart(tmp_path: Path) -> None: + source = tmp_path / "input.png" + source.write_bytes(b"image-input") + + def handler(request: httpx.Request) -> httpx.Response: + body = request.content.decode("latin-1") + assert 'name="input_images"' in body + assert 'name="prompt"' in body + return httpx.Response(202, json={"job_id": "flux_" + "b" * 32, "status": "queued"}) + + job = await FluxProviderAdapter(client=_client(handler)).submit( + payload={"prompt": "Edit this image"}, + idempotency_key="generation-request-id", + input_path=source, + input_mime_type="image/png", + ) + assert job.external_job_id.startswith("flux_") + + +@pytest.mark.asyncio +async def test_flux_rejects_invalid_requests_and_input_assets() -> None: + adapter = FluxProviderAdapter(client=None) + for invalid in ({"prompt": " "}, {"modality": "video"}): + with pytest.raises(Exception): + await adapter.validate_request(_payload(**invalid)) + + with pytest.raises(GenerationValidationError): + await adapter.validate_input_asset( + _payload(), SimpleNamespace(mime_type="video/mp4", file_size=100) + ) + with pytest.raises(GenerationValidationError): + await adapter.validate_input_asset( + _payload(), SimpleNamespace(mime_type="image/png", file_size=21 * 1024 * 1024) + ) + + +@pytest.mark.parametrize( + "field,value", + [ + ("negative_prompt", "unsupported"), + ("scheduler", "unsupported"), + ("width", 1023), + ("height", 1032), + ], +) +def test_flux_schema_rejects_unsupported_or_invalid_parameters(field: str, value: object) -> None: + raw = _payload().model_dump() + flux = dict(raw["flux"] or {}) + flux[field] = value + raw["flux"] = flux + with pytest.raises(ValueError): + GenerationRequestCreate.model_validate(raw) + + +@pytest.mark.asyncio +async def test_flux_rejects_controls_for_another_provider() -> None: + payload = _payload(wan={"duration_seconds": 1.0}) + with pytest.raises(GenerationCapabilityUnsupportedError): + await FluxProviderAdapter(client=None).validate_request(payload) + + +@pytest.mark.asyncio +async def test_flux_completed_job_maps_a_safe_png_output_and_retrieves_it() -> None: + job_id = "flux_" + "c" * 32 + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/output"): + return httpx.Response(200, content=b"png-output") + return httpx.Response( + 200, + json={ + "job_id": job_id, + "status": "completed", + "output": {"type": "image", "filename": "output.png"}, + }, + ) + + adapter = FluxProviderAdapter(client=_client(handler)) + job = await adapter.get_job(external_job_id=job_id) + assert job.output is not None + assert job.output.mime_type == "image/png" + assert job.output.download_path == f"/v1/jobs/{job_id}/output" + output = await adapter.retrieve_output(external_job_id=job_id) + async with adapter.stream_output(output) as chunks: + assert b"".join([chunk async for chunk in chunks]) == b"png-output" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [429, 502, 503, 504]) +async def test_flux_polling_uses_shared_bounded_transient_retry(status_code: int) -> None: + calls = 0 + job_id = "flux_" + "d" * 32 + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + if calls < 3: + return httpx.Response(status_code, json={"detail": {"token": "never-store"}}) + return httpx.Response(200, json={"job_id": job_id, "status": "running"}) + + job = await FluxProviderAdapter(client=_client(handler, retries=2)).get_job( + external_job_id=job_id + ) + assert job.status is WorkerJobStatus.RUNNING + assert calls == 3 + + +@pytest.mark.asyncio +async def test_flux_permanent_error_is_not_retried_and_cancellation_is_accurate() -> None: + job_id = "flux_" + "e" * 32 + calls = 0 + + def permanent_handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(400, json={"detail": {"code": "FLUX_REQUEST_INVALID"}}) + + with pytest.raises(GenerationWorkerError) as raised: + await FluxProviderAdapter(client=_client(permanent_handler, retries=3)).get_job( + external_job_id=job_id + ) + assert raised.value.category is WorkerErrorCategory.INVALID_REQUEST + assert calls == 1 + + def queued_handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"job_id": job_id, "status": "cancelled"}) + + def running_handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 409, + json={"detail": {"code": "FLUX_JOB_NOT_CANCELLABLE", "status": "running"}}, + ) + + assert ( + await FluxProviderAdapter(client=_client(queued_handler)).cancel(external_job_id=job_id) + ).status is WorkerCancellationStatus.CANCELLED + assert ( + await FluxProviderAdapter(client=_client(running_handler)).cancel(external_job_id=job_id) + ).status is WorkerCancellationStatus.FAILED + + +def test_flux_configuration_is_optional_and_does_not_change_wan_configuration() -> None: + disabled = FluxProviderAdapter.from_settings(Settings(_env_file=None)) + invalid = FluxProviderAdapter.from_settings( + Settings(_env_file=None, flux_space_url="https://flux-worker.example") + ) + assert not disabled.available + assert not invalid.available + assert invalid.configuration_error is not None diff --git a/tests/test_generation_foundation.py b/tests/test_generation_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..10348f0b710aa4a7a91803af805c7ecc4097c414 --- /dev/null +++ b/tests/test_generation_foundation.py @@ -0,0 +1,509 @@ +from __future__ import annotations + +import base64 +from contextlib import asynccontextmanager +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient + +from app.container import build_container +from app.ai.schemas import AiGenerateImageRequest +from app.core.config import Settings +from app.generation.domain.capabilities import ( + GenerationModelCapability, + GenerationProviderCapabilities, +) +from app.generation.domain.enums import ( + GenerationJobStatus, + GenerationModality, + WorkerCancellationStatus, + WorkerHealthStatus, + WorkerJobStatus, + WorkerReadinessStatus, +) +from app.generation.domain.errors import ( + GenerationIdempotencyConflictError, + GenerationInputAssetNotFoundError, + GenerationJobNotFoundError, + GenerationProviderJobConflictError, +) +from app.generation.domain.runtime import ( + WorkerCancellationResult, + WorkerHealth, + WorkerInfo, + WorkerJob, + WorkerOutput, + WorkerReadiness, +) +from app.generation.model_registry import ( + GenerationModelRegistration, + GenerationModelRegistry, +) +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.providers.registry import GenerationProviderRegistry +from app.generation.schemas.requests import GenerationRequestCreate +from app.security.schemas import APIKeyCreate +from main import create_app + + +def generation_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=True, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + generation_enabled=True, + ) + + +class AvailableTestProvider(GenerationProviderAdapter): + capabilities = GenerationProviderCapabilities( + provider="test-generation", + name="Test generation adapter", + implementation_status="test", + models=[ + GenerationModelCapability( + id="test-image-v1", + name="Test image v1", + modality=GenerationModality.IMAGE, + input_asset_supported=True, + ) + ], + ) + + def __init__(self) -> None: + self.cancellation_result = WorkerCancellationResult( + status=WorkerCancellationStatus.REQUESTED + ) + + @property + def available(self) -> bool: + return True + + async def validate_request(self, payload: GenerationRequestCreate) -> dict[str, object]: + return {"prompt": payload.prompt} + + async def health(self) -> WorkerHealth: + return WorkerHealth(status=WorkerHealthStatus.HEALTHY) + + async def info(self) -> WorkerInfo: + return WorkerInfo( + id="test-generation-worker", + name="Test generation worker", + media_types=[GenerationModality.IMAGE], + models=[ + { + "id": "test-image-v1", + "name": "Test image v1", + "media_types": [GenerationModality.IMAGE], + } + ], + status=WorkerHealthStatus.HEALTHY, + ) + + async def ready(self) -> WorkerReadiness: + return WorkerReadiness( + status=WorkerReadinessStatus.READY, + model_loaded=True, + model_ids=["test-image-v1"], + ) + + async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult: + assert external_job_id == "worker-job-1" + return self.cancellation_result + + async def get_job(self, *, external_job_id: str) -> WorkerJob: + assert external_job_id == "worker-job-1" + return WorkerJob( + external_job_id=external_job_id, + status=WorkerJobStatus.COMPLETED, + output=WorkerOutput( + output_type=GenerationModality.IMAGE, + mime_type="image/png", + provider_output_id="worker-output-1", + download_path="/v1/outputs/worker-output-1", + ), + ) + + @asynccontextmanager + async def stream_output(self, output: WorkerOutput): + assert output.provider_output_id == "worker-output-1" + + async def chunks(): + yield base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ" + "AAAADUlEQVQIHWP4z8DwHwAFgAI/ScL9aQAAAABJRU5ErkJggg==" + ) + + yield chunks() + + +async def create_context(container, name: str): + _, secret = await container.api_keys.create( + APIKeyCreate( + name=name, + environment="test", + role=None, + scopes=[ + "generation:providers:read", + "generation:requests:read", + "generation:requests:create", + "generation:jobs:cancel", + ], + ), + created_by="tests", + ) + return await container.api_keys.authenticate(secret) + + +def request_payload(*, prompt: str = "A test image") -> GenerationRequestCreate: + return GenerationRequestCreate( + provider="test-generation", + model_id="test-image-v1", + modality=GenerationModality.IMAGE, + prompt=prompt, + ) + + +@pytest.fixture +async def generation_container(tmp_path: Path): + container = build_container(generation_settings(tmp_path)) + await container.security_database.initialize() + provider = AvailableTestProvider() + container.generation.providers = GenerationProviderRegistry([provider]) + container.generation.models = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id=provider.provider, + model=provider.capabilities.models[0], + configuration_reference="test-generation-worker", + ) + ] + ) + await container.generation.initialize() + await container.generation.refresh_provider_runtime(provider.provider) + try: + yield container + finally: + await container.security_database.close() + + +@pytest.mark.asyncio +async def test_optional_generation_providers_start_unavailable_without_configuration( + tmp_path: Path, +) -> None: + container = build_container(generation_settings(tmp_path)) + await container.security_database.initialize() + await container.generation.initialize() + try: + providers = container.generation.list_providers() + assert [provider.capabilities.provider for provider in providers] == ["flux", "wan"] + assert not any(provider.available for provider in providers) + assert not container.generation.get_model("flux", "flux.2-klein-4b").available + assert not container.generation.get_model("wan", "wan2.2").available + finally: + await container.security_database.close() + + +@pytest.mark.asyncio +async def test_ai_studio_advertises_and_isolates_real_generation_history( + generation_container, +) -> None: + context = await create_context(generation_container, "AI Studio") + capabilities = generation_container.ai.capabilities() + image_tool = next(tool for tool in capabilities.tools if tool.operation == "generate_image") + video_tool = next(tool for tool in capabilities.tools if tool.operation == "generate_video") + assert image_tool.available + assert not video_tool.available + + ordinary = await generation_container.generation.create( + workspace_id=context.workspace_id, + user_id=context.user_id, + payload=request_payload(prompt="ordinary generation"), + idempotency_key="ordinary-generation-key", + ) + ai_job = await generation_container.ai.create( + workspace_id=context.workspace_id, + user_id=context.user_id, + api_key_id=context.api_key_id, + request_id="ai-request", + payload=AiGenerateImageRequest( + operation="generate_image", + prompt="AI Studio generation", + ), + idempotency_key="ai-studio-generation-key", + ) + history = await generation_container.ai.history( + workspace_id=context.workspace_id, + user_id=context.user_id, + offset=0, + limit=25, + ) + assert [item.generation_id for item in history.items] == [ai_job.generation_id] + assert ordinary.id not in {item.generation_id for item in history.items} + + +def test_application_starts_with_optional_providers_disabled_when_unconfigured( + tmp_path: Path, +) -> None: + """No worker URL/token is needed merely to start the application.""" + + with TestClient(create_app(generation_settings(tmp_path))) as client: + providers = client.app.state.container.generation.list_providers() + assert [provider.capabilities.provider for provider in providers] == ["flux", "wan"] + models = client.app.state.container.generation.list_models() + assert [model.model.id for model in models] == ["flux.2-klein-4b", "wan2.2"] + assert not any(model.available for model in models) + + +@pytest.mark.asyncio +async def test_provider_discovery_requires_a_verified_model(generation_container) -> None: + """A configured adapter is not publicly usable before runtime verification.""" + + provider_id = "test-generation" + generation_container.generation.models.mark_unavailable(provider_id) + assert not generation_container.generation.get_provider(provider_id).available + assert not generation_container.generation.list_providers()[0].available + + await generation_container.generation.refresh_provider_runtime(provider_id) + assert generation_container.generation.get_provider(provider_id).available + + +@pytest.mark.asyncio +async def test_generation_request_idempotency_and_cancel(generation_container) -> None: + context = await create_context(generation_container, "Generation A") + workspace_id = str(context.workspace_id) + user_id = str(context.user_id) + + first = await generation_container.generation.create( + workspace_id=workspace_id, + user_id=user_id, + payload=request_payload(), + idempotency_key="generation-request-key", + ) + replay = await generation_container.generation.create( + workspace_id=workspace_id, + user_id=user_id, + payload=request_payload(), + idempotency_key="generation-request-key", + ) + assert replay.id == first.id + assert replay.job.id == first.job.id + with pytest.raises(GenerationIdempotencyConflictError): + await generation_container.generation.create( + workspace_id=workspace_id, + user_id=user_id, + payload=request_payload(prompt="Different request"), + idempotency_key="generation-request-key", + ) + + cancelled = await generation_container.generation.cancel(workspace_id, user_id, first.job.id) + assert cancelled.status is GenerationJobStatus.CANCELLED + retrieved = await generation_container.generation.get_request(workspace_id, user_id, first.id) + assert retrieved.status is GenerationJobStatus.CANCELLED + + +@pytest.mark.asyncio +async def test_generation_records_are_workspace_isolated(generation_container) -> None: + context_a = await create_context(generation_container, "Generation A") + context_b = await create_context(generation_container, "Generation B") + created = await generation_container.generation.create( + workspace_id=str(context_a.workspace_id), + user_id=str(context_a.user_id), + payload=request_payload(), + idempotency_key="generation-isolation-key", + ) + with pytest.raises(GenerationJobNotFoundError): + await generation_container.generation.get_job( + str(context_b.workspace_id), str(context_b.user_id), created.job.id + ) + assert ( + await generation_container.generation.list_requests( + str(context_b.workspace_id), str(context_b.user_id) + ) + == [] + ) + + +@pytest.mark.asyncio +async def test_generation_rejects_another_workspace_canonical_input_asset( + generation_container, +) -> None: + context_a = await create_context(generation_container, "Generation A") + context_b = await create_context(generation_container, "Generation B") + request_id = "00000000-0000-0000-0000-000000000010" + output_dir = generation_container.settings.output_dir / request_id + output_dir.mkdir(parents=True) + output = output_dir / "owned-input.png" + output.write_bytes(b"canonical image") + asset = await generation_container.assets.register_output( + workspace_id=str(context_a.workspace_id), + user_id=str(context_a.user_id), + request_id=request_id, + path=output, + mime_type="image/png", + ) + with pytest.raises(GenerationInputAssetNotFoundError): + await generation_container.generation.create( + workspace_id=str(context_b.workspace_id), + user_id=str(context_b.user_id), + payload=GenerationRequestCreate( + provider="test-generation", + model_id="test-image-v1", + modality=GenerationModality.IMAGE, + prompt="Use another workspace asset", + input_asset_id=asset.id, + ), + idempotency_key="generation-cross-asset-key", + ) + + +@pytest.mark.parametrize("forbidden_field", ["provider_payload", "worker_url", "output_url"]) +def test_generation_request_schema_rejects_client_supplied_provider_controls( + forbidden_field: str, +) -> None: + payload: dict[str, object] = { + "provider": "test-generation", + "model_id": "test-image-v1", + "modality": "image", + "prompt": "A test image", + } + payload[forbidden_field] = {"unsafe": True} + with pytest.raises(ValueError): + GenerationRequestCreate.model_validate(payload) + + +@pytest.mark.asyncio +async def test_remote_cancellation_preserves_requested_and_confirmed_states( + generation_container, +) -> None: + context = await create_context(generation_container, "Generation cancellation") + workspace_id = str(context.workspace_id) + user_id = str(context.user_id) + created = await generation_container.generation.create( + workspace_id=workspace_id, + user_id=user_id, + payload=request_payload(), + idempotency_key="generation-cancellation-key", + ) + await generation_container.generation.repository.transition_job( + workspace_id, + created.job.id, + GenerationJobStatus.SUBMITTING, + user_id=user_id, + ) + await generation_container.generation.bind_provider_job( + workspace_id=workspace_id, + user_id=user_id, + job_id=created.job.id, + worker_job_id="worker-job-1", + ) + await generation_container.generation.repository.transition_job( + workspace_id, + created.job.id, + GenerationJobStatus.RUNNING, + user_id=user_id, + ) + + requested = await generation_container.generation.cancel(workspace_id, user_id, created.job.id) + assert requested.status is GenerationJobStatus.CANCEL_REQUESTED + + provider = generation_container.generation.providers.get("test-generation") + assert isinstance(provider, AvailableTestProvider) + provider.cancellation_result = WorkerCancellationResult( + status=WorkerCancellationStatus.CANCELLED + ) + confirmed = await generation_container.generation.cancel(workspace_id, user_id, created.job.id) + assert confirmed.status is GenerationJobStatus.CANCELLED + + +@pytest.mark.asyncio +async def test_provider_job_binding_and_output_ingestion_are_workspace_scoped( + generation_container, +) -> None: + context_a = await create_context(generation_container, "Generation output A") + context_b = await create_context(generation_container, "Generation output B") + workspace_a, user_a = str(context_a.workspace_id), str(context_a.user_id) + workspace_b, user_b = str(context_b.workspace_id), str(context_b.user_id) + job_a = await generation_container.generation.create( + workspace_id=workspace_a, + user_id=user_a, + payload=request_payload(), + idempotency_key="generation-output-a", + ) + job_b = await generation_container.generation.create( + workspace_id=workspace_b, + user_id=user_b, + payload=request_payload(), + idempotency_key="generation-output-b", + ) + for workspace_id, user_id, job_id in ( + (workspace_a, user_a, job_a.job.id), + (workspace_b, user_b, job_b.job.id), + ): + await generation_container.generation.repository.transition_job( + workspace_id, + job_id, + GenerationJobStatus.SUBMITTING, + user_id=user_id, + ) + + await generation_container.generation.bind_provider_job( + workspace_id=workspace_a, + user_id=user_a, + job_id=job_a.job.id, + worker_job_id="worker-job-1", + ) + with pytest.raises(GenerationProviderJobConflictError): + await generation_container.generation.bind_provider_job( + workspace_id=workspace_b, + user_id=user_b, + job_id=job_b.job.id, + worker_job_id="worker-job-1", + ) + + await generation_container.generation.repository.transition_job( + workspace_a, + job_a.job.id, + GenerationJobStatus.RUNNING, + user_id=user_a, + ) + completed = await generation_container.generation.ingest_completed_provider_output( + workspace_id=workspace_a, + user_id=user_a, + job_id=job_a.job.id, + ) + assert completed.status is GenerationJobStatus.SUCCEEDED + assert completed.output_asset_id is not None + output_asset = await generation_container.assets.get_owned_by_id( + workspace_id=workspace_a, + user_id=user_a, + asset_id=completed.output_asset_id, + ) + assert output_asset.mime_type == "image/png" + assert output_asset.metadata_json["generation"]["media"]["resolution"] == { + "width": 1, + "height": 1, + } + assert ( + await generation_container.generation.ingest_completed_provider_output( + workspace_id=workspace_a, + user_id=user_a, + job_id=job_a.job.id, + ) + == completed + ) + with pytest.raises(GenerationJobNotFoundError): + await generation_container.generation.ingest_completed_provider_output( + workspace_id=workspace_b, + user_id=user_b, + job_id=job_a.job.id, + ) diff --git a/tests/test_generation_provider_runtime.py b/tests/test_generation_provider_runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..0149304aee33afbb32f1ca3cd00f7fa0c9bc7366 --- /dev/null +++ b/tests/test_generation_provider_runtime.py @@ -0,0 +1,371 @@ +from __future__ import annotations + +from collections.abc import Callable + +import httpx +import pytest +from pydantic import ValidationError + +from app.generation.domain.capabilities import ( + GenerationModelCapability, + GenerationProviderCapabilities, +) +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerHealthStatus, + WorkerReadinessStatus, +) +from app.generation.domain.errors import GenerationWorkerError +from app.generation.domain.retry import GenerationRetryPolicy +from app.generation.domain.runtime import WorkerInfo, WorkerOutput, WorkerReadiness +from app.generation.model_registry import ( + GenerationModelRegistration, + GenerationModelRegistry, +) +from app.generation.providers.base import GenerationProviderAdapter +from app.generation.providers.registry import GenerationProviderRegistry +from app.generation.providers.worker_client import RemoteWorkerClient + + +def worker_client( + handler: Callable[[httpx.Request], httpx.Response] | None = None, + *, + retries: int = 2, + sleep_calls: list[float] | None = None, +) -> RemoteWorkerClient: + async def sleep(delay: float) -> None: + if sleep_calls is not None: + sleep_calls.append(delay) + + client = httpx.AsyncClient( + transport=httpx.MockTransport( + handler + or (lambda _: httpx.Response(200, json={"status": "ok"})) + ) + ) + return RemoteWorkerClient( + base_url="https://worker.example", + bearer_token="test-worker-token", + connect_timeout_seconds=1, + request_timeout_seconds=1, + read_timeout_seconds=1, + retry_policy=GenerationRetryPolicy(max_retries=retries, backoff_seconds=0), + http_client=client, + sleep=sleep, + ) + + +class RuntimeTestProvider(GenerationProviderAdapter): + capabilities = GenerationProviderCapabilities( + provider="runtime-test", + name="Runtime test provider", + implementation_status="test", + models=[ + GenerationModelCapability( + id="runtime-image-v1", + name="Runtime image v1", + modality=GenerationModality.IMAGE, + ) + ], + ) + + +def test_provider_and_model_registration_starts_unavailable() -> None: + provider = RuntimeTestProvider() + providers = GenerationProviderRegistry([provider]) + assert providers.get("runtime-test") is provider + models = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id=provider.provider, + model=provider.capabilities.models[0], + configuration_reference="runtime-test-config", + metadata={ + "access_token": "must-not-survive", + "diagnostic": ( + "Bearer must-not-survive " + "https://worker.example/output?sig=secret" + ), + "download_url": "https://worker.example/output?sig=secret", + }, + ) + ] + ) + view = models.get(provider.provider, "runtime-image-v1") + assert not view.available + assert "access_token" not in view.metadata + assert "download_url" not in view.metadata + assert "must-not-survive" not in str(view.metadata) + + +def test_model_availability_requires_readiness_info_and_configuration() -> None: + model = GenerationModelCapability( + id="runtime-image-v1", name="Runtime", modality=GenerationModality.IMAGE + ) + registry = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id="runtime-test", + model=model, + configuration_reference="runtime-test-config", + ) + ] + ) + info = WorkerInfo( + id="runtime-test-worker", + name="Runtime worker", + media_types=[GenerationModality.IMAGE], + models=[ + { + "id": model.id, + "name": model.name, + "media_types": [GenerationModality.IMAGE], + } + ], + ) + not_ready = WorkerReadiness( + status=WorkerReadinessStatus.STARTING, + model_loaded=False, + model_ids=[model.id], + ) + assert not registry.verify_readiness( + provider_id="runtime-test", + worker_info=info, + readiness=not_ready, + provider_configured=True, + )[0].available + ready = WorkerReadiness( + status=WorkerReadinessStatus.READY, model_loaded=True, model_ids=[model.id] + ) + assert registry.verify_readiness( + provider_id="runtime-test", + worker_info=info, + readiness=ready, + provider_configured=True, + )[0].available + + +@pytest.mark.asyncio +async def test_worker_health_readiness_info_and_bearer_authentication() -> None: + seen_headers: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen_headers.append(request.headers.get("authorization", "")) + if request.url.path == "/health": + return httpx.Response(200, json={"status": "ok"}) + if request.url.path == "/ready": + return httpx.Response( + 200, + json={"status": "ready", "model_loaded": True, "model": "model-v1"}, + ) + return httpx.Response( + 200, + json={"id": "model-v1", "name": "Worker model", "type": "image", "status": "ready"}, + ) + + client = worker_client(handler) + assert (await client.health()).status is WorkerHealthStatus.HEALTHY + readiness = await client.ready() + assert readiness.status is WorkerReadinessStatus.READY + assert readiness.model_ids == ["model-v1"] + info = await client.info() + assert info.media_types == [GenerationModality.IMAGE] + assert info.models[0].id == "model-v1" + assert seen_headers == ["Bearer test-worker-token"] * 3 + + +@pytest.mark.asyncio +async def test_timeout_and_connection_failure_are_retryable_and_safe() -> None: + request = httpx.Request("GET", "https://worker.example/health") + for exception, category in ( + (httpx.ReadTimeout("secret-token", request=request), WorkerErrorCategory.TIMEOUT), + ( + httpx.ConnectError("Bearer test-worker-token", request=request), + WorkerErrorCategory.WORKER_UNAVAILABLE, + ), + ): + calls = 0 + + def handler(_: httpx.Request, error: Exception = exception) -> httpx.Response: + nonlocal calls + calls += 1 + raise error + + client = worker_client(handler, retries=1) + with pytest.raises(GenerationWorkerError) as raised: + await client.health() + assert raised.value.category is category + assert "test-worker-token" not in str(raised.value) + assert calls == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [429, 502, 503, 504]) +async def test_retryable_http_failures_use_bounded_retry(status_code: int) -> None: + calls = 0 + delays: list[float] = [] + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + if calls < 3: + return httpx.Response(status_code, json={"secret": "not surfaced"}) + return httpx.Response(200, json={"status": "ok"}) + + client = worker_client(handler, retries=2, sleep_calls=delays) + assert (await client.health()).status is WorkerHealthStatus.HEALTHY + assert calls == 3 + assert delays == [0, 0] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [400, 401]) +async def test_non_retryable_http_failures_do_not_retry(status_code: int) -> None: + calls = 0 + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(status_code) + + client = worker_client(handler, retries=3) + with pytest.raises(GenerationWorkerError) as raised: + await client.health() + assert calls == 1 + assert raised.value.http_status == status_code + + +@pytest.mark.asyncio +async def test_unexpected_exception_is_not_automatically_retryable() -> None: + calls = 0 + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + raise RuntimeError("programming failure with secret-token") + + client = worker_client(handler, retries=3) + with pytest.raises(GenerationWorkerError) as raised: + await client.health() + assert raised.value.category is WorkerErrorCategory.UNKNOWN_ERROR + assert calls == 1 + assert "secret-token" not in str(raised.value) + + +@pytest.mark.asyncio +async def test_worker_cancellation_and_output_contract() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "POST": + return httpx.Response(202, json={"status": "cancellation_requested"}) + return httpx.Response( + 200, + json={ + "job_id": "job-1", + "status": "completed", + "output": { + "type": "image", + "mime_type": "image/png", + "id": "output-1", + "download_path": "/v1/outputs/output-1", + "filename": "output.png", + }, + }, + ) + + client = worker_client(handler) + cancellation = await client.cancel("job-1") + assert cancellation.status is WorkerCancellationStatus.REQUESTED + output = await client.retrieve_output("job-1") + assert output.provider_output_id == "output-1" + assert output.download_path == "/v1/outputs/output-1" + with pytest.raises(ValidationError): + WorkerOutput( + output_type=GenerationModality.IMAGE, + mime_type="image/png", + provider_output_id="output-1", + download_path="https://attacker.example/output.png", + ) + with pytest.raises(ValidationError): + WorkerOutput( + output_type=GenerationModality.IMAGE, + mime_type="image/png", + provider_output_id="output-1", + download_path="/v1/outputs/%2e%2e/secrets", + ) + + +@pytest.mark.asyncio +async def test_empty_successful_cancellation_response_means_requested_not_cancelled() -> None: + client = worker_client(lambda _: httpx.Response(204)) + result = await client.cancel("job-1") + assert result.status is WorkerCancellationStatus.REQUESTED + + +@pytest.mark.asyncio +async def test_output_stream_is_scoped_to_the_configured_worker_origin() -> None: + client = worker_client(lambda _: httpx.Response(200, content=b"worker-output")) + output = WorkerOutput( + output_type=GenerationModality.IMAGE, + mime_type="image/png", + provider_output_id="output-1", + download_path="/v1/outputs/output-1", + ) + async with client.stream_output(output) as chunks: + received = b"".join([chunk async for chunk in chunks]) + assert received == b"worker-output" + + +@pytest.mark.asyncio +async def test_worker_info_requires_a_discovered_model_match_for_availability() -> None: + model = GenerationModelCapability( + id="runtime-image-v1", name="Runtime", modality=GenerationModality.IMAGE + ) + registry = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id="runtime-test", + model=model, + configuration_reference="runtime-test-config", + ) + ] + ) + readiness = WorkerReadiness( + status=WorkerReadinessStatus.READY, model_loaded=True, model_ids=[model.id] + ) + undiscovered = WorkerInfo( + id="worker", + name="Worker", + media_types=[GenerationModality.IMAGE], + models=[{"id": "other-model", "name": "Other", "media_types": ["image"]}], + ) + assert not registry.verify_readiness( + provider_id="runtime-test", + worker_info=undiscovered, + readiness=readiness, + provider_configured=True, + )[0].available + + +def test_worker_url_and_path_validation_blocks_ssrf_and_traversal() -> None: + policy = GenerationRetryPolicy(max_retries=0, backoff_seconds=0) + for url in ( + "http://example.com", + "https://10.0.0.1", + "http://169.254.169.254", + "https://169.254.169.254", + "https://worker.example/%2e%2e/internal", + "file:///etc/passwd", + ): + with pytest.raises(ValueError): + RemoteWorkerClient( + base_url=url, + bearer_token=None, + connect_timeout_seconds=1, + request_timeout_seconds=1, + read_timeout_seconds=1, + retry_policy=policy, + ) + with pytest.raises(GenerationWorkerError): + RemoteWorkerClient._safe_external_id("job/../../metadata") diff --git a/tests/test_generation_wan.py b/tests/test_generation_wan.py new file mode 100644 index 0000000000000000000000000000000000000000..efbd555cc5f761f053f6bb1858b2b9c6efc977a3 --- /dev/null +++ b/tests/test_generation_wan.py @@ -0,0 +1,346 @@ +"""Mocked protocol tests for the audited WAN 2.2 worker integration.""" + +from __future__ import annotations + +from collections.abc import Callable +from pathlib import Path + +import httpx +import pytest + +from app.core.config import Settings +from app.generation.domain.enums import ( + GenerationModality, + WorkerCancellationStatus, + WorkerErrorCategory, + WorkerJobStatus, +) +from app.generation.domain.errors import GenerationWorkerError +from app.generation.domain.retry import GenerationRetryPolicy +from app.generation.model_registry import GenerationModelRegistration, GenerationModelRegistry +from app.generation.providers.wan import ( + WAN_MODEL_CAPABILITY, + WAN_MODEL_ID, + WAN_PROVIDER_ID, + WanProviderAdapter, +) +from app.generation.providers.worker_client import RemoteWorkerClient +from app.generation.schemas.requests import GenerationRequestCreate + + +def _client( + handler: Callable[[httpx.Request], httpx.Response], *, retries: int = 2 +) -> RemoteWorkerClient: + return RemoteWorkerClient( + base_url="https://wan-worker.example", + bearer_token="x" * 32, + connect_timeout_seconds=1, + request_timeout_seconds=1, + read_timeout_seconds=1, + retry_policy=GenerationRetryPolicy(max_retries=retries, backoff_seconds=0), + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + sleep=lambda _: _no_sleep(), + ) + + +async def _no_sleep() -> None: + return None + + +def _payload(**overrides: object) -> GenerationRequestCreate: + value: dict[str, object] = { + "provider": WAN_PROVIDER_ID, + "model_id": WAN_MODEL_ID, + "modality": "video", + "input_asset_id": "11111111-1111-4111-8111-111111111111", + "prompt": "Slow cinematic cloud movement", + "wan": { + "duration_seconds": 0.5, + "steps": 4, + "guidance_scale": 1.0, + "guidance_scale_2": 1.0, + "seed": 42, + "randomize_seed": False, + }, + } + value.update(overrides) + return GenerationRequestCreate.model_validate(value) + + +def _info() -> dict[str, object]: + return { + "id": "wan2.2", + "name": "WAN 2.2 FP8 AOTI Faster", + "type": "video", + "task": "image-to-video", + "status": "ready", + "model_id": "Wan-AI/Wan2.2-I2V-A14B-Diffusers", + "fps": 16, + } + + +@pytest.mark.asyncio +async def test_wan_exact_model_discovery_and_readiness() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/health": + return httpx.Response(200, json={"status": "ok", "service": "mediarouter-wan-worker"}) + if request.url.path == "/ready": + return httpx.Response( + 200, + json={ + "status": "ready", + "model_loaded": True, + "model": "wan2.2", + "accepting_jobs": True, + }, + ) + return httpx.Response(200, json=_info()) + + adapter = WanProviderAdapter(client=_client(handler)) + registry = GenerationModelRegistry( + [ + GenerationModelRegistration( + provider_id=WAN_PROVIDER_ID, + model=WAN_MODEL_CAPABILITY, + configuration_reference="wan-space", + ) + ] + ) + models = registry.verify_readiness( + provider_id=WAN_PROVIDER_ID, + worker_info=await adapter.info(), + readiness=await adapter.ready(), + provider_configured=adapter.available, + ) + assert models[0].model.id == WAN_MODEL_ID + assert models[0].model.modality is GenerationModality.VIDEO + assert models[0].available + + +@pytest.mark.asyncio +async def test_wan_not_ready_and_model_mismatch_are_not_advertised() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/ready": + return httpx.Response( + 503, + json={ + "status": "not_ready", + "model_loaded": False, + "model": "other-model", + "accepting_jobs": False, + }, + ) + return httpx.Response( + 200, json={"status": "ok"} if request.url.path == "/health" else _info() + ) + + adapter = WanProviderAdapter(client=_client(handler)) + with pytest.raises(GenerationWorkerError) as raised: + await adapter.ready() + assert raised.value.category is WorkerErrorCategory.WORKER_NOT_READY + + +@pytest.mark.asyncio +async def test_wan_model_identity_mismatch_remains_unavailable() -> None: + wrong_info = { + **_info(), + "id": "different-wan-model", + "name": "Different model", + } + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/ready": + return httpx.Response( + 200, + json={ + "status": "ready", + "model_loaded": True, + "model": WAN_MODEL_ID, + "accepting_jobs": True, + }, + ) + return httpx.Response(200, json=wrong_info) + + adapter = WanProviderAdapter(client=_client(handler)) + with pytest.raises(GenerationWorkerError) as raised: + await adapter.info() + assert raised.value.category is WorkerErrorCategory.PROVIDER_ERROR + + +@pytest.mark.asyncio +async def test_wan_submission_is_multipart_and_has_no_automatic_retry(tmp_path: Path) -> None: + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(202, json={"job_id": "wan_" + "a" * 32, "status": "queued"}) + + source = tmp_path / "input.png" + source.write_bytes(b"not-decoded-in-adapter-test") + adapter = WanProviderAdapter(client=_client(handler)) + job = await adapter.submit( + payload={"prompt": "slow movement", "wan": {"duration_seconds": 0.5, "steps": 4}}, + idempotency_key="generation-request-id", + input_path=source, + input_mime_type="image/png", + ) + assert job.status is WorkerJobStatus.QUEUED + assert job.external_job_id.startswith("wan_") + assert seen[0].headers["authorization"] == "Bearer " + "x" * 32 + body = seen[0].content.decode("latin-1") + assert 'name="image"' in body + assert 'name="duration_seconds"' in body + assert 'name="width"' not in body + + +@pytest.mark.asyncio +async def test_wan_submission_connection_ambiguity_is_not_retried(tmp_path: Path) -> None: + calls = 0 + request = httpx.Request("POST", "https://wan-worker.example/v1/generate") + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + raise httpx.ConnectError("Bearer " + "x" * 32, request=request) + + source = tmp_path / "input.png" + source.write_bytes(b"input") + adapter = WanProviderAdapter(client=_client(handler, retries=3)) + with pytest.raises(GenerationWorkerError) as raised: + await adapter.submit( + payload={"prompt": "slow movement"}, + idempotency_key="generation-request-id", + input_path=source, + input_mime_type="image/png", + ) + assert raised.value.category is WorkerErrorCategory.WORKER_UNAVAILABLE + assert calls == 1 + assert "Bearer" not in str(raised.value) + + +@pytest.mark.asyncio +async def test_wan_completed_job_maps_a_safe_video_output() -> None: + job_id = "wan_" + "b" * 32 + + def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={ + "job_id": job_id, + "status": "completed", + "output": {"type": "video", "filename": f"{job_id}.mp4"}, + }, + ) + + adapter = WanProviderAdapter(client=_client(handler)) + job = await adapter.get_job(external_job_id=job_id) + assert job.output is not None + assert job.output.mime_type == "video/mp4" + assert job.output.provider_output_id == job_id + assert job.output.download_path == f"/v1/jobs/{job_id}/output" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [429, 502, 503, 504]) +async def test_wan_polling_uses_shared_bounded_transient_retry(status_code: int) -> None: + calls = 0 + job_id = "wan_" + "d" * 32 + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + if calls < 3: + return httpx.Response(status_code, json={"detail": {"token": "never-store"}}) + return httpx.Response(200, json={"job_id": job_id, "status": "running"}) + + job = await WanProviderAdapter(client=_client(handler, retries=2)).get_job( + external_job_id=job_id + ) + assert job.status is WorkerJobStatus.RUNNING + assert calls == 3 + + +@pytest.mark.asyncio +async def test_wan_polling_does_not_retry_permanent_client_errors() -> None: + calls = 0 + job_id = "wan_" + "e" * 32 + + def handler(_: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(400, json={"detail": {"code": "WAN_PARAMETERS_INVALID"}}) + + with pytest.raises(GenerationWorkerError) as raised: + await WanProviderAdapter(client=_client(handler, retries=3)).get_job( + external_job_id=job_id + ) + assert raised.value.category is WorkerErrorCategory.INVALID_REQUEST + assert calls == 1 + + +@pytest.mark.asyncio +async def test_wan_cancellation_only_confirms_queued_worker_cancellation() -> None: + def queued_handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"job_id": "wan_" + "c" * 32, "status": "cancelled"}) + + def running_handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 409, + json={ + "detail": { + "code": "WAN_JOB_NOT_CANCELLABLE", + "message": "A running job cannot be cancelled.", + "status": "running", + } + }, + ) + + assert ( + await WanProviderAdapter(client=_client(queued_handler)).cancel( + external_job_id="wan_" + "c" * 32 + ) + ).status is WorkerCancellationStatus.CANCELLED + assert ( + await WanProviderAdapter(client=_client(running_handler)).cancel( + external_job_id="wan_" + "c" * 32 + ) + ).status is WorkerCancellationStatus.FAILED + + +@pytest.mark.parametrize( + "invalid", + [ + {"prompt": " "}, + {"modality": "image"}, + {"input_asset_id": None}, + ], +) +@pytest.mark.asyncio +async def test_wan_request_validation_rejects_invalid_required_values( + invalid: dict[str, object] +) -> None: + adapter = WanProviderAdapter(client=None) + with pytest.raises(Exception): + payload = _payload(**invalid) + await adapter.validate_request(payload) + + +@pytest.mark.parametrize("field", ["width", "height", "num_frames", "provider_payload"]) +def test_wan_schema_rejects_unsupported_parameters(field: str) -> None: + raw = _payload().model_dump() + wan = dict(raw["wan"] or {}) + wan[field] = 1 + raw["wan"] = wan + with pytest.raises(ValueError): + GenerationRequestCreate.model_validate(raw) + + +def test_wan_configuration_is_optional_and_never_enables_flux() -> None: + disabled = WanProviderAdapter.from_settings(Settings(_env_file=None)) + invalid = WanProviderAdapter.from_settings( + Settings(_env_file=None, wan_space_url="https://wan-worker.example") + ) + assert not disabled.available + assert not invalid.available + assert invalid.configuration_error is not None + assert WAN_PROVIDER_ID == "wan" diff --git a/tests/test_health.py b/tests/test_health.py new file mode 100644 index 0000000000000000000000000000000000000000..a6df0c858fbd711e13b81321fd3b0c31623551b8 --- /dev/null +++ b/tests/test_health.py @@ -0,0 +1,13 @@ +from fastapi.testclient import TestClient + +from main import create_app + + +def test_health_endpoint(settings) -> None: + with TestClient(create_app(settings)) as client: + response = client.get("/health") + assert response.status_code == 200 + payload = response.json() + assert payload["success"] is True + assert payload["metadata"]["status"] == "healthy" + assert response.headers["x-request-id"] == payload["request_id"] diff --git a/tests/test_input_resolver.py b/tests/test_input_resolver.py new file mode 100644 index 0000000000000000000000000000000000000000..23713b72f9d186bf6b196e20a493a029e2cc75e7 --- /dev/null +++ b/tests/test_input_resolver.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +import base64 +from uuid import uuid4 + +import pytest +from starlette.requests import Request + +from app.container import build_container +from app.core.exceptions import InputError +from app.models.media import MediaSource + + +def json_request(payload: bytes) -> Request: + sent = False + + async def receive(): + nonlocal sent + if sent: + return {"type": "http.disconnect"} + sent = True + return {"type": "http.request", "body": payload, "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/probe", + "headers": [(b"content-type", b"application/json")], + "query_string": b"", + }, + receive, + ) + request.state.request_id = str(uuid4()) + return request + + +async def test_resolves_json_base64(settings) -> None: + container = build_container(settings) + encoded = base64.b64encode(b"ID3-not-real-audio").decode() + request = json_request( + ('{"base64":"%s","filename":"sample.mp3","format":"wav"}' % encoded).encode() + ) + resolved = await container.resolver.resolve(request) + assert resolved.primary.source is MediaSource.JSON_BASE64 + assert resolved.primary.filename == "sample.mp3" + assert resolved.primary.temp_path.read_bytes() == b"ID3-not-real-audio" + assert resolved.params["filename"] == "sample.mp3" + assert resolved.params["format"] == "wav" + + +async def test_resolves_n8n_binary_property(settings) -> None: + container = build_container(settings) + encoded = base64.b64encode(b"audio").decode() + payload = ( + '{"binary":{"audio":{"data":"%s","fileName":"voice.mp3",' + '"mimeType":"audio/mpeg"}}}' % encoded + ).encode() + resolved = await container.resolver.resolve(json_request(payload)) + assert resolved.primary.source is MediaSource.N8N_BINARY + assert resolved.primary.filename == "voice.mp3" + assert resolved.primary.temp_path.read_bytes() == b"audio" + + +async def test_resolves_nested_template_input(settings) -> None: + container = build_container(settings) + encoded = base64.b64encode(b"RIFF-template-audio").decode() + request = json_request( + ( + '{"template":"mp3","input":{"base64":"%s",' + '"filename":"source.wav","mime_type":"audio/wav"},"parameters":{}}' % encoded + ).encode() + ) + + resolved = await container.resolver.resolve(request) + + assert resolved.primary.source is MediaSource.JSON_BASE64 + assert resolved.primary.filename == "source.wav" + assert resolved.params["template"] == "mp3" + assert resolved.params["parameters"] == {} + + +async def test_resolve_payload_copies_managed_temp_file(settings) -> None: + settings.output_dir.mkdir(parents=True) + source = settings.output_dir / "previous" / "clip.mp3" + source.parent.mkdir() + source.write_bytes(b"ID3-managed-media") + container = build_container(settings) + + resolved = await container.resolver.resolve_payload({"temp_path": str(source)}, str(uuid4())) + + assert resolved.primary.source is MediaSource.LOCAL_PATH + assert resolved.primary.temp_path != source + assert resolved.primary.temp_path.read_bytes() == source.read_bytes() + + +async def test_resolve_payload_rejects_unmanaged_path(settings, tmp_path) -> None: + source = tmp_path / "outside.mp3" + source.write_bytes(b"ID3-unmanaged-media") + container = build_container(settings) + + with pytest.raises(InputError, match="TEMP_DIR or OUTPUT_DIR"): + await container.resolver.resolve_payload({"temp_path": str(source)}, str(uuid4())) diff --git a/tests/test_linkedin_foundation.py b/tests/test_linkedin_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..4cfa0fb69447fc621a85395b42d30cb2e8c2bcf1 --- /dev/null +++ b/tests/test_linkedin_foundation.py @@ -0,0 +1,565 @@ +"""Phase 6A LinkedIn OIDC and organization-discovery coverage. + +All LinkedIn traffic is mocked. Normal CI needs no developer application, +member credential, organization role, or interactive authorization flow. +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from pydantic import ValidationError +from sqlalchemy import func, select + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialCapabilityUnsupportedError, + SocialOAuthStateError, + SocialPermissionDeniedError, + SocialReauthRequiredError, +) +from app.social.models import OAuthState, SocialAccountToken +from app.social.providers.linkedin import LINKEDIN_API_VERSION, LinkedInProvider +from app.social.schemas.accounts import SocialAccountConnectRequest + +_REDIRECT_URI = "https://api.example.com/v1/social/accounts/linkedin/callback" + + +def linkedin_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="phase-6a-linkedin-test-encryption-material", + social_oauth_redirect_base_url="https://api.example.com", + linkedin_client_id="linkedin-client-id", + linkedin_client_secret="linkedin-client-secret", + linkedin_redirect_uri=_REDIRECT_URI, + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +async def test_linkedin_member_authorization_uses_official_oidc_without_pkce( + tmp_path: Path, +) -> None: + provider = LinkedInProvider(linkedin_settings(tmp_path)) + try: + url = await provider.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + ) + with pytest.raises(SocialPermissionDeniedError): + await provider.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + code_challenge="undocumented-pkce-challenge", + ) + with pytest.raises(SocialCapabilityUnsupportedError): + await provider.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + additional_scopes=["w_member_social"], + ) + finally: + await provider.close() + + parsed = urlparse(url) + query = parse_qs(parsed.query) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == ( + "https://www.linkedin.com/oauth/v2/authorization" + ) + assert query == { + "client_id": ["linkedin-client-id"], + "redirect_uri": [_REDIRECT_URI], + "response_type": ["code"], + "state": ["s" * 43], + "scope": ["openid profile"], + } + assert "code_challenge" not in query + + +async def test_linkedin_exchange_member_and_organization_discovery_use_official_apis( + tmp_path: Path, +) -> None: + calls: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + if request.url.path == "/oauth/v2/accessToken": + assert request.url.host == "www.linkedin.com" + form = parse_qs(request.content.decode()) + assert form == { + "grant_type": ["authorization_code"], + "code": ["authorization-code"], + "redirect_uri": [_REDIRECT_URI], + "client_id": ["linkedin-client-id"], + "client_secret": ["linkedin-client-secret"], + } + return httpx.Response( + 200, + json={ + "access_token": "linkedin-access-token", + "expires_in": 5184000, + "token_type": "Bearer", + }, + ) + assert request.headers["authorization"] == "Bearer linkedin-access-token" + if request.url.path == "/v2/userinfo": + return httpx.Response( + 200, + json={ + "sub": "oidc-member-subject_123", + "name": "Ada Lovelace", + "given_name": "Ada", + "family_name": "Lovelace", + "picture": "https://media.licdn.com/member.jpg", + "locale": {"country": "US", "language": "en"}, + "email": "not-persisted@example.com", + "email_verified": True, + }, + ) + assert request.headers["linkedin-version"] == LINKEDIN_API_VERSION + assert request.headers["x-restli-protocol-version"] == "2.0.0" + if request.url.path == "/rest/organizationAcls": + assert parse_qs(request.url.query.decode()) == { + "q": ["roleAssignee"], + "role": ["ADMINISTRATOR"], + "state": ["APPROVED"], + "count": ["100"], + "start": ["0"], + } + return httpx.Response( + 200, + json={ + "elements": [ + {"organization": "urn:li:organization:123456"}, + { + "organizationTarget": "urn:li:organization:789012" + }, + ], + "paging": {"start": 0, "count": 2, "total": 2, "links": []}, + }, + ) + organization_id = request.url.path.rsplit("/", 1)[-1] + return httpx.Response( + 200, + json={ + "id": int(organization_id), + "localizedName": f"Organization {organization_id}", + "vanityName": f"organization-{organization_id}", + "logoV2": { + "digitalmediaAsset": "urn:li:digitalmediaAsset:logo_asset", + "original~": { + "elements": [ + { + "identifiers": [ + { + "identifier": f"https://media.licdn.com/{organization_id}.png" + } + ] + } + ] + }, + }, + "primaryOrganizationType": "NONE", + "defaultLocale": {"country": "US", "language": "en"}, + "localizedWebsite": "https://example.com", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(linkedin_settings(tmp_path), http_client=client) + try: + token = await provider.exchange_code( + code="authorization-code", + redirect_uri=_REDIRECT_URI, + ) + accounts = await provider.discover_accounts( + token, account_type="linkedin_organization" + ) + finally: + await client.aclose() + + member, first, second = accounts + assert member["external_account_id"] == "oidc-member-subject_123" + assert member["account_type"] == "linkedin_member" + assert member["connection_status"] == "connected" + assert member["metadata"]["email_verified"] is True + assert "email" not in member["metadata"] + assert [first["external_account_id"], second["external_account_id"]] == [ + "123456", + "789012", + ] + assert first["account_type"] == "linkedin_organization" + assert first["connection_status"] == "pending" + assert first["avatar_url"] == "https://media.licdn.com/123456.png" + assert first["metadata"]["parent_member_id"] == "oidc-member-subject_123" + assert first["metadata"]["logo_asset"] == ( + "urn:li:digitalmediaAsset:logo_asset" + ) + assert calls == [ + "/oauth/v2/accessToken", + "/v2/userinfo", + "/rest/organizationAcls", + "/rest/organizations/123456", + "/rest/organizations/789012", + ] + + +async def test_linkedin_organization_scope_is_explicit_and_bound_to_state( + tmp_path: Path, +) -> None: + container = build_container(linkedin_settings(tmp_path)) + await container.social.initialize() + try: + member = await container.social.oauth.connect( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(), + ) + member_query = parse_qs(urlparse(member.authorization_url or "").query) + assert member_query["scope"] == ["openid profile"] + assert "code_challenge" not in member_query + + organization = await container.social.oauth.connect( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest( + account_type="linkedin_organization" + ), + ) + organization_query = parse_qs( + urlparse(organization.authorization_url or "").query + ) + assert organization_query["scope"] == [ + "openid profile rw_organization_admin" + ] + assert "w_organization_social" not in organization_query["scope"][0] + state = await container.social.oauth.states.consume( + state=organization_query["state"][0], provider="linkedin" + ) + assert state.workspace_id == "workspace-a" + assert state.user_id == "user-a" + assert state.requested_account_type == "linkedin_organization" + assert state.requested_scopes == [ + "openid", + "profile", + "rw_organization_admin", + ] + + with pytest.raises(SocialCapabilityUnsupportedError): + await container.social.oauth.connect( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(account_type="organization"), + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_linkedin_callback_is_duplicate_safe_selectable_and_workspace_bound( + tmp_path: Path, +) -> None: + container = build_container(linkedin_settings(tmp_path)) + await container.social.initialize() + adapter = container.social.accounts.providers.get("linkedin") + assert isinstance(adapter, LinkedInProvider) + await adapter._client.aclose() + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/oauth/v2/accessToken": + return httpx.Response( + 200, + json={ + "access_token": "linkedin-token-that-must-remain-encrypted", + "expires_in": 3600, + "token_type": "Bearer", + }, + ) + if request.url.path == "/v2/userinfo": + return httpx.Response( + 200, + json={"sub": "stable-member-sub", "name": "Workspace Member"}, + ) + if request.url.path == "/rest/organizationAcls": + return httpx.Response( + 200, + json={ + "elements": [ + {"organization": "urn:li:organization:123456"} + ], + "paging": {"start": 0, "count": 1, "total": 1}, + }, + ) + return httpx.Response( + 200, + json={ + "id": 123456, + "localizedName": "Workspace Organization", + "vanityName": "workspace-organization", + }, + ) + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + try: + connect = await container.social.oauth.connect( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest( + account_type="linkedin_organization" + ), + ) + state = parse_qs(urlparse(connect.authorization_url or "").query)["state"][0] + member = await container.social.oauth.callback( + provider="linkedin", state=state, code="first-code" + ) + assert member.account_type == "linkedin_member" + assert member.status.value == "connected" + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.callback( + provider="linkedin", state=state, code="replayed-code" + ) + + accounts = await container.social.accounts.list("workspace-a") + assert len(accounts) == 2 + organization = next( + item + for item in accounts + if item.account_type == "linkedin_organization" + ) + assert organization.status.value == "pending" + assert "linkedin-token-that-must-remain-encrypted" not in ( + organization.model_dump_json() + ) + + with pytest.raises(SocialAccountNotFoundError): + await container.social.accounts.select_discovered( + "workspace-b", [organization.id] + ) + selected = await container.social.accounts.select_discovered( + "workspace-a", [organization.id, organization.id] + ) + assert len(selected) == 1 + assert selected[0].status.value == "connected" + + second_connect = await container.social.oauth.connect( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest( + account_type="linkedin_organization" + ), + ) + second_state = parse_qs( + urlparse(second_connect.authorization_url or "").query + )["state"][0] + await container.social.oauth.callback( + provider="linkedin", state=second_state, code="second-code" + ) + duplicate_safe = await container.social.accounts.list("workspace-a") + assert len(duplicate_safe) == 2 + assert next( + item + for item in duplicate_safe + if item.account_type == "linkedin_organization" + ).status.value == "connected" + + async with container.social.database.session("workspace-a") as session: + token_count = await session.scalar(select(func.count(SocialAccountToken.id))) + encrypted_payloads = list( + ( + await session.scalars( + select(SocialAccountToken.encrypted_payload) + ) + ).all() + ) + assert token_count == 2 + assert all(encrypted_payloads) + assert all( + "linkedin-token-that-must-remain-encrypted" not in str(payload) + for payload in encrypted_payloads + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_linkedin_state_redirect_expiry_and_provider_binding( + tmp_path: Path, +) -> None: + container = build_container(linkedin_settings(tmp_path)) + await container.social.initialize() + try: + assert container.social.oauth._redirect_uri("linkedin", None) == _REDIRECT_URI + with pytest.raises(SocialPermissionDeniedError): + container.social.oauth._redirect_uri( + "linkedin", + "https://attacker.example/v1/social/accounts/linkedin/callback", + ) + state = await container.social.oauth.states.create( + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=_REDIRECT_URI, + requested_account_type="linkedin_member", + requested_scopes=["openid", "profile"], + ) + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=state.state, provider="x" + ) + consumed = await container.social.oauth.states.consume( + state=state.state, provider="linkedin" + ) + assert consumed.workspace_id == "workspace-a" + + expired = OAuthState( + state="expired-linkedin-state-value-that-is-long-enough", + provider="linkedin", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=_REDIRECT_URI, + requested_account_type="linkedin_member", + requested_scopes=["openid", "profile"], + expires_at=datetime.now(timezone.utc) - timedelta(seconds=1), + ) + async with container.social.database.session("workspace-a") as session: + session.add(expired) + await session.commit() + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=expired.state, provider="linkedin" + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_linkedin_invalid_code_and_refresh_rules_do_not_leak_secrets( + tmp_path: Path, +) -> None: + secret_code = "linkedin-code-that-must-not-leak" + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + json={ + "error": "invalid_grant", + "error_description": f"bad code {secret_code}", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(linkedin_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialPermissionDeniedError) as raised: + await provider.exchange_code( + code=secret_code, redirect_uri=_REDIRECT_URI + ) + with pytest.raises(SocialReauthRequiredError): + await provider.refresh_token({"access_token": "expired"}) + finally: + await client.aclose() + assert secret_code not in str(raised.value) + assert "linkedin-client-secret" not in str(raised.value) + + +async def test_linkedin_refresh_is_used_only_when_provider_issued_it( + tmp_path: Path, +) -> None: + async def handler(request: httpx.Request) -> httpx.Response: + form = parse_qs(request.content.decode()) + assert form == { + "grant_type": ["refresh_token"], + "refresh_token": ["partner-refresh-token"], + "client_id": ["linkedin-client-id"], + "client_secret": ["linkedin-client-secret"], + } + return httpx.Response( + 200, + json={ + "access_token": "refreshed-access-token", + "expires_in": 5184000, + "token_type": "Bearer", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(linkedin_settings(tmp_path), http_client=client) + try: + refreshed = await provider.refresh_token( + { + "access_token": "expired-access-token", + "refresh_token": "partner-refresh-token", + } + ) + finally: + await client.aclose() + + assert refreshed["access_token"] == "refreshed-access-token" + assert refreshed["refresh_token"] == "partner-refresh-token" + + +async def test_linkedin_capabilities_are_discovery_only(tmp_path: Path) -> None: + container = build_container(linkedin_settings(tmp_path)) + try: + linkedin = container.social.accounts.get_provider("linkedin") + assert linkedin.available + assert linkedin.configured + assert linkedin.capabilities.implementation_status == "implemented" + assert linkedin.capabilities.account_types == [ + "linkedin_member", + "linkedin_organization", + ] + assert linkedin.capabilities.required_scopes == ["openid", "profile"] + assert linkedin.capabilities.optional_scopes == ["rw_organization_admin"] + assert not linkedin.capabilities.video + assert not linkedin.capabilities.image + assert not linkedin.capabilities.direct_publish + assert not linkedin.capabilities.draft_upload + assert not linkedin.capabilities.scheduled_publish + assert not linkedin.capabilities.analytics + assert not linkedin.capabilities.delete_post + assert not linkedin.capabilities.personal_publishing + assert not linkedin.capabilities.organization_publishing + assert linkedin.capabilities.publish_metadata_schema == {} + finally: + await container.social.close() + await container.security_database.close() + + +def test_linkedin_redirect_configuration_is_fail_closed() -> None: + invalid = [ + "https://attacker.example/not-the-linkedin-callback", + "http://api.example.com/v1/social/accounts/linkedin/callback", + "https://api.example.com/v1/social/accounts/linkedin/callback?next=bad", + "ftp://localhost/v1/social/accounts/linkedin/callback", + ] + for redirect in invalid: + with pytest.raises(ValidationError): + Settings(_env_file=None, linkedin_redirect_uri=redirect) + local = Settings( + _env_file=None, + linkedin_redirect_uri=( + "http://localhost/v1/social/accounts/linkedin/callback" + ), + ) + assert local.linkedin_redirect_uri.startswith("http://localhost/") diff --git a/tests/test_linkedin_live.py b/tests/test_linkedin_live.py new file mode 100644 index 0000000000000000000000000000000000000000..94d325e94f2a3d63c9d93714d4cb142b28d0bbcb --- /dev/null +++ b/tests/test_linkedin_live.py @@ -0,0 +1,327 @@ +"""Opt-in, destructive LinkedIn integration verification. + +Normal CI always skips this module. Run it only with a dedicated LinkedIn +member and, when organization coverage is required, a dedicated organization. +Provider credentials are read from the process environment and never logged. +""" + +from __future__ import annotations + +import asyncio +import os +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import pytest + + +pytestmark = pytest.mark.skipif( + os.getenv("RUN_LINKEDIN_INTEGRATION_TESTS", "").lower() != "true", + reason=( + "LinkedIn live integration is NOT VERIFIED; set " + "RUN_LINKEDIN_INTEGRATION_TESTS=true with dedicated credentials." + ), +) + + +def _required(name: str) -> str: + value = os.getenv(name, "").strip() + if not value: + pytest.skip(f"LinkedIn live integration is NOT VERIFIED; missing {name}.") + return value + + +def _granted_scopes() -> list[str]: + value = os.getenv("LINKEDIN_LIVE_TEST_GRANTED_SCOPES", "") + return list(dict.fromkeys(value.replace(",", " ").split())) + + +def _require_base_configuration() -> None: + for name in ( + "LINKEDIN_CLIENT_ID", + "LINKEDIN_CLIENT_SECRET", + "LINKEDIN_REDIRECT_URI", + "LINKEDIN_LIVE_TEST_ACCESS_TOKEN", + ): + _required(name) + + +def _require_publish_consent() -> None: + _require_base_configuration() + if os.getenv("LINKEDIN_LIVE_TEST_ALLOW_PUBLISH", "").lower() != "true": + pytest.skip( + "Set LINKEDIN_LIVE_TEST_ALLOW_PUBLISH=true to create test posts." + ) + if os.getenv("LINKEDIN_LIVE_TEST_DELETE", "").lower() != "true": + pytest.skip( + "Set LINKEDIN_LIVE_TEST_DELETE=true to require deletion of test posts." + ) + + +def _settings(): + from app.core.config import Settings + + return Settings( + _env_file=None, + auth_enabled=False, + linkedin_client_id=_required("LINKEDIN_CLIENT_ID"), + linkedin_client_secret=_required("LINKEDIN_CLIENT_SECRET"), + linkedin_redirect_uri=_required("LINKEDIN_REDIRECT_URI"), + linkedin_publishing_enabled=True, + whisper_model="tiny", + ) + + +def _token() -> dict[str, object]: + return { + "access_token": _required("LINKEDIN_LIVE_TEST_ACCESS_TOKEN"), + "_mediarouter_granted_scopes": _granted_scopes(), + } + + +async def _identity(provider, token: dict[str, object]) -> tuple[str, str]: + account_type = os.getenv( + "LINKEDIN_LIVE_TEST_ACCOUNT_TYPE", "linkedin_member" + ).strip() + if account_type == "linkedin_member": + member = await provider.get_account(token) + return account_type, str(member["external_account_id"]) + if account_type != "linkedin_organization": + pytest.skip( + "LINKEDIN_LIVE_TEST_ACCOUNT_TYPE must be linkedin_member or " + "linkedin_organization." + ) + organization_id = _required("LINKEDIN_LIVE_TEST_ORGANIZATION_ID") + discovered = await provider.discover_accounts( + token, account_type="linkedin_organization" + ) + organizations = { + str(account["external_account_id"]): account + for account in discovered + if account.get("account_type") == "linkedin_organization" + } + if organization_id not in organizations: + pytest.fail( + "The configured LinkedIn organization was not returned by official " + "organization-access discovery." + ) + return account_type, organization_id + + +async def _publish_text(provider, token: dict[str, object]) -> dict[str, object]: + account_type, account_id = await _identity(provider, token) + state: dict[str, object] = {} + + async def persist(value: dict[str, object]) -> None: + state.clear() + state.update(value) + + return await provider.publish( + token, + { + "provider_account_id": account_id, + "provider_account_type": account_type, + "linkedin_post_metadata": { + "post_type": "text", + "commentary": ( + "MediaRouter Phase 6C live integration verification" + ), + }, + "upload": {"identity_type": "none"}, + "provider_state": state, + "persist_provider_state": persist, + }, + ) + + +def test_linkedin_live_configuration_requires_explicit_opt_in() -> None: + assert os.getenv("RUN_LINKEDIN_INTEGRATION_TESTS", "").lower() == "true" + _require_base_configuration() + + +async def test_linkedin_live_authorization_url_and_account_discovery() -> None: + """Verify the official authorization contract and current identity token.""" + + _require_base_configuration() + from app.social.providers.linkedin import LinkedInProvider + + provider = LinkedInProvider(_settings()) + token = _token() + account_type = os.getenv( + "LINKEDIN_LIVE_TEST_ACCOUNT_TYPE", "linkedin_member" + ).strip() + additional_scopes = provider.account_type_scopes(account_type) + try: + authorization_url = await provider.get_authorization_url( + state="phase6c-live-linkedin-state-value-that-is-long-enough", + redirect_uri=_required("LINKEDIN_REDIRECT_URI"), + additional_scopes=additional_scopes, + ) + parsed = urlparse(authorization_url) + assert parsed.scheme == "https" + assert parsed.netloc == "www.linkedin.com" + assert parsed.path == "/oauth/v2/authorization" + assert "code_challenge" not in parse_qs(parsed.query) + discovered_type, external_id = await _identity(provider, token) + assert discovered_type == account_type + assert external_id + finally: + await provider.close() + + +async def test_linkedin_live_authorization_code_exchange_when_supplied() -> None: + """A fresh one-time browser code is optional and never required by CI.""" + + _require_base_configuration() + code = os.getenv("LINKEDIN_LIVE_TEST_AUTHORIZATION_CODE", "").strip() + if not code: + pytest.skip( + "LinkedIn OAuth code exchange is NOT VERIFIED; provide a fresh " + "LINKEDIN_LIVE_TEST_AUTHORIZATION_CODE." + ) + from app.social.providers.linkedin import LinkedInProvider + + provider = LinkedInProvider(_settings()) + try: + token = await provider.exchange_code( + code=code, + redirect_uri=_required("LINKEDIN_REDIRECT_URI"), + ) + account = await provider.get_account(token) + assert account["external_account_id"] + finally: + await provider.close() + + +async def test_linkedin_live_text_publish_status_and_delete() -> None: + _require_publish_consent() + from app.social.providers.linkedin import LinkedInProvider + + provider = LinkedInProvider(_settings()) + token = _token() + account_type = os.getenv( + "LINKEDIN_LIVE_TEST_ACCOUNT_TYPE", "linkedin_member" + ).strip() + read_scope = { + "linkedin_member": "r_member_social", + "linkedin_organization": "r_organization_social", + }.get(account_type) + if read_scope is None or read_scope not in _granted_scopes(): + await provider.close() + pytest.skip( + "LinkedIn live status reconciliation is NOT VERIFIED; declare the " + f"approved {read_scope or 'account read'} scope." + ) + external_id: str | None = None + try: + result = await _publish_text(provider, token) + external_id = str(result["id"]) + status: dict[str, object] | None = None + for _ in range(12): + status = await provider.get_publish_status(token, external_id) + if status.get("status") in {"published", "failed", "deleted"}: + break + await asyncio.sleep(5) + assert status is not None and status.get("status") == "published" + finally: + if external_id: + await provider.delete_post(token, external_id) + await provider.close() + + +async def test_linkedin_live_media_publish_when_asset_is_supplied() -> None: + _require_publish_consent() + media_value = os.getenv("LINKEDIN_LIVE_TEST_MEDIA_PATH", "").strip() + if not media_value: + pytest.skip( + "LinkedIn media publishing is NOT VERIFIED; set " + "LINKEDIN_LIVE_TEST_MEDIA_PATH to a dedicated image or MP4 asset." + ) + from app.services.ffprobe_service import FFprobeService + from app.services.validator import MediaValidator + from app.social.providers.linkedin import LinkedInProvider + + media_path = Path(media_value).expanduser().resolve() + if not media_path.is_file(): + pytest.skip("LINKEDIN_LIVE_TEST_MEDIA_PATH is not a readable file.") + settings = _settings() + provider = LinkedInProvider(settings) + token = _token() + state: dict[str, object] = {} + + async def persist(value: dict[str, object]) -> None: + state.clear() + state.update(value) + + external_id: str | None = None + try: + account_type, account_id = await _identity(provider, token) + probe = await FFprobeService(settings).probe(media_path) + mime_type = MediaValidator(settings).infer_mime(media_path) + post_type = "image" if mime_type.startswith("image/") else "video" + media = { + "path": media_path, + "mime_type": mime_type, + "file_size": media_path.stat().st_size, + "probe": probe, + "provider_account_id": account_id, + "provider_account_type": account_type, + "linkedin_post_metadata": { + "post_type": post_type, + "commentary": "MediaRouter Phase 6C media verification", + }, + "provider_state": state, + "persist_provider_state": persist, + } + await provider.validate_media(media) + uploaded = await provider.upload_media(token, media) + result = await provider.publish( + token, + { + "provider_account_id": account_id, + "provider_account_type": account_type, + "linkedin_post_metadata": media["linkedin_post_metadata"], + "upload": uploaded, + "provider_state": state, + "persist_provider_state": persist, + }, + ) + external_id = str(result["id"]) + assert external_id.startswith("urn:li:") + finally: + if external_id: + await provider.delete_post(token, external_id) + await provider.close() + + +async def test_linkedin_live_analytics_when_authorized() -> None: + _require_publish_consent() + from app.social.providers.linkedin import LinkedInProvider + + provider = LinkedInProvider(_settings()) + token = _token() + external_id: str | None = None + try: + account_type, account_id = await _identity(provider, token) + required_scopes = provider.analytics_scopes(account_type) + if not required_scopes or required_scopes[0] not in _granted_scopes(): + pytest.skip( + "LinkedIn analytics are NOT VERIFIED; the dedicated token does " + "not declare the required analytics grant." + ) + result = await _publish_text(provider, token) + external_id = str(result["id"]) + metrics = await provider.get_metrics( + { + **token, + "_mediarouter_account_type": account_type, + "_mediarouter_external_account_id": account_id, + }, + external_id, + ) + assert metrics["status"] == "available" + assert isinstance(metrics.get("raw_metrics"), dict) + finally: + if external_id: + await provider.delete_post(token, external_id) + await provider.close() diff --git a/tests/test_linkedin_production.py b/tests/test_linkedin_production.py new file mode 100644 index 0000000000000000000000000000000000000000..42d6b1179ffab1ec0b3cdb89e120057658e42737 --- /dev/null +++ b/tests/test_linkedin_production.py @@ -0,0 +1,787 @@ +"""Phase 6C LinkedIn analytics, security, tenancy, and certification tests. + +Normal CI uses SQLite and mocked official LinkedIn REST traffic. Destructive +live verification is isolated in ``test_linkedin_live.py`` and is opt-in. +""" + +from __future__ import annotations + +import json +import logging +from datetime import datetime, timedelta, timezone +from pathlib import Path +from types import SimpleNamespace +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.core.logger import JsonFormatter +from app.mcp.registry import MCPRegistry +from app.mcp.server import create_mcp_server +from app.security.context import AuthContext, auth_context, http_auth_applied +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialJobNotFoundError, + SocialMediaInvalidError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialReauthRequiredError, +) +from app.social.domain.retry import classify_retry +from app.social.models import ( + SocialAccount, + SocialAuditEvent, + SocialJob, + SocialMediaAsset, + SocialPost, + SocialPostMetric, + SocialPostTarget, +) +from app.social.providers.linkedin import LINKEDIN_API_VERSION, LinkedInProvider +from app.social.schemas.accounts import SocialAccountConnectRequest, SocialAccountView +from app.social.schemas.jobs import SocialJobView +from app.social.workers.publisher import SocialPublisher + + +_REDIRECT_URI = "https://api.example.com/v1/social/accounts/linkedin/callback" + + +def phase6c_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "phase-6c-linkedin-encryption-material", + "social_oauth_redirect_base_url": "https://api.example.com", + "linkedin_client_id": "linkedin-client-id", + "linkedin_client_secret": "linkedin-client-secret", + "linkedin_redirect_uri": _REDIRECT_URI, + "linkedin_publishing_enabled": True, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +@pytest.fixture +async def phase6c_container(tmp_path: Path): + container = build_container(phase6c_settings(tmp_path)) + await container.social.initialize() + try: + yield container + finally: + await container.social.close() + await container.security_database.close() + + +async def _connected_account( + container: object, + workspace_id: str, + *, + account_type: str, + external_id: str, + scopes: list[str], +) -> SocialAccount: + social = container.social # type: ignore[attr-defined] + account = await social.accounts.repository.create( + SocialAccount( + workspace_id=workspace_id, + provider="linkedin", + account_type=account_type, + external_account_id=external_id, + display_name="LinkedIn production test", + status="connected", + metadata_json=( + {"roles": ["ADMINISTRATOR"]} + if account_type == "linkedin_organization" + else {} + ), + ) + ) + await social.accounts.tokens.store( + workspace_id, + account.id, + { + "access_token": "linkedin-provider-secret", + "refresh_token": "linkedin-refresh-secret", + }, + expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + scopes=scopes, + token_type="bearer", + ) + return account + + +async def test_linkedin_member_analytics_uses_official_endpoint_and_normalizes( + tmp_path: Path, +) -> None: + secret = "member-analytics-secret" + external_id = "urn:li:share:7325786486870552578" + counts = { + "IMPRESSION": 101, + "MEMBERS_REACHED": 88, + "REACTION": 22, + "COMMENT": 3, + "RESHARE": 4, + } + requested: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + assert request.url.path == "/rest/memberCreatorPostAnalytics" + assert request.headers["authorization"] == f"Bearer {secret}" + assert request.headers["linkedin-version"] == LINKEDIN_API_VERSION + assert request.headers["x-restli-protocol-version"] == "2.0.0" + query_type = request.url.params["queryType"] + requested.append(query_type) + assert request.url.params["q"] == "entity" + assert request.url.params["entity"] == f"(share:{external_id})" + assert request.url.params["aggregation"] == "TOTAL" + return httpx.Response( + 200, + json={ + "elements": [{ + "count": counts[query_type], + "targetEntity": {"share": external_id}, + "metricType": {"type": query_type}, + "access_token": secret, + }], + "paging": {"count": 10, "start": 0}, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(phase6c_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics( + { + "access_token": secret, + "_mediarouter_account_type": "linkedin_member", + "_mediarouter_external_account_id": "member_123", + }, + external_id, + ) + finally: + await client.aclose() + + assert requested == [ + "IMPRESSION", + "MEMBERS_REACHED", + "REACTION", + "COMMENT", + "RESHARE", + ] + assert result["status"] == "available" + assert result["impressions"] == 101 + assert result["likes"] == 22 + assert result["comments"] == 3 + assert result["shares"] == 4 + assert result["raw_metrics"]["members_reached"] == 88 + assert secret not in json.dumps(result) + + +async def test_linkedin_member_analytics_never_invents_an_omitted_metric( + tmp_path: Path, +) -> None: + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"elements": []}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(phase6c_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError): + await provider.get_metrics( + { + "access_token": "member-missing-metric-secret", + "_mediarouter_account_type": "linkedin_member", + "_mediarouter_external_account_id": "member_123", + }, + "urn:li:share:7325786486870552578", + ) + finally: + await client.aclose() + + +@pytest.mark.parametrize( + ("external_id", "query_key"), + [ + ("urn:li:share:7132564752928563200", "shares"), + ("urn:li:ugcPost:7132564752928563201", "ugcPosts[0]"), + ], +) +async def test_linkedin_organization_analytics_uses_official_share_statistics( + tmp_path: Path, + external_id: str, + query_key: str, +) -> None: + secret = "organization-analytics-secret" + organization_urn = "urn:li:organization:5515715" + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/rest/organizationalEntityShareStatistics" + assert request.headers["authorization"] == f"Bearer {secret}" + assert request.url.params["q"] == "organizationalEntity" + assert request.url.params["organizationalEntity"] == organization_urn + if query_key == "shares": + assert request.url.params[query_key] == f"List({external_id})" + else: + assert request.url.params[query_key] == external_id + field = "share" if external_id.startswith("urn:li:share:") else "ugcPost" + return httpx.Response( + 200, + json={ + "elements": [{ + "organizationalEntity": organization_urn, + field: external_id, + "totalShareStatistics": { + "clickCount": 7, + "commentCount": 3, + "engagement": 0.125, + "impressionCount": 101, + "likeCount": 22, + "shareCount": 4, + "refresh_token": secret, + }, + }], + "paging": {"count": 10, "start": 0}, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(phase6c_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics( + { + "access_token": secret, + "_mediarouter_account_type": "linkedin_organization", + "_mediarouter_external_account_id": "5515715", + }, + external_id, + ) + finally: + await client.aclose() + + assert result["status"] == "available" + assert result["impressions"] == 101 + assert result["likes"] == 22 + assert result["comments"] == 3 + assert result["shares"] == 4 + assert result["engagement_rate"] == 0.125 + assert result["raw_metrics"]["click_count"] == 7 + assert secret not in json.dumps(result) + + +async def test_linkedin_analytics_scopes_are_explicit_and_account_specific( + phase6c_container, +) -> None: + social = phase6c_container.social + normal = await social.oauth.connect( + provider="linkedin", + workspace_id="workspace-scope", + user_id="user-scope", + payload=SocialAccountConnectRequest(account_type="linkedin_member"), + ) + member = await social.oauth.connect( + provider="linkedin", + workspace_id="workspace-scope", + user_id="user-scope", + payload=SocialAccountConnectRequest( + account_type="linkedin_member", + authorization_purpose="analytics", + ), + ) + organization = await social.oauth.connect( + provider="linkedin", + workspace_id="workspace-scope", + user_id="user-scope", + payload=SocialAccountConnectRequest( + account_type="linkedin_organization", + authorization_purpose="analytics", + ), + ) + normal_scopes = parse_qs(urlparse(str(normal.authorization_url)).query)["scope"][0].split() + member_scopes = parse_qs(urlparse(str(member.authorization_url)).query)["scope"][0].split() + organization_scopes = parse_qs( + urlparse(str(organization.authorization_url)).query + )["scope"][0].split() + + assert normal_scopes == ["openid", "profile"] + assert member_scopes == ["openid", "profile", "r_member_postAnalytics"] + assert organization_scopes == ["openid", "profile", "rw_organization_admin"] + capabilities = social.accounts.providers.get("linkedin").capabilities + assert capabilities.analytics + assert capabilities.account_type_analytics_scopes == { + "linkedin_member": ["r_member_postAnalytics"], + "linkedin_organization": ["rw_organization_admin"], + } + + +async def test_linkedin_analytics_persists_normalized_metrics_and_raw_data( + phase6c_container, +) -> None: + social = phase6c_container.social + account = await _connected_account( + phase6c_container, + "workspace-analytics", + account_type="linkedin_organization", + external_id="5515715", + scopes=["openid", "profile", "rw_organization_admin"], + ) + post, targets = await social.publishing.posts.create( + SocialPost( + workspace_id="workspace-analytics", + status="published", + publish_mode="now", + ), + [SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="linkedin", + status="published", + external_post_id="urn:li:share:7132564752928563200", + )], + ) + adapter = social.accounts.providers.get("linkedin") + assert isinstance(adapter, LinkedInProvider) + await adapter._client.aclose() + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.headers["authorization"] == "Bearer linkedin-provider-secret" + return httpx.Response( + 200, + json={ + "elements": [{ + "organizationalEntity": "urn:li:organization:5515715", + "share": targets[0].external_post_id, + "totalShareStatistics": { + "clickCount": 9, + "commentCount": 4, + "engagement": 0.25, + "impressionCount": 120, + "likeCount": 30, + "shareCount": 5, + }, + }], + }, + ) + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + result = await social.analytics.post("workspace-analytics", post.id) + + assert result["unavailable"] == [] + assert result["metrics"][0]["impressions"] == 120 + assert result["metrics"][0]["engagement_rate"] == 0.25 + async with social.database.session("workspace-analytics") as session: + persisted = await session.scalar( + select(SocialPostMetric).where( + SocialPostMetric.social_post_target_id == targets[0].id + ) + ) + assert persisted is not None + assert persisted.likes == 30 + assert persisted.raw_metrics["click_count"] == 9 + assert "linkedin-provider-secret" not in json.dumps(result, default=str) + + +async def test_linkedin_missing_analytics_scope_fails_closed_without_provider_call( + phase6c_container, +) -> None: + social = phase6c_container.social + account = await _connected_account( + phase6c_container, + "workspace-no-analytics", + account_type="linkedin_member", + external_id="member_analytics", + scopes=["openid", "profile", "w_member_social"], + ) + post, _ = await social.publishing.posts.create( + SocialPost( + workspace_id="workspace-no-analytics", + status="published", + publish_mode="now", + ), + [SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="linkedin", + status="published", + external_post_id="urn:li:share:7132564752928563202", + )], + ) + result = await social.analytics.post("workspace-no-analytics", post.id) + assert result["metrics"] == [] + assert result["unavailable"] == [{ + "provider": "linkedin", + "status": "unavailable", + "reason": "LINKEDIN_ANALYTICS_ADDITIONAL_AUTHORIZATION_REQUIRED", + "required_scopes": ["r_member_postAnalytics"], + }] + + +async def test_linkedin_workspace_isolation_covers_all_phase6c_resources( + phase6c_container, +) -> None: + social = phase6c_container.social + member = await _connected_account( + phase6c_container, + "workspace-a", + account_type="linkedin_member", + external_id="member_a", + scopes=["openid", "profile", "r_member_postAnalytics"], + ) + organization = await _connected_account( + phase6c_container, + "workspace-a", + account_type="linkedin_organization", + external_id="5515715", + scopes=["openid", "profile", "rw_organization_admin"], + ) + post, targets = await social.publishing.posts.create( + SocialPost(workspace_id="workspace-a", status="published", publish_mode="now"), + [SocialPostTarget( + social_post_id="", + social_account_id=organization.id, + provider="linkedin", + status="published", + external_post_id="urn:li:share:7132564752928563203", + )], + ) + asset = await social.media_assets.repository.create( + SocialMediaAsset( + workspace_id="workspace-a", + request_id="11111111-1111-4111-8111-111111111111", + filename="owned.mp4", + mime_type="video/mp4", + file_size=80_000, + ) + ) + jobs = await social.jobs.repository.create_many([ + SocialJob( + workspace_id="workspace-a", + social_post_id=post.id, + social_post_target_id=targets[0].id, + provider="linkedin", + status="queued", + idempotency_key="workspace-a-job", + ) + ]) + async with social.database.session("workspace-a") as session: + session.add(SocialPostMetric( + social_post_id=post.id, + social_post_target_id=targets[0].id, + provider="linkedin", + impressions=1, + raw_metrics={"source": "official"}, + )) + await session.commit() + + for account_id in (member.id, organization.id): + with pytest.raises(SocialAccountNotFoundError): + await social.accounts.repository.get("workspace-b", account_id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.posts.get("workspace-b", post.id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.posts.set_target_status( + "workspace-b", targets[0].id, "failed" + ) + with pytest.raises(SocialJobNotFoundError): + await social.jobs.repository.get("workspace-b", jobs[0].id) + with pytest.raises(SocialMediaInvalidError): + await social.media_assets.repository.get("workspace-b", asset.id) + with pytest.raises(SocialPostNotFoundError): + await social.analytics.post("workspace-b", post.id) + + +async def test_linkedin_revoked_and_expired_credentials_require_reauthorization( + phase6c_container, +) -> None: + social = phase6c_container.social + revoked = await _connected_account( + phase6c_container, + "workspace-token", + account_type="linkedin_member", + external_id="member_revoked", + scopes=["openid", "profile", "r_member_postAnalytics"], + ) + await social.accounts.tokens.revoke("workspace-token", revoked.id) + with pytest.raises(SocialReauthRequiredError): + await social.oauth.token_for_request( + workspace_id="workspace-token", account_id=revoked.id + ) + + expired = await _connected_account( + phase6c_container, + "workspace-token", + account_type="linkedin_member", + external_id="member_expired", + scopes=["openid", "profile", "r_member_postAnalytics"], + ) + await social.accounts.tokens.store( + "workspace-token", + expired.id, + {"access_token": "expired-linkedin-token"}, + expires_at=datetime.now(timezone.utc) - timedelta(minutes=1), + scopes=["openid", "profile", "r_member_postAnalytics"], + ) + with pytest.raises(SocialReauthRequiredError): + await social.oauth.token_for_request( + workspace_id="workspace-token", account_id=expired.id + ) + + +@pytest.mark.parametrize( + ("status_code", "retryable", "reauth"), + [ + (429, True, False), + (500, True, False), + (502, True, False), + (503, True, False), + (504, True, False), + (401, True, True), + (403, False, False), + (400, False, False), + ], +) +def test_linkedin_retry_policy_is_bounded_and_classified( + status_code: int, retryable: bool, reauth: bool +) -> None: + decision = classify_retry(status_code=status_code, attempt=1) + assert decision.retryable is retryable + assert decision.refresh_token_first is reauth + if status_code == 401: + assert not classify_retry(status_code=401, attempt=2).retryable + + +async def test_linkedin_transient_retry_stops_at_job_attempt_limit() -> None: + transitions: list[str] = [] + + class Jobs: + async def complete_attempt(self, *_: object, **__: object) -> None: + return None + + async def transition( + self, _: str, __: str, status: str, **___: object + ) -> SocialJob: + transitions.append(status) + return job + + class Audit: + async def record(self, **_: object) -> None: + return None + + job = SocialJob( + id="linkedin-job-limit", + workspace_id="workspace-limit", + social_post_id="linkedin-post-limit", + provider="linkedin", + status="publishing", + attempt_count=5, + max_attempts=5, + ) + publisher = SocialPublisher( + SimpleNamespace( + jobs=SimpleNamespace(repository=Jobs()), + audit=Audit(), + ) + ) + await publisher._handle_failure( + "workspace-limit", + job, + "linkedin-attempt-limit", + SocialProviderUnavailableError("temporary LinkedIn failure"), + ) + assert transitions == ["failed"] + + +async def test_linkedin_status_reconciliation_covers_all_normalized_states( + tmp_path: Path, +) -> None: + external_ids = { + "urn:li:share:7132564752928563210": "published", + "urn:li:share:7132564752928563211": "processing", + "urn:li:share:7132564752928563212": "failed", + "urn:li:share:7132564752928563213": "deleted", + "urn:li:share:7132564752928563214": "unavailable", + } + + async def handler(request: httpx.Request) -> httpx.Response: + encoded = request.url.path.rsplit("/", 1)[-1] + external_id = next(key for key in external_ids if key.split(":")[-1] in encoded) + expected = external_ids[external_id] + if expected == "deleted": + return httpx.Response(404, json={"message": "not found"}) + lifecycle = { + "published": "PUBLISHED", + "processing": "PUBLISH_REQUESTED", + "failed": "PUBLISH_FAILED", + "unavailable": "UNKNOWN_PROVIDER_STATE", + }[expected] + return httpx.Response(200, json={"lifecycleState": lifecycle}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(phase6c_settings(tmp_path), http_client=client) + token = { + "access_token": "linkedin-status-secret", + "_mediarouter_granted_scopes": ["r_organization_social"], + } + try: + results = { + external_id: await provider.get_publish_status(token, external_id) + for external_id in external_ids + } + finally: + await client.aclose() + assert {key: value["status"] for key, value in results.items()} == external_ids + assert "linkedin-status-secret" not in json.dumps(results) + + +async def test_linkedin_public_views_logs_and_audit_boundaries_redact_credentials( + phase6c_container, +) -> None: + secret = "linkedin-secret-never-expose" + account = SocialAccount( + workspace_id="workspace-security", + provider="linkedin", + account_type="linkedin_member", + external_account_id="member_secure", + status="connected", + metadata_json={"access_token": secret, "name": "Safe member"}, + ) + job = SocialJob( + workspace_id="workspace-security", + social_post_id="post-security", + provider="linkedin", + status="queued", + payload_json={"refresh_token": secret}, + provider_state_encrypted=secret, + ) + assert secret not in SocialAccountView.from_record(account).model_dump_json() + assert secret not in SocialJobView.from_record(job).model_dump_json() + record = logging.LogRecord( + "linkedin-security", + logging.ERROR, + __file__, + 1, + f"Authorization: Bearer {secret}", + (), + None, + ) + record.provider_payload = { + "refresh_token": secret, + "message": f"access_token={secret}", + } + assert secret not in JsonFormatter().format(record) + + await phase6c_container.social.audit.record( + workspace_id="workspace-security", + event_type="SOCIAL_LINKEDIN_SECURITY_TEST", + provider="linkedin", + metadata={ + "client_secret": secret, + "message": f"Authorization: Bearer {secret}", + }, + ) + async with phase6c_container.social.database.session( + "workspace-security" + ) as session: + audit = await session.scalar( + select(SocialAuditEvent).where( + SocialAuditEvent.event_type + == "SOCIAL_LINKEDIN_SECURITY_TEST" + ) + ) + assert audit is not None + assert secret not in json.dumps(audit.metadata_json) + + +async def test_linkedin_mcp_contract_enforces_scope_and_never_exposes_credentials( + phase6c_container, +) -> None: + server = create_mcp_server(phase6c_container) + tools = {tool.name for tool in await server.list_tools()} + assert { + "social.list_providers", + "social.get_capabilities", + "social.list_accounts", + "social.create_post", + "social.publish_post", + "social.schedule_post", + "social.get_job", + "social.get_analytics", + } <= tools + context = AuthContext( + api_key_id="workspace-linkedin", + key_name="phase-6c", + key_prefix="mp_test", + environment="test", + role="viewer", + scopes=frozenset({"social:accounts:read"}), + requests_per_minute=100, + concurrent_jobs=2, + uploads_per_hour=10, + processing_bytes_per_day=1_000_000, + expires_at=None, + ) + auth_token = auth_context.set(context) + http_token = http_auth_applied.set(True) + called = False + + async def forbidden_action() -> dict[str, object]: + nonlocal called + called = True + return {"access_token": "must-not-appear"} + + try: + result = await MCPRegistry(phase6c_container).run_metadata_tool( + "social.get_analytics", + forbidden_action, + required_scope="social:analytics:read", + ) + finally: + http_auth_applied.reset(http_token) + auth_context.reset(auth_token) + assert result["success"] is False + assert result["error"]["code"] == "FORBIDDEN" + assert not called + assert "must-not-appear" not in json.dumps(result) + + +async def test_linkedin_analytics_timeout_is_retryable_and_secret_safe( + tmp_path: Path, +) -> None: + secret = "linkedin-timeout-secret" + + async def handler(request: httpx.Request) -> httpx.Response: + raise httpx.ReadTimeout( + f"Authorization: Bearer {secret}", request=request + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(phase6c_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError) as raised: + await provider.get_metrics( + { + "access_token": secret, + "_mediarouter_account_type": "linkedin_organization", + "_mediarouter_external_account_id": "5515715", + }, + "urn:li:share:7132564752928563220", + ) + finally: + await client.aclose() + assert secret not in str(raised.value) + assert classify_retry( + status_code=raised.value.status_code, attempt=1 + ).retryable diff --git a/tests/test_linkedin_publishing.py b/tests/test_linkedin_publishing.py new file mode 100644 index 0000000000000000000000000000000000000000..764a42ad354937152b864fa1f7e51ae4d2e735d5 --- /dev/null +++ b/tests/test_linkedin_publishing.py @@ -0,0 +1,745 @@ +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from pydantic import ValidationError + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialCapabilityUnsupportedError, + SocialIdempotencyConflictError, + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialPublishFailedError, + SocialRateLimitedError, + SocialReauthRequiredError, +) +from app.social.models import SocialAccount +from app.social.providers.linkedin import LINKEDIN_API_VERSION, LinkedInProvider +from app.social.schemas.linkedin import LinkedInPostMetadata +from app.social.schemas.posts import SocialPostCreate +from app.social.workers.publisher import SocialPublisher + + +_REDIRECT_URI = "https://api.example.com/v1/social/accounts/linkedin/callback" + + +def publishing_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "phase-6b-linkedin-test-encryption-material", + "social_oauth_redirect_base_url": "https://api.example.com", + "linkedin_client_id": "linkedin-client-id", + "linkedin_client_secret": "linkedin-client-secret", + "linkedin_redirect_uri": _REDIRECT_URI, + "linkedin_publishing_enabled": True, + "linkedin_media_processing_poll_seconds": 1, + "linkedin_media_processing_timeout_seconds": 30, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +def image_probe(*, codec: str = "png", frames: int = 1) -> dict[str, object]: + return { + "container": "png_pipe", + "duration": None, + "fps": 1.0, + "resolution": {"width": 1200, "height": 675}, + "video_streams": [{"codec": codec, "frame_count": frames}], + "audio_streams": [], + } + + +def video_probe(**overrides: object) -> dict[str, object]: + values: dict[str, object] = { + "container": "mov,mp4,m4a,3gp,3g2,mj2", + "duration": 15.0, + "fps": 29.97, + "resolution": {"width": 1280, "height": 720}, + "video_streams": [{"codec": "h264"}], + "audio_streams": [{"codec": "aac", "sample_rate": 48_000}], + } + values.update(overrides) + return values + + +def linkedin_post_payload( + account_id: str, + *, + commentary: str = "A production-safe LinkedIn update", + publish_mode: str = "draft", + scheduled_at: datetime | None = None, +) -> SocialPostCreate: + value: dict[str, object] = { + "publish_mode": publish_mode, + "targets": [{ + "social_account_id": account_id, + "caption": {"commentary": commentary}, + "linkedin": { + "post_type": "text", + "commentary": commentary, + }, + }], + } + if scheduled_at is not None: + value.update({"scheduled_at": scheduled_at, "timezone": "Africa/Lagos"}) + return SocialPostCreate.model_validate(value) + + +async def connected_linkedin_account( + container: object, + workspace_id: str, + *, + scopes: list[str] | None = None, +) -> SocialAccount: + social = container.social # type: ignore[attr-defined] + account = await social.accounts.repository.create( + SocialAccount( + workspace_id=workspace_id, + provider="linkedin", + account_type="linkedin_organization", + external_account_id="5515715", + username="mediarouter", + display_name="MediaRouter", + status="connected", + metadata_json={ + "organization_urn": "urn:li:organization:5515715", + "roles": ["ADMINISTRATOR"], + }, + ) + ) + await social.accounts.tokens.store( + workspace_id, + account.id, + {"access_token": "linkedin-provider-token"}, + expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + scopes=scopes + or [ + "openid", + "profile", + "rw_organization_admin", + "w_organization_social", + ], + token_type="bearer", + ) + return account + + +async def test_capabilities_and_oauth_scopes_are_account_specific_and_gated( + tmp_path: Path, +) -> None: + disabled = LinkedInProvider( + publishing_settings(tmp_path, linkedin_publishing_enabled=False) + ) + enabled = LinkedInProvider(publishing_settings(tmp_path)) + try: + assert not disabled.capabilities.direct_publish + assert disabled.capabilities.account_type_publishing_scopes == {} + assert enabled.capabilities.text + assert enabled.capabilities.image + assert enabled.capabilities.video + assert enabled.capabilities.link + assert enabled.capabilities.scheduled_publish + assert not enabled.capabilities.native_scheduling + assert enabled.capabilities.delete_post + assert enabled.publishing_scopes("linkedin_member") == ["w_member_social"] + assert enabled.publishing_scopes("linkedin_organization") == [ + "w_organization_social" + ] + with pytest.raises(SocialCapabilityUnsupportedError): + enabled.publishing_scopes("unsupported") + + member = await enabled.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + additional_scopes=enabled.publishing_scopes("linkedin_member"), + ) + organization = await enabled.get_authorization_url( + state="o" * 43, + redirect_uri=_REDIRECT_URI, + additional_scopes=[ + *enabled.account_type_scopes("linkedin_organization"), + *enabled.publishing_scopes("linkedin_organization"), + ], + ) + assert parse_qs(urlparse(member).query)["scope"] == [ + "openid profile w_member_social" + ] + assert parse_qs(urlparse(organization).query)["scope"] == [ + "openid profile rw_organization_admin w_organization_social" + ] + finally: + await disabled.close() + await enabled.close() + + +def test_linkedin_metadata_is_typed_and_media_requirement_is_explicit() -> None: + text = LinkedInPostMetadata.model_validate( + {"post_type": "text", "commentary": "Production post"} + ) + assert text.to_post_body(author_urn="urn:li:person:member1")[ + "lifecycleState" + ] == "PUBLISHED" + with pytest.raises(ValidationError): + LinkedInPostMetadata.model_validate( + {"post_type": "text", "commentary": "ok", "provider_payload": {}} + ) + with pytest.raises(ValidationError): + LinkedInPostMetadata.model_validate({"post_type": "link"}) + with pytest.raises(ValidationError): + SocialPostCreate.model_validate( + { + "targets": [ + { + "social_account_id": "linkedin-account", + "linkedin": {"post_type": "video"}, + } + ] + } + ) + link = SocialPostCreate.model_validate( + { + "targets": [ + { + "social_account_id": "linkedin-account", + "linkedin": { + "post_type": "link", + "link": { + "source": "https://example.com/article", + "title": "Explicit title", + "description": "Explicit description", + }, + }, + } + ] + } + ) + assert link.media_asset_id is None + + +@pytest.mark.parametrize( + ("account_type", "account_id", "expected_author"), + [ + ("linkedin_member", "member_123", "urn:li:person:member_123"), + ( + "linkedin_organization", + "5515715", + "urn:li:organization:5515715", + ), + ], +) +async def test_member_and_organization_text_publishing_use_posts_api( + tmp_path: Path, + account_type: str, + account_id: str, + expected_author: str, +) -> None: + requests: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.url == "https://api.linkedin.com/rest/posts" + assert request.headers["linkedin-version"] == LINKEDIN_API_VERSION + assert request.headers["x-restli-protocol-version"] == "2.0.0" + assert request.headers["authorization"] == "Bearer linkedin-token" + body = json.loads(request.content) + assert body["author"] == expected_author + assert body["commentary"] == "Production post" + assert body["distribution"]["feedDistribution"] == "MAIN_FEED" + return httpx.Response( + 201, + headers={"x-restli-id": "urn:li:share:6844785523593134080"}, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + states: list[dict[str, object] | None] = [] + + async def persist(value: dict[str, object] | None) -> None: + states.append(value) + + try: + result = await provider.publish( + {"access_token": "linkedin-token"}, + { + "provider_account_id": account_id, + "provider_account_type": account_type, + "linkedin_post_metadata": { + "post_type": "text", + "commentary": "Production post", + }, + "provider_state": {"linkedin_post_submission_attempted": False}, + "persist_provider_state": persist, + "upload": {"identity_type": "none"}, + }, + ) + assert result["id"] == "urn:li:share:6844785523593134080" + assert result["status"] == "published" + assert len(requests) == 1 + assert states[-1]["linkedin_post_submission_attempted"] is True # type: ignore[index] + finally: + await client.aclose() + + +async def test_image_upload_streams_with_oauth_and_creates_media_urn( + tmp_path: Path, +) -> None: + path = tmp_path / "image.png" + path.write_bytes(b"image" * 1024) + expires = int((datetime.now(timezone.utc) + timedelta(hours=1)).timestamp() * 1000) + calls: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(f"{request.method} {request.url.path}") + if request.url.path == "/rest/images": + assert request.url.params["action"] == "initializeUpload" + assert json.loads(request.content)["initializeUploadRequest"]["owner"] == ( + "urn:li:organization:5515715" + ) + return httpx.Response( + 200, + json={ + "value": { + "uploadUrlExpiresAt": expires, + "uploadUrl": "https://www.linkedin.com/dms-uploads/image/upload", + "image": "urn:li:image:C4E10AQFoyyAjHPMQuQ", + } + }, + ) + assert request.url == "https://www.linkedin.com/dms-uploads/image/upload" + assert request.method == "PUT" + assert request.headers["authorization"] == "Bearer linkedin-token" + assert len(request.content) == path.stat().st_size + return httpx.Response(201) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + state: dict[str, object] = {} + + async def persist(value: dict[str, object] | None) -> None: + state.clear() + state.update(value or {}) + + try: + result = await provider.upload_media( + {"access_token": "linkedin-token"}, + { + "path": path, + "file_size": path.stat().st_size, + "mime_type": "image/png", + "probe": image_probe(), + "provider_account_id": "5515715", + "provider_account_type": "linkedin_organization", + "linkedin_post_metadata": { + "post_type": "image", + "commentary": "Image post", + "image_alt_text": "Accessible description", + }, + "provider_state": {}, + "persist_provider_state": persist, + }, + ) + assert result["id"] == "urn:li:image:C4E10AQFoyyAjHPMQuQ" + assert state["linkedin_image_uploaded"] is True + assert calls == ["POST /rest/images", "PUT /dms-uploads/image/upload"] + assert "linkedin-token" not in str(state) + finally: + await client.aclose() + + +async def test_video_multipart_upload_streams_ranges_and_finalizes( + tmp_path: Path, +) -> None: + path = tmp_path / "video.mp4" + path.write_bytes(b"a" * 80_000) + expires = int((datetime.now(timezone.utc) + timedelta(hours=1)).timestamp() * 1000) + put_bodies: list[bytes] = [] + finalized: list[dict[str, object]] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/rest/videos" and request.url.params.get("action") == "initializeUpload": + return httpx.Response( + 200, + json={ + "value": { + "video": "urn:li:video:C4E10AQEfKKMV9a1d-g", + "uploadToken": "opaque-upload-token", + "uploadUrlsExpireAt": expires, + "uploadInstructions": [ + { + "firstByte": 0, + "lastByte": 39_999, + "uploadUrl": "https://www.linkedin.com/dms-uploads/video/part-0", + }, + { + "firstByte": 40_000, + "lastByte": 79_999, + "uploadUrl": "https://www.linkedin.com/dms-uploads/video/part-1", + }, + ], + } + }, + ) + if request.method == "PUT": + assert "authorization" not in request.headers + put_bodies.append(request.content) + return httpx.Response(200, headers={"ETag": f'"part-{len(put_bodies)}"'}) + assert request.url.params["action"] == "finalizeUpload" + finalized.append(json.loads(request.content)["finalizeUploadRequest"]) + return httpx.Response(200, json={}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + state: dict[str, object] = {} + + async def persist(value: dict[str, object] | None) -> None: + state.clear() + state.update(value or {}) + + try: + result = await provider.upload_media( + { + "access_token": "linkedin-token", + "_mediarouter_granted_scopes": ["w_member_social"], + }, + { + "path": path, + "file_size": path.stat().st_size, + "mime_type": "video/mp4", + "probe": video_probe(), + "provider_account_id": "member_123", + "provider_account_type": "linkedin_member", + "linkedin_post_metadata": { + "post_type": "video", + "commentary": "Video post", + }, + "provider_state": {}, + "persist_provider_state": persist, + }, + ) + assert result["id"] == "urn:li:video:C4E10AQEfKKMV9a1d-g" + assert [len(body) for body in put_bodies] == [40_000, 40_000] + assert finalized[0]["uploadedPartIds"] == ["part-1", "part-2"] + assert state["linkedin_video_finalized"] is True + assert "linkedin-token" not in str(state) + finally: + await client.aclose() + + +async def test_status_reconciliation_and_idempotent_delete_use_encoded_post_urn( + tmp_path: Path, +) -> None: + methods: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + methods.append(request.method) + assert request.url.raw_path.decode().split("?", 1)[0].endswith( + "/urn%3Ali%3Ashare%3A6844785523593134080" + ) + if request.method == "GET": + assert request.url.params["viewContext"] == "AUTHOR" + return httpx.Response( + 200, + json={ + "id": "urn:li:share:6844785523593134080", + "lifecycleState": "PUBLISH_REQUESTED", + }, + ) + assert request.headers["x-restli-method"] == "DELETE" + return httpx.Response(204) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + token = { + "access_token": "linkedin-token", + "_mediarouter_granted_scopes": [ + "w_organization_social", + "r_organization_social", + ], + } + try: + status = await provider.get_publish_status( + token, "urn:li:share:6844785523593134080" + ) + await provider.delete_post( + token, "urn:li:share:6844785523593134080" + ) + assert status["status"] == "processing" + assert methods == ["GET", "DELETE"] + finally: + await client.aclose() + + +async def test_uncertain_linkedin_create_outcome_never_resubmits( + tmp_path: Path, +) -> None: + create_calls = 0 + state: dict[str, object] = { + "linkedin_post_submission_attempted": False, + "linkedin_publish_started_at": datetime.now(timezone.utc).isoformat(), + } + + async def persist(value: dict[str, object] | None) -> None: + state.clear() + state.update(value or {}) + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal create_calls + create_calls += 1 + raise httpx.ReadTimeout( + "response lost after provider acceptance", request=request + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + payload = { + "provider_account_id": "member_123", + "provider_account_type": "linkedin_member", + "linkedin_post_metadata": { + "post_type": "text", + "commentary": "One logical post", + }, + "provider_state": state, + "persist_provider_state": persist, + "upload": {"identity_type": "none"}, + } + try: + with pytest.raises(SocialProviderUnavailableError): + await provider.publish( + {"access_token": "linkedin-token"}, payload + ) + assert state["linkedin_post_submission_attempted"] is True + with pytest.raises(SocialProviderUnavailableError, match="uncertain"): + await provider.reconcile_pending_publish( + {"access_token": "linkedin-token"}, + {**payload, "provider_state": state}, + ) + with pytest.raises(SocialProviderUnavailableError, match="uncertain"): + await provider.publish( + {"access_token": "linkedin-token"}, + {**payload, "provider_state": state}, + ) + assert create_calls == 1 + finally: + await client.aclose() + + +async def test_media_validation_status_delete_errors_and_unknown_outcome( + tmp_path: Path, +) -> None: + path = tmp_path / "video.mp4" + path.write_bytes(b"a" * 80_000) + provider = LinkedInProvider(publishing_settings(tmp_path)) + try: + with pytest.raises(SocialMediaInvalidError): + await provider.validate_media( + { + "path": path, + "file_size": path.stat().st_size, + "mime_type": "video/mp4", + "probe": video_probe(fps=30.0), + "linkedin_post_metadata": {"post_type": "video"}, + } + ) + with pytest.raises(SocialProviderUnavailableError): + await provider.reconcile_pending_publish( + {"access_token": "secret"}, + { + "provider_state": { + "linkedin_post_submission_attempted": True + } + }, + ) + disabled = LinkedInProvider( + publishing_settings(tmp_path, linkedin_publishing_enabled=False) + ) + try: + assert disabled.publishing_scopes("linkedin_member") == [] + finally: + await disabled.close() + finally: + await provider.close() + + +async def test_linkedin_idempotency_scheduling_authorization_and_workspace_isolation( + tmp_path: Path, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + try: + account = await connected_linkedin_account(container, "workspace-a") + payload = linkedin_post_payload(account.id) + first = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="linkedin-create-key", + ) + duplicate = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="linkedin-create-key", + ) + assert duplicate.id == first.id + with pytest.raises(SocialIdempotencyConflictError): + await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=linkedin_post_payload( + account.id, commentary="Different payload" + ), + idempotency_key="linkedin-create-key", + ) + scheduled = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=linkedin_post_payload( + account.id, + publish_mode="schedule", + scheduled_at=datetime.now(timezone.utc) + timedelta(hours=1), + ), + idempotency_key="linkedin-schedule-key", + ) + assert scheduled.status.value == "scheduled" + with pytest.raises(SocialAccountNotFoundError): + await container.social.publishing.create( + workspace_id="workspace-b", + user_id="user-b", + payload=linkedin_post_payload(account.id), + idempotency_key="workspace-b-key", + ) + with pytest.raises(SocialPostNotFoundError): + await container.social.publishing.delete("workspace-b", first.id) + + read_only = await connected_linkedin_account( + container, + "workspace-read-only", + scopes=["openid", "profile", "rw_organization_admin"], + ) + with pytest.raises(SocialPermissionDeniedError): + await container.social.publishing.create( + workspace_id="workspace-read-only", + user_id="user-read-only", + payload=linkedin_post_payload( + read_only.id, publish_mode="now" + ), + idempotency_key="missing-linkedin-write-scope", + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_linkedin_worker_publishes_once_and_persists_safe_identity( + tmp_path: Path, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + adapter = container.social.accounts.providers.get("linkedin") + assert isinstance(adapter, LinkedInProvider) + await adapter._client.aclose() + create_calls = 0 + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal create_calls + assert request.url.path == "/rest/posts" + create_calls += 1 + return httpx.Response( + 201, + headers={"x-restli-id": "urn:li:share:6844785523593134081"}, + ) + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + try: + account = await connected_linkedin_account( + container, "workspace-worker" + ) + post = await container.social.publishing.create( + workspace_id="workspace-worker", + user_id="worker-user", + payload=linkedin_post_payload( + account.id, + commentary="Worker post", + publish_mode="now", + ), + idempotency_key="linkedin-worker-key", + ) + jobs = await container.social.jobs.list("workspace-worker") + assert len(jobs) == 1 + await SocialPublisher(container.social).process( + "workspace-worker", jobs[0].id + ) + stored = await container.social.publishing.get( + "workspace-worker", post.id + ) + stored_job = await container.social.jobs.get( + "workspace-worker", jobs[0].id + ) + assert stored.status.value == "published" + assert ( + stored.targets[0].external_post_id + == "urn:li:share:6844785523593134081" + ) + assert stored_job.status.value == "published" + assert create_calls == 1 + serialized = stored.model_dump_json() + stored_job.model_dump_json() + assert "linkedin-provider-token" not in serialized + finally: + await container.social.close() + await container.security_database.close() + + +@pytest.mark.parametrize( + ("status", "error"), + [ + (400, SocialPublishFailedError), + (401, SocialReauthRequiredError), + (403, SocialPermissionDeniedError), + (429, SocialRateLimitedError), + (500, SocialProviderUnavailableError), + (502, SocialProviderUnavailableError), + (503, SocialProviderUnavailableError), + (504, SocialProviderUnavailableError), + ], +) +async def test_linkedin_error_normalization( + tmp_path: Path, status: int, error: type[Exception] +) -> None: + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(status, json={"message": "secret-provider-detail"}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = LinkedInProvider(publishing_settings(tmp_path), http_client=client) + try: + with pytest.raises(error) as raised: + await provider.delete_post( + {"access_token": "linkedin-token"}, + "urn:li:share:6844785523593134080", + ) + assert "linkedin-token" not in str(raised.value) + assert "secret-provider-detail" not in str(raised.value) + finally: + await client.aclose() + SocialIdempotencyConflictError, + SocialPostNotFoundError, diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 0000000000000000000000000000000000000000..74f00ad8fd79135a4238d516b611eed571354d51 --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,110 @@ +from __future__ import annotations + +from fastapi.testclient import TestClient + +from app.container import build_container +from app.mcp.registry import ( + AUDIO_TOOLS, + IMAGE_TOOLS, + PROBE_TOOLS, + SYSTEM_TOOLS, + TEMPLATE_TOOLS, + VIDEO_TOOLS, + WHISPER_TOOLS, + YTDLP_TOOLS, + MCPRegistry, + MediaInput, +) +from app.mcp.server import create_mcp_server +from main import create_app + + +async def test_mcp_registers_all_tools_resources_and_prompts(settings) -> None: + server = create_mcp_server(build_container(settings)) + + tools = {tool.name for tool in await server.list_tools()} + expected_tools = set( + VIDEO_TOOLS + + AUDIO_TOOLS + + IMAGE_TOOLS + + WHISPER_TOOLS + + YTDLP_TOOLS + + PROBE_TOOLS + + SYSTEM_TOOLS + + TEMPLATE_TOOLS + ) + expected_tools.update( + { + "social.list_providers", + "social.get_capabilities", + "social.list_media_assets", + "social.register_media_asset", + "social.list_accounts", + "social.get_account", + "social.create_post", + "social.publish_post", + "social.schedule_post", + "social.cancel_post", + "social.get_post", + "social.get_job", + "social.get_analytics", + "ai.capabilities", + "ai.generate", + "ai.list_jobs", + "ai.get_job", + "ai.cancel_job", + } + ) + assert tools == expected_tools + + resources = {str(resource.uri) for resource in await server.list_resources()} + assert resources == { + "media://operations", + "media://formats", + "media://codecs", + "media://health", + "media://configuration", + "media://version", + } + + prompts = {prompt.name for prompt in await server.list_prompts()} + assert prompts == { + "compress_for_social_media", + "youtube_to_mp3", + "download_and_transcribe", + "generate_subtitles", + "extract_audio", + "make_thumbnail", + "probe_media", + "instagram_reel", + "tiktok_video", + "podcast_audio", + } + + +async def test_mcp_errors_use_safe_structured_envelope(settings) -> None: + registry = MCPRegistry(build_container(settings)) + + response = await registry.run_probe( + "probe_media", + MediaInput(temp_path=str(settings.temp_dir / "missing.mp4")), + ) + + assert response["success"] is False + assert response["request_id"] + assert response["processing_time"] >= 0 + assert response["error"]["code"] == "INVALID_INPUT" + assert "traceback" not in str(response).lower() + + +def test_rest_and_mcp_coexist_in_one_application(settings) -> None: + application = create_app(settings) + mounts = {getattr(route, "path", None) for route in application.routes} + assert "/mcp" in mounts + + with TestClient(application) as client: + response = client.get("/health") + + assert response.status_code == 200 + assert response.json()["success"] is True + assert application.state.mcp_server is not None diff --git a/tests/test_membership_role_contract.py b/tests/test_membership_role_contract.py new file mode 100644 index 0000000000000000000000000000000000000000..d2738e547fffe3f4978726e15b508f151d12b5ff --- /dev/null +++ b/tests/test_membership_role_contract.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +import unittest + +from app.security.context import AuthContext +from app.security.errors import ForbiddenError +from app.security.service import APIKeyService + + +class MembershipRoleContractTests(unittest.TestCase): + def test_membership_role_is_informational_and_does_not_bypass_scope_authorization( + self, + ) -> None: + context = AuthContext( + api_key_id="key", + key_name="Contract Key", + key_prefix="mp_test_contract", + environment="test", + role="viewer", + scopes=frozenset({"admin"}), + requests_per_minute=60, + concurrent_jobs=1, + uploads_per_hour=1, + processing_bytes_per_day=1024, + membership_role="owner", + ) + + APIKeyService.authorize(context, "projects:write") + + def test_authoritative_scopes_remain_required_regardless_of_membership_role( + self, + ) -> None: + context = AuthContext( + api_key_id="key", + key_name="Contract Key", + key_prefix="mp_test_contract", + environment="test", + role="viewer", + scopes=frozenset(), + requests_per_minute=60, + concurrent_jobs=1, + uploads_per_hour=1, + processing_bytes_per_day=1024, + membership_role="owner", + ) + + with self.assertRaises(ForbiddenError): + APIKeyService.authorize(context, "projects:write") + diff --git a/tests/test_meta_production.py b/tests/test_meta_production.py new file mode 100644 index 0000000000000000000000000000000000000000..18780ece0d025ea09d959e9d6576469df5b4a528 --- /dev/null +++ b/tests/test_meta_production.py @@ -0,0 +1,268 @@ +"""Phase 3C unit coverage for Meta analytics and boundary hardening. + +All Graph calls use MockTransport. Live tests remain opt-in so normal CI never +requires a Page, professional account, browser consent, or Meta credentials. +""" + +from __future__ import annotations + +import os +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import SocialReauthRequiredError +from app.social.models import SocialAccount, SocialJob, SocialPost, SocialPostTarget +from app.social.providers.facebook import FacebookProvider +from app.social.providers.instagram import InstagramProvider +from app.social.schemas.accounts import SocialAccountView +from app.social.schemas.jobs import SocialJobView + + +def meta_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + meta_app_id="meta-app-id", + meta_app_secret="meta-app-secret", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +@pytest.fixture +async def meta_container(tmp_path: Path): + container = build_container(meta_settings(tmp_path)) + await container.social.initialize() + try: + yield container + finally: + await container.social.close() + await container.security_database.close() + + +async def test_facebook_page_metrics_use_v25_bearer_auth_and_normalize(tmp_path: Path) -> None: + async def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v25.0/page-post-id" + assert request.headers["authorization"] == "Bearer token-that-must-not-enter-url" + assert "access_token" not in request.url.query.decode() + return httpx.Response( + 200, + json={ + "created_time": "2026-07-31T00:00:00+0000", + "insights": { + "data": [ + {"name": "post_impressions", "values": [{"value": 42}]}, + {"name": "post_video_views", "values": [{"value": 11}]}, + ] + }, + "reactions": {"summary": {"total_count": 7}}, + "comments": {"summary": {"total_count": 3}}, + "shares": {"count": 2}, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = FacebookProvider(meta_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics({"access_token": "token-that-must-not-enter-url"}, "page-post-id") + finally: + await client.aclose() + assert result["status"] == "available" + assert result["impressions"] == 42 + assert result["views"] == 11 + assert result["likes"] == 7 + assert result["comments"] == 3 + assert result["shares"] == 2 + + +async def test_instagram_reel_metrics_are_media_type_specific(tmp_path: Path) -> None: + requests: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.headers["authorization"] == "Bearer meta-token" + if request.url.path == "/v25.0/ig-media-id": + return httpx.Response( + 200, + json={ + "media_product_type": "REELS", + "media_type": "VIDEO", + "timestamp": "2026-07-31T00:00:00+0000", + "permalink": "https://www.instagram.com/reel/example/", + }, + ) + assert request.url.path == "/v25.0/ig-media-id/insights" + assert parse_qs(request.url.query.decode())["metric"] == ["views,reach,likes,comments,shares,saved"] + return httpx.Response( + 200, + json={ + "data": [ + {"name": "views", "values": [{"value": 100}]}, + {"name": "reach", "values": [{"value": 80}]}, + {"name": "likes", "values": [{"value": 20}]}, + {"name": "comments", "values": [{"value": 4}]}, + {"name": "shares", "values": [{"value": 2}]}, + {"name": "saved", "values": [{"value": 9}]}, + ] + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = InstagramProvider(meta_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics({"access_token": "meta-token"}, "ig-media-id") + finally: + await client.aclose() + assert len(requests) == 2 + assert result["status"] == "available" + assert result["views"] == 100 + assert result["shares"] == 2 + assert result["url"] == "https://www.instagram.com/reel/example/" + + +async def test_meta_graph_authentication_error_is_safe_and_reauth_required(tmp_path: Path) -> None: + secret = "never-return-this-access-token" + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(400, json={"error": {"code": 190, "message": secret}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = FacebookProvider(meta_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialReauthRequiredError) as raised: + await provider.get_metrics({"access_token": secret}, "page-post-id") + finally: + await client.aclose() + assert secret not in str(raised.value) + + +async def test_meta_analytics_requires_explicit_authorization_and_persists_safe_snapshot(meta_container) -> None: + account = await meta_container.social.accounts.repository.create( + SocialAccount( + workspace_id="workspace-meta", + provider="facebook", + account_type="facebook_page", + external_account_id="page-id", + status="connected", + ) + ) + await meta_container.social.accounts.tokens.store( + "workspace-meta", account.id, {"access_token": "stored-token"}, scopes=["pages_read_engagement"] + ) + readiness = await meta_container.social.analytics.account("workspace-meta", account.id) + assert readiness["status"] == "unavailable" + assert readiness["reason"] == "META_ANALYTICS_ADDITIONAL_AUTHORIZATION_REQUIRED" + assert readiness["required_scopes"] == ["read_insights"] + + await meta_container.social.accounts.tokens.store( + "workspace-meta", + account.id, + {"access_token": "stored-token"}, + scopes=["pages_read_engagement", "read_insights"], + ) + post, targets = await meta_container.social.publishing.posts.create( + SocialPost(workspace_id="workspace-meta", media_asset_id="owned-asset"), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="facebook", + status="published", + external_post_id="page-post-id", + ) + ], + ) + assert targets + adapter = meta_container.social.accounts.providers.get("facebook") + + async def metrics(_: dict[str, object], __: str) -> dict[str, object]: + return { + "status": "available", + "views": 8, + "impressions": 12, + "likes": 3, + "comments": 1, + "shares": 2, + "raw_metrics": {"access_token": "must-not-leak", "provider_value": 8}, + } + + adapter.get_metrics = metrics # type: ignore[method-assign] + result = await meta_container.social.analytics.post("workspace-meta", post.id) + assert result["metrics"][0]["views"] == 8 + assert "access_token" not in str(result) + assert result["metrics"][0]["raw_metrics"] == {"provider_value": 8} + + +async def test_meta_analytics_consent_is_explicit_and_never_added_to_normal_connection(tmp_path: Path) -> None: + provider = FacebookProvider(meta_settings(tmp_path)) + try: + normal = await provider.get_authorization_url( + state="s" * 32, redirect_uri="https://api.example/callback" + ) + analytics = await provider.get_authorization_url( + state="a" * 32, + redirect_uri="https://api.example/callback", + additional_scopes=provider.capabilities.analytics_required_scopes, + ) + finally: + await provider.close() + assert "read_insights" not in parse_qs(urlparse(normal).query).get("scope", [""])[0] + requested = parse_qs(urlparse(analytics).query)["scope"][0].split() + assert requested == ["pages_read_engagement", "read_insights"] + + +def test_public_social_views_remove_token_like_data() -> None: + secret = "never-expose-me" + account = SocialAccount( + workspace_id="workspace-a", + provider="facebook", + account_type="facebook_page", + external_account_id="page-a", + status="connected", + metadata_json={"access_token": secret, "nested": {"client_secret": secret}, "name": "Page"}, + ) + job = SocialJob( + workspace_id="workspace-a", + social_post_id="post-a", + provider="facebook", + status="queued", + payload_json={"access_token": secret, "media_asset_id": "asset-a"}, + ) + assert secret not in SocialAccountView.from_record(account).model_dump_json() + assert secret not in SocialJobView.from_record(job).model_dump_json() + + +@pytest.mark.skipif( + os.getenv("RUN_META_INTEGRATION_TESTS") != "true", + reason="Set RUN_META_INTEGRATION_TESTS=true with dedicated Meta test credentials.", +) +async def test_live_meta_page_post_insights() -> None: + """Optional live smoke test; OAuth/publishing require separate manual consent setup. + + Required CI-secret variables are deliberately not named or logged by the + application. This test uses a dedicated Page post and never publishes. + """ + + token = os.environ.get("META_TEST_PAGE_ACCESS_TOKEN") + post_id = os.environ.get("META_TEST_PAGE_POST_ID") + if not token or not post_id: + pytest.skip("META_TEST_PAGE_ACCESS_TOKEN and META_TEST_PAGE_POST_ID are not configured.") + settings = Settings(_env_file=None, meta_graph_api_version="v25.0", whisper_model="tiny") + provider = FacebookProvider(settings) + try: + result = await provider.get_metrics({"access_token": token}, post_id) + finally: + await provider.close() + assert result["status"] in {"available", "unavailable"} diff --git a/tests/test_notifications.py b/tests/test_notifications.py new file mode 100644 index 0000000000000000000000000000000000000000..ed5c236a76ac8bbdd30f927fc4b1655cb896234a --- /dev/null +++ b/tests/test_notifications.py @@ -0,0 +1,21 @@ +import pytest +from app.projects.repositories.notification_repository import NotificationRepository +from app.projects.services.notification_service import NotificationService + +@pytest.mark.asyncio +async def test_notification_preferences(db_session): + repo = NotificationRepository(db_session) + service = NotificationService(repo) + + workspace_id = "ws1" + user_id = "user1" + event = "approval_request" + + # Test initial fetch (should be empty or defaults) + prefs = await service.get_preferences(workspace_id, user_id) + + # Test update + await service.update_preference(workspace_id, user_id, event, False) + + prefs = await service.get_preferences(workspace_id, user_id) + assert any(p['event_type'] == event and not p['enabled'] for p in prefs) diff --git a/tests/test_postgres_rls.py b/tests/test_postgres_rls.py new file mode 100644 index 0000000000000000000000000000000000000000..5423b36e9fb3559a227f165a6b0ded9a187fe717 --- /dev/null +++ b/tests/test_postgres_rls.py @@ -0,0 +1,298 @@ +"""Executable PostgreSQL RLS verification for the production tenant boundary. + +Run only against an isolated disposable database: + + SOCIAL_TEST_ADMIN_DATABASE_URL=postgresql://... # BYPASSRLS migration/seeding role + SOCIAL_TEST_TENANT_DATABASE_URL=postgresql://... # non-owner, non-BYPASSRLS role + pytest tests/test_postgres_rls.py + +The test never falls back to SQLite because SQLite cannot validate PostgreSQL +policies. It intentionally does not run in normal CI without those dedicated +credentials. +""" + +from __future__ import annotations + +import os +from pathlib import Path + +import pytest + +admin_url = os.getenv("SOCIAL_TEST_ADMIN_DATABASE_URL", "").strip() +tenant_url = os.getenv("SOCIAL_TEST_TENANT_DATABASE_URL", "").strip() +pytestmark = pytest.mark.skipif( + not (admin_url and tenant_url), + reason="PostgreSQL RLS integration credentials are not configured.", +) + + +def _asyncpg_url(value: str) -> str: + return value.replace("postgresql+asyncpg://", "postgresql://", 1) + + +async def _apply_migrations(connection: object, directory: Path) -> None: + for migration in sorted(directory.glob("*.sql")): + await connection.execute(migration.read_text(encoding="utf-8")) # type: ignore[attr-defined] + + +async def _tenant_context(connection: object, workspace_id: str, user_id: str) -> None: + await connection.execute("select set_config('app.workspace_id', $1, false)", workspace_id) # type: ignore[attr-defined] + await connection.execute("select set_config('app.user_id', $1, false)", user_id) # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_postgres_rls_rejects_cross_workspace_reads_and_writes() -> None: + asyncpg = pytest.importorskip("asyncpg") + root = Path(__file__).resolve().parents[1] + admin = await asyncpg.connect(_asyncpg_url(admin_url)) + tenant = await asyncpg.connect(_asyncpg_url(tenant_url)) + try: + # The configured database must be disposable and owned by the + # migration role. Never point these variables at a customer database. + await _apply_migrations(admin, root / "app" / "security" / "migrations") + await _apply_migrations(admin, root / "app" / "projects" / "migrations") + await _apply_migrations(admin, root / "app" / "social" / "migrations") + + role = await tenant.fetchrow( + "select r.rolsuper, r.rolbypassrls from pg_roles r where r.rolname = current_user" + ) + assert role is not None + assert not role["rolsuper"] and not role["rolbypassrls"] + + # Seed two fully independent tenants as the dedicated privileged role. + await admin.execute( + """ + insert into api_keys (id,name,key_prefix,key_hash,environment,status,scopes,created_at,requests_per_minute,concurrent_jobs,uploads_per_hour,processing_bytes_per_day) + values ('key-a','A','mp_test_aaaaaaaa','aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa','test','active','[]'::jsonb,now(),1,1,1,1048576), + ('key-b','B','mp_test_bbbbbbbb','bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb','test','active','[]'::jsonb,now(),1,1,1,1048576); + insert into users (id,subject,display_name) values ('user-a','test:user-a','A'),('user-b','test:user-b','B'); + insert into workspaces (id,slug,name) values ('workspace-a','test-a','A'),('workspace-b','test-b','B'); + insert into workspace_memberships (id,workspace_id,user_id,role) values ('membership-a','workspace-a','user-a','owner'),('membership-b','workspace-b','user-b','owner'); + insert into api_key_principals (id,api_key_id,workspace_id,user_id,membership_id) values ('principal-a','key-a','workspace-a','user-a','membership-a'),('principal-b','key-b','workspace-b','user-b','membership-b'); + insert into media_assets (id,workspace_id,request_id,filename,mime_type,file_size,sha256) values + ('asset-a','workspace-a','00000000-0000-0000-0000-000000000001','a.mp4','video/mp4',1,repeat('a',64)), + ('asset-b','workspace-b','00000000-0000-0000-0000-000000000002','b.mp4','video/mp4',1,repeat('b',64)); + insert into media_asset_variants (id,workspace_id,asset_id,request_id,filename,mime_type,file_size,sha256) values + ('asset-variant-a','workspace-a','asset-a','00000000-0000-0000-0000-000000000003','a-variant.mp4','video/mp4',1,repeat('c',64)), + ('asset-variant-b','workspace-b','asset-b','00000000-0000-0000-0000-000000000004','b-variant.mp4','video/mp4',1,repeat('d',64)); + insert into generation_requests (id,workspace_id,created_by_user_id,provider,model_id,modality,spec,request_fingerprint,idempotency_key) values + ('generation-request-a','workspace-a','user-a','test-provider','test-image','image','{"prompt":"A"}'::jsonb,repeat('1',64),'generation-key-a'), + ('generation-request-b','workspace-b','user-b','test-provider','test-image','image','{"prompt":"B"}'::jsonb,repeat('2',64),'generation-key-b'); + insert into generation_jobs (id,generation_request_id,workspace_id,provider,status,output_asset_id) values + ('generation-job-a','generation-request-a','workspace-a','test-provider','queued','asset-a'), + ('generation-job-b','generation-request-b','workspace-b','test-provider','queued','asset-b'); + insert into generation_job_attempts (id,generation_job_id,attempt_number,status) values + ('generation-attempt-a','generation-job-a',1,'started'),('generation-attempt-b','generation-job-b',1,'started'); + insert into projects (id,workspace_id,created_by,name) values + ('project-a','workspace-a','user-a','A'),('project-b','workspace-b','user-b','B'); + update media_assets set project_id = 'project-a' where id = 'asset-a'; + update media_assets set project_id = 'project-b' where id = 'asset-b'; + insert into project_generation_jobs (id,workspace_id,project_id,generation_job_id,attached_by) values + ('project-generation-a','workspace-a','project-a','generation-job-a','user-a'), + ('project-generation-b','workspace-b','project-b','generation-job-b','user-b'); + insert into project_editor_states (id,workspace_id,project_id,revision,schema_version,state,updated_by) values + ('editor-a','workspace-a','project-a',1,1,'{"schemaVersion":1,"projectId":"project-a","timeline":{"timeUnit":"milliseconds","tracks":[],"transitions":[],"markers":[]},"renderSettings":{"format":"mp4","width":1920,"height":1080,"frameRate":30}}'::jsonb,'user-a'), + ('editor-b','workspace-b','project-b',1,1,'{"schemaVersion":1,"projectId":"project-b","timeline":{"timeUnit":"milliseconds","tracks":[],"transitions":[],"markers":[]},"renderSettings":{"format":"mp4","width":1920,"height":1080,"frameRate":30}}'::jsonb,'user-b'); + insert into project_render_jobs (id,workspace_id,project_id,editor_revision,editor_schema_version,editor_state,render_settings,request_fingerprint,idempotency_key,requested_by) values + ('render-a','workspace-a','project-a',1,1,'{"schemaVersion":1,"projectId":"project-a"}'::jsonb,'{"format":"mp4"}'::jsonb,repeat('a',64),'render-key-a','user-a'), + ('render-b','workspace-b','project-b',1,1,'{"schemaVersion":1,"projectId":"project-b"}'::jsonb,'{"format":"mp4"}'::jsonb,repeat('b',64),'render-key-b','user-b'); + insert into audit_events (id,workspace_id,actor_user_id,event_type,entity_type,entity_id) values + ('project-audit-a','workspace-a','user-a','project.created','project','project-a'), + ('project-audit-b','workspace-b','user-b','project.created','project','project-b'); + + insert into social_accounts (id,workspace_id,provider,account_type,external_account_id,status) values + ('account-a','workspace-a','youtube','channel','a','connected'),('account-b','workspace-b','youtube','channel','b','connected'); + insert into social_account_tokens (id,social_account_id,encrypted_payload) values ('token-a','account-a','opaque'),('token-b','account-b','opaque'); + insert into social_account_capabilities (id,social_account_id,capability,enabled) values ('cap-a','account-a','publish',true),('cap-b','account-b','publish',true); + insert into media_variants (id,workspace_id,source_asset_id) values ('variant-a','workspace-a','asset-a'),('variant-b','workspace-b','asset-b'); + insert into social_media_assets (id,workspace_id,canonical_asset_id,request_id,filename,mime_type,file_size) values + ('social-asset-a','workspace-a','asset-a','00000000-0000-0000-0000-000000000001','a.mp4','video/mp4',1), + ('social-asset-b','workspace-b','asset-b','00000000-0000-0000-0000-000000000002','b.mp4','video/mp4',1); + insert into social_campaigns (id,workspace_id,name) values ('campaign-a','workspace-a','A'),('campaign-b','workspace-b','B'); + insert into social_webhook_events (id,provider,event_type,external_event_id,workspace_id) values ('webhook-a','youtube','TEST','webhook-a','workspace-a'),('webhook-b','youtube','TEST','webhook-b','workspace-b'); + insert into social_posts (id,workspace_id,campaign_id,media_asset_id,source_variant_id,status,publish_mode) values + ('post-a','workspace-a','campaign-a','social-asset-a','variant-a','draft','draft'), + ('post-b','workspace-b','campaign-b','social-asset-b','variant-b','draft','draft'); + insert into social_post_targets (id,social_post_id,social_account_id,provider) values ('target-a','post-a','account-a','youtube'),('target-b','post-b','account-b','youtube'); + insert into social_post_media (id,social_post_id,media_variant_id,media_asset_id) values ('post-media-a','post-a','variant-a','social-asset-a'),('post-media-b','post-b','variant-b','social-asset-b'); + insert into social_schedules (id,social_post_id,scheduled_at,timezone) values ('schedule-a','post-a',now(),'UTC'),('schedule-b','post-b',now(),'UTC'); + insert into social_jobs (id,workspace_id,social_post_id,social_post_target_id,provider,status) values ('job-a','workspace-a','post-a','target-a','youtube','queued'),('job-b','workspace-b','post-b','target-b','youtube','queued'); + insert into social_job_attempts (id,social_job_id,attempt_number,status) values ('attempt-a','job-a',1,'started'),('attempt-b','job-b',1,'started'); + insert into social_post_metrics (id,social_post_id,social_post_target_id,provider) values ('metric-a','post-a','target-a','youtube'),('metric-b','post-b','target-b','youtube'); + insert into social_audit_events (id,workspace_id,event_type) values ('audit-a','workspace-a','TEST'),('audit-b','workspace-b','TEST'); + """ + ) + + await _tenant_context(tenant, "workspace-a", "user-a") + for table in ( + "users", + "workspaces", + "workspace_memberships", + "api_key_principals", + "media_assets", + "media_asset_variants", + "generation_requests", + "generation_jobs", + "generation_job_attempts", + "social_accounts", + "projects", + "project_generation_jobs", + "project_editor_states", + "project_render_jobs", + "audit_events", + "social_account_tokens", + "social_account_capabilities", + "media_variants", + "social_media_assets", + "social_posts", + "social_post_targets", + "social_post_media", + "social_schedules", + "social_jobs", + "social_job_attempts", + "social_post_metrics", + "social_audit_events", + "social_campaigns", + "social_webhook_events", + ): + assert await tenant.fetchval(f"select count(*) from {table}") == 1, table + + # RLS WITH CHECK rejects direct reassignment to tenant B. The relation + # integrity triggers independently reject cross-tenant child links. + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into social_accounts (id,workspace_id,provider,account_type,external_account_id,status) values ('blocked-account','workspace-b','youtube','channel','blocked','connected')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into workspace_memberships (id,workspace_id,user_id,role) values ('blocked-membership','workspace-b','user-a','member')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into media_assets (id,workspace_id,request_id,filename,mime_type,file_size,sha256) values ('blocked-asset','workspace-b','00000000-0000-0000-0000-000000000006','blocked.mp4','video/mp4',1,repeat('f',64))" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into media_asset_variants (id,workspace_id,asset_id,request_id,filename,mime_type,file_size,sha256) values ('blocked-variant','workspace-a','asset-b','00000000-0000-0000-0000-000000000005','blocked.mp4','video/mp4',1,repeat('e',64))" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into generation_requests (id,workspace_id,created_by_user_id,provider,model_id,modality,spec,request_fingerprint,idempotency_key) values ('blocked-generation-request','workspace-b','user-a','test-provider','test-image','image','{}'::jsonb,repeat('3',64),'blocked-generation-key')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into projects (id,workspace_id,created_by,name) values ('blocked-project','workspace-b','user-a','Blocked')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update projects set thumbnail_asset_id = 'asset-b' where id = 'project-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update media_assets set project_id = 'project-b' where id = 'asset-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into project_generation_jobs (id,workspace_id,project_id,generation_job_id,attached_by) values ('blocked-project-job','workspace-a','project-a','generation-job-b','user-a')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into project_editor_states (id,workspace_id,project_id,revision,schema_version,state,updated_by) values ('blocked-editor','workspace-b','project-b',1,1,'{}'::jsonb,'user-a')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into project_render_jobs (id,workspace_id,project_id,editor_revision,editor_schema_version,editor_state,render_settings,request_fingerprint,idempotency_key,requested_by) values ('blocked-render','workspace-a','project-b',1,1,'{}'::jsonb,'{}'::jsonb,repeat('c',64),'blocked-render','user-a')" + ) + with pytest.raises(asyncpg.PostgresError): + await admin.execute( + "insert into project_generation_jobs (id,workspace_id,project_id,generation_job_id,attached_by) values ('blocked-project-job-admin','workspace-a','project-a','generation-job-b','user-a')" + ) + # The provider-runtime migration prevents a trusted worker recovery + # process from binding one opaque worker job to two tenant jobs. + await admin.execute( + "update generation_jobs set external_job_id = 'worker-job-a' where id = 'generation-job-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await admin.execute( + "update generation_jobs set external_job_id = 'worker-job-a' where id = 'generation-job-b'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update generation_jobs set output_asset_id = 'asset-b' where id = 'generation-job-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "insert into social_post_targets (id,social_post_id,social_account_id,provider) values ('blocked-target','post-a','account-b','youtube')" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update social_posts set workspace_id = 'workspace-b' where id = 'post-a'" + ) + # RLS USING makes an attempted write to tenant B affect zero rows. + assert ( + await tenant.execute("update social_jobs set status = 'failed' where id = 'job-b'") + == "UPDATE 0" + ) + assert ( + await tenant.execute("update projects set name = 'Blocked' where id = 'project-b'") + == "UPDATE 0" + ) + assert await tenant.execute("delete from projects where id = 'project-b'") == "DELETE 0" + assert ( + await tenant.execute( + "delete from project_generation_jobs where id = 'project-generation-b'" + ) + == "DELETE 0" + ) + assert ( + await tenant.fetchval( + "select count(*) from project_editor_states where project_id = 'project-b'" + ) + == 0 + ) + assert ( + await tenant.execute( + "update project_editor_states set revision = 2 where id = 'editor-b'" + ) + == "UPDATE 0" + ) + assert ( + await tenant.execute( + "update project_render_jobs set status = 'cancelled' where id = 'render-b'" + ) + == "UPDATE 0" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update project_editor_states set revision = 3 where id = 'editor-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update project_render_jobs set render_settings = '{\"format\":\"webm\"}'::jsonb where id = 'render-a'" + ) + with pytest.raises(asyncpg.PostgresError): + await tenant.execute( + "update project_render_jobs set status = 'completed', completed_at = now() where id = 'render-a'" + ) + await admin.execute( + "insert into users (id,subject,display_name) values ('viewer-a','test:viewer-a','Viewer'); " + "insert into workspace_memberships (id,workspace_id,user_id,role) values ('viewer-membership-a','workspace-a','viewer-a','viewer')" + ) + await _tenant_context(tenant, "workspace-a", "viewer-a") + assert await tenant.fetchval("select count(*) from project_editor_states") == 1 + assert await tenant.fetchval("select count(*) from project_render_jobs") == 1 + assert ( + await tenant.execute( + "update project_editor_states set revision = 2 where id = 'editor-a'" + ) + == "UPDATE 0" + ) + assert ( + await tenant.execute( + "update project_render_jobs set status = 'cancelled' where id = 'render-a'" + ) + == "UPDATE 0" + ) + finally: + await tenant.close() + await admin.close() diff --git a/tests/test_production_configuration.py b/tests/test_production_configuration.py new file mode 100644 index 0000000000000000000000000000000000000000..58990d3889b3421bf77c6c88afd315ff33e96243 --- /dev/null +++ b/tests/test_production_configuration.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from app.core.config import Settings + + +def production_settings(**overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "app_environment": "production", + "database_url": "postgresql+asyncpg://security@db.example/mediarouter", + "security_database_role": "mediarouter_security_service", + "cors_allowed_origins": "https://app.example.vercel.app", + # Social persistence has an additional tenant/worker role boundary. + # Disable it in the minimal production fixture; dedicated coverage + # below verifies the enabled contract. + "social_enabled": False, + } + values.update(overrides) + return Settings(**values) + + +def test_production_configuration_accepts_external_postgres_and_https_cors() -> None: + settings = production_settings( + cors_allowed_origins=( + "https://app.example.vercel.app, https://preview.example.vercel.app/" + ) + ) + + assert settings.allowed_cors_origins == ( + "https://app.example.vercel.app", + "https://preview.example.vercel.app", + ) + + +@pytest.mark.parametrize( + "database_url", + [ + "sqlite+aiosqlite:////app/data/mediarouter.db", + "postgresql+asyncpg://security@localhost/mediarouter", + "", + ], +) +def test_production_configuration_rejects_local_or_missing_database( + database_url: str, +) -> None: + with pytest.raises(ValidationError, match="external PostgreSQL"): + production_settings(database_url=database_url) + + +@pytest.mark.parametrize("origin", ["*", "https://*.vercel.app"]) +def test_production_configuration_rejects_wildcard_cors(origin: str) -> None: + with pytest.raises(ValidationError, match="CORS_ALLOWED_ORIGINS"): + production_settings(cors_allowed_origins=origin) + + +def test_production_configuration_rejects_automatic_migrations() -> None: + with pytest.raises(ValidationError, match="AUTO_MIGRATE must be false"): + production_settings(security_auto_migrate=True) + + +def test_social_enabled_requires_explicit_tenant_and_worker_boundaries() -> None: + with pytest.raises(ValidationError, match="SOCIAL_DATABASE_URL"): + production_settings(social_enabled=True) + + settings = production_settings( + social_enabled=True, + social_database_url="postgresql+asyncpg://tenant@db.example/mediarouter", + social_tenant_database_role="mediarouter_tenant", + social_worker_database_url=( + "postgresql+asyncpg://social_worker@db.example/mediarouter" + ), + social_worker_database_role="mediarouter_social_worker", + ) + assert settings.social_enabled is True + + +def test_development_keeps_the_existing_sqlite_contract() -> None: + settings = Settings(_env_file=None, cors_allowed_origins="http://localhost:3000/") + + assert settings.database_url.startswith("sqlite") + assert settings.allowed_cors_origins == ("http://localhost:3000",) diff --git a/tests/test_project_collaboration_authorization.py b/tests/test_project_collaboration_authorization.py new file mode 100644 index 0000000000000000000000000000000000000000..26ba75595317599a36338cf2d087ef2fc2c4fa87 --- /dev/null +++ b/tests/test_project_collaboration_authorization.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import unittest + +from app.projects.repositories.collaboration_repository import CollaborationRepository +from app.projects.services.collaboration_service import CollaborationService + + +class CollaborationAuthorizationTests(unittest.TestCase): + def test_project_collaborator_list_requires_project_workspace(self) -> None: + repository = CollaborationRepository(None) + service = CollaborationService(repository) + + with self.assertRaises(Exception): + service.list_project_collaborators("project-from-another-workspace") + + def test_project_collaborator_removal_requires_project_workspace(self) -> None: + repository = CollaborationRepository(None) + service = CollaborationService(repository) + + with self.assertRaises(Exception): + service.remove_project_collaborator("project-from-another-workspace", "user-id") + diff --git a/tests/test_project_models_imports.py b/tests/test_project_models_imports.py new file mode 100644 index 0000000000000000000000000000000000000000..aad35d0ed77726e350824da0a9217a07350ebce7 --- /dev/null +++ b/tests/test_project_models_imports.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +from importlib.machinery import PathFinder + +import pytest + + +def _load_project_collaboration_symbols(): + pytest.importorskip("sqlalchemy") + + import app.projects.models + import app.projects.models.collaboration + + from app.projects.models import ( + Project, + ProjectGenerationJob, + ProjectEditorState, + ProjectRenderJob, + ) + from app.projects.models.collaboration import ( + ApprovalRequest, + ApprovalWorkflow, + CollaborationActivity, + Invitation, + ProjectCollaborator, + ReviewComment, + Team, + TeamMember, + ) + + return { + "project_models": app.projects.models, + "project": Project, + "project_generation_job": ProjectGenerationJob, + "project_editor_state": ProjectEditorState, + "project_render_job": ProjectRenderJob, + "team": Team, + "team_member": TeamMember, + "invitation": Invitation, + "approval_workflow": ApprovalWorkflow, + "approval_request": ApprovalRequest, + "review_comment": ReviewComment, + "project_collaborator": ProjectCollaborator, + "collaboration_activity": CollaborationActivity, + } + + +_SYMBOLS = _load_project_collaboration_symbols() + + +def test_project_model_package_imports_are_resilient() -> None: + assert _SYMBOLS["project"].__name__ == "Project" + assert _SYMBOLS["project_generation_job"].__name__ == "ProjectGenerationJob" + assert _SYMBOLS["project_editor_state"].__name__ == "ProjectEditorState" + assert _SYMBOLS["project_render_job"].__name__ == "ProjectRenderJob" + + +def test_collaboration_submodule_exports_expected_models() -> None: + assert _SYMBOLS["team"].__name__ == "Team" + assert _SYMBOLS["team_member"].__name__ == "TeamMember" + assert _SYMBOLS["invitation"].__name__ == "Invitation" + assert _SYMBOLS["approval_workflow"].__name__ == "ApprovalWorkflow" + assert _SYMBOLS["approval_request"].__name__ == "ApprovalRequest" + assert _SYMBOLS["review_comment"].__name__ == "ReviewComment" + assert _SYMBOLS["project_collaborator"].__name__ == "ProjectCollaborator" + assert _SYMBOLS["collaboration_activity"].__name__ == "CollaborationActivity" + + +def test_production_style_startup_imports_are_resolvable() -> None: + assert _SYMBOLS["project_models"].Project is not None + assert _SYMBOLS["project_models"].collaboration.Team is not None + assert _SYMBOLS["project_models"].Project is not None + + +def test_models_module_resolution_avoids_collision() -> None: + parent_paths = [str(path) for path in __import__("app.projects", fromlist=[""]).__path__] + package_spec = PathFinder.find_spec("app.projects.models", parent_paths) + assert package_spec is not None + assert package_spec.origin.endswith("__init__.py") + assert package_spec.submodule_search_locations + + collaboration_spec = PathFinder.find_spec( + "app.projects.models.collaboration", + list(package_spec.submodule_search_locations), + ) + assert collaboration_spec is not None + assert collaboration_spec.origin.endswith("collaboration.py") + + +def test_main_app_symbol_is_registered() -> None: + pytest.importorskip("sqlalchemy") + import main + + assert main.app is not None diff --git a/tests/test_projects_foundation.py b/tests/test_projects_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..39ba7556e2aa297ac1da1d3f4abc676d0cb02c62 --- /dev/null +++ b/tests/test_projects_foundation.py @@ -0,0 +1,713 @@ +from __future__ import annotations + +from pathlib import Path +from uuid import uuid4 + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.generation.models import GenerationJob, GenerationRequest +from app.projects.errors import ( + ProjectAlreadyArchivedError, + ProjectAssetConflictError, + ProjectAssetNotFoundError, + ProjectJobNotFoundError, + ProjectNotFoundError, + ProjectThumbnailInvalidError, +) +from app.projects.schemas import ProjectCreate, ProjectStatus, ProjectUpdate +from app.security.models import AuditEvent +from app.security.policy import ScopePolicy +from app.security.schemas import APIKeyCreate +from app.security.service import APIKeyService +from main import create_app + + +def project_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": True, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "test-only-encryption-material", + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + "auth_default_requests_per_minute": 10_000, + } + values.update(overrides) + return Settings(**values) + + +async def _project_context(container: object, name: str, scopes: list[str]): + record, secret = await container.api_keys.create( # type: ignore[attr-defined] + APIKeyCreate(name=name, environment="test", role=None, scopes=scopes), + created_by="tests", + ) + return record, secret, await container.api_keys.authenticate(secret) # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_project_service_lifecycle_isolation_pagination_thumbnail_and_audit( + tmp_path: Path, +) -> None: + container = build_container(project_settings(tmp_path)) + await container.security_database.initialize() + scopes = [ + "projects:read", + "projects:create", + "projects:update", + "projects:delete", + ] + try: + key_a, _, actor_a = await _project_context(container, "Workspace A", scopes) + _, _, actor_b = await _project_context(container, "Workspace B", scopes) + assert actor_a.workspace_id != actor_b.workspace_id + + output_id = str(uuid4()) + output = container.settings.output_dir / output_id + output.mkdir(parents=True) + thumbnail_path = output / "thumbnail.png" + thumbnail_path.write_bytes(b"canonical thumbnail") + thumbnail = await container.assets.register_output( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + request_id=output_id, + path=thumbnail_path, + mime_type="image/png", + ) + + first = await container.projects.create( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate( + name=" Podcast Episode 41 ", + description="Primary project", + thumbnail_asset_id=thumbnail.id, + metadata={"aspect_ratio": "16:9"}, + ), + ) + assert first.name == "Podcast Episode 41" + assert first.status is ProjectStatus.ACTIVE + assert first.workspace_id == actor_a.workspace_id + assert first.created_by == actor_a.user_id + assert first.thumbnail_asset_id == thumbnail.id + + with pytest.raises(ProjectNotFoundError): + await container.projects.get( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + project_id=first.id, + ) + with pytest.raises(ProjectNotFoundError): + await container.projects.update( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + api_key_id=actor_b.api_key_id, + request_id=str(uuid4()), + project_id=first.id, + payload=ProjectUpdate(name="IDOR update"), + ) + with pytest.raises(ProjectNotFoundError): + await container.projects.delete( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + api_key_id=actor_b.api_key_id, + request_id=str(uuid4()), + project_id=first.id, + ) + with pytest.raises(ProjectThumbnailInvalidError): + await container.projects.create( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + api_key_id=actor_b.api_key_id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Foreign thumbnail", thumbnail_asset_id=thumbnail.id), + ) + + second = await container.projects.create( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Second project"), + ) + third = await container.projects.create( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Third project"), + ) + page_one = await container.projects.list( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + status=ProjectStatus.ACTIVE, + search=None, + limit=2, + cursor=None, + ) + assert len(page_one.items) == 2 + assert page_one.next_cursor + page_two = await container.projects.list( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + status=ProjectStatus.ACTIVE, + search=None, + limit=2, + cursor=page_one.next_cursor, + ) + assert len(page_two.items) == 1 + assert {item.id for item in page_one.items + page_two.items} == { + first.id, + second.id, + third.id, + } + searched = await container.projects.list( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + status=ProjectStatus.ACTIVE, + search="podcast", + limit=10, + cursor=None, + ) + assert [item.id for item in searched.items] == [first.id] + + updated = await container.projects.update( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=first.id, + payload=ProjectUpdate(description="Updated description"), + ) + assert updated.description == "Updated description" + archived = await container.projects.update( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=first.id, + payload=ProjectUpdate(status=ProjectStatus.ARCHIVED), + ) + assert archived.status is ProjectStatus.ARCHIVED + assert archived.archived_at is not None + with pytest.raises(ProjectAlreadyArchivedError): + await container.projects.update( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=first.id, + payload=ProjectUpdate(name="Blocked"), + ) + + await container.projects.delete( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=second.id, + ) + async with container.security_database.tenant_session( + workspace_id=str(actor_a.workspace_id), user_id=str(actor_a.user_id) + ) as session: + events = list( + ( + await session.scalars( + select(AuditEvent).where( + AuditEvent.workspace_id == actor_a.workspace_id, + AuditEvent.entity_type == "project", + ) + ) + ).all() + ) + assert {event.event_type for event in events} >= { + "project.created", + "project.updated", + "project.archived", + "project.deleted", + } + assert all("aspect_ratio" not in event.metadata_json for event in events) + finally: + await container.security_database.close() + + +@pytest.mark.asyncio +async def test_project_resource_service_ownership_lifecycle_and_audit( + tmp_path: Path, +) -> None: + container = build_container(project_settings(tmp_path)) + await container.security_database.initialize() + scopes = ["projects:read", "projects:create", "projects:update"] + try: + key_a, _, actor_a = await _project_context(container, "Resource A", scopes) + key_b, _, actor_b = await _project_context(container, "Resource B", scopes) + project_a = await container.projects.create( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Project A"), + ) + second_a = await container.projects.create( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Second A"), + ) + project_b = await container.projects.create( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + api_key_id=key_b.id, + request_id=str(uuid4()), + payload=ProjectCreate(name="Project B"), + ) + + async def canonical_asset(actor: object, name: str): + request_id = str(uuid4()) + directory = container.settings.output_dir / request_id + directory.mkdir(parents=True) + path = directory / name + path.write_bytes(name.encode()) + return await container.assets.register_output( + workspace_id=str(actor.workspace_id), # type: ignore[attr-defined] + user_id=str(actor.user_id), # type: ignore[attr-defined] + request_id=request_id, + path=path, + mime_type="video/mp4", + ) + + asset_a = await canonical_asset(actor_a, "asset-a.mp4") + asset_b = await canonical_asset(actor_b, "asset-b.mp4") + attached = await container.projects.attach_asset( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + asset_id=asset_a.id, + ) + assert attached.project_id == project_a.id + assert [ + item.id + for item in ( + await container.projects.list_assets( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + project_id=project_a.id, + ) + ).items + ] == [asset_a.id] + with pytest.raises(ProjectAssetConflictError): + await container.projects.attach_asset( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=second_a.id, + asset_id=asset_a.id, + ) + with pytest.raises(ProjectAssetNotFoundError): + await container.projects.attach_asset( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + asset_id=asset_b.id, + ) + with pytest.raises(ProjectNotFoundError): + await container.projects.list_assets( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + project_id=project_a.id, + ) + + generation_request_a = GenerationRequest( + workspace_id=str(actor_a.workspace_id), + created_by_user_id=str(actor_a.user_id), + provider="test", + model_id="test/model", + modality="video", + spec_json={}, + request_fingerprint="a" * 64, + idempotency_key=str(uuid4()), + status="queued", + ) + generation_request_b = GenerationRequest( + workspace_id=str(actor_b.workspace_id), + created_by_user_id=str(actor_b.user_id), + provider="test", + model_id="test/model", + modality="video", + spec_json={}, + request_fingerprint="b" * 64, + idempotency_key=str(uuid4()), + status="queued", + ) + async with container.security_database.session() as session: + session.add_all([generation_request_a, generation_request_b]) + await session.flush() + job_a = GenerationJob( + generation_request_id=generation_request_a.id, + workspace_id=str(actor_a.workspace_id), + provider="test", + status="queued", + ) + job_b = GenerationJob( + generation_request_id=generation_request_b.id, + workspace_id=str(actor_b.workspace_id), + provider="test", + status="queued", + ) + session.add_all([job_a, job_b]) + await session.commit() + + linked_job = await container.projects.attach_generation_job( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + generation_job_id=job_a.id, + ) + assert linked_job.id == job_a.id + assert [ + item.id + for item in ( + await container.projects.list_generation_jobs( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + project_id=project_a.id, + ) + ).items + ] == [job_a.id] + with pytest.raises(ProjectJobNotFoundError): + await container.projects.attach_generation_job( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + generation_job_id=job_b.id, + ) + with pytest.raises(ProjectNotFoundError): + await container.projects.list_generation_jobs( + workspace_id=str(actor_b.workspace_id), + user_id=str(actor_b.user_id), + project_id=project_a.id, + ) + + await container.projects.detach_generation_job( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + generation_job_id=job_a.id, + ) + await container.projects.detach_asset( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + asset_id=asset_a.id, + ) + await container.projects.delete( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + ) + with pytest.raises(ProjectAlreadyArchivedError): + await container.projects.attach_asset( + workspace_id=str(actor_a.workspace_id), + user_id=str(actor_a.user_id), + api_key_id=key_a.id, + request_id=str(uuid4()), + project_id=project_a.id, + asset_id=asset_a.id, + ) + + async with container.security_database.tenant_session( + workspace_id=str(actor_a.workspace_id), user_id=str(actor_a.user_id) + ) as session: + events = list( + ( + await session.scalars( + select(AuditEvent).where( + AuditEvent.workspace_id == actor_a.workspace_id, + AuditEvent.entity_id == project_a.id, + ) + ) + ).all() + ) + assert {event.event_type for event in events} >= { + "project.asset_attached", + "project.asset_detached", + "project.job_attached", + "project.job_detached", + } + assert all( + set(event.metadata_json) <= {"resource_id", "disposition", "has_thumbnail"} + for event in events + ) + assert project_b.workspace_id == actor_b.workspace_id + finally: + await container.security_database.close() + + +def _create_key(client: TestClient, admin_headers: dict[str, str], scopes: list[str]) -> str: + response = client.post( + "/v1/api-keys", + headers=admin_headers, + json={ + "name": f"Project key {uuid4()}", + "environment": "test", + "role": None, + "scopes": scopes, + }, + ) + assert response.status_code == 201 + return response.json()["api_key"] + + +def test_project_http_permissions_openapi_validation_and_archive(tmp_path: Path) -> None: + material = APIKeyService.generate_material("test") + settings = project_settings( + tmp_path, + auth_bootstrap_key_hash=material.key_hash, + auth_bootstrap_key_prefix=material.key_prefix, + auth_bootstrap_environment="test", + ) + admin_headers = {"Authorization": f"Bearer {material.api_key}"} + with TestClient(create_app(settings)) as client: + schema = client.get("/openapi.json").json() + assert set(schema["paths"]["/v1/projects"]) >= {"get", "post"} + assert set(schema["paths"]["/v1/projects/{project_id}"]) >= { + "get", + "patch", + "delete", + } + assert set(schema["paths"]["/v1/projects/{project_id}/assets"]) >= {"get", "post"} + assert "delete" in schema["paths"]["/v1/projects/{project_id}/assets/{asset_id}"] + assert set(schema["paths"]["/v1/projects/{project_id}/jobs"]) >= {"get", "post"} + assert "delete" in schema["paths"]["/v1/projects/{project_id}/jobs/{job_id}"] + assert set(schema["paths"]["/v1/projects/{project_id}/editor"]) >= {"get", "put"} + assert set(schema["paths"]["/v1/projects/{project_id}/renders"]) >= {"get", "post"} + assert "get" in schema["paths"]["/v1/projects/{project_id}/renders/{render_id}"] + assert "post" in schema["paths"]["/v1/projects/{project_id}/renders/{render_id}/cancel"] + assert "ProjectCreate" in schema["components"]["schemas"] + assert "ProjectUpdate" in schema["components"]["schemas"] + assert "ProjectResponse" in schema["components"]["schemas"] + assert "ProjectListResponse" in schema["components"]["schemas"] + assert "ProjectAssetResponse" in schema["components"]["schemas"] + assert "ProjectGenerationJobResponse" in schema["components"]["schemas"] + assert "EditorSaveRequest" in schema["components"]["schemas"] + assert "EditorStateResponse" in schema["components"]["schemas"] + assert "ProjectRenderCreate" in schema["components"]["schemas"] + assert "ProjectRenderResponse" in schema["components"]["schemas"] + + assert client.get("/v1/projects").status_code == 401 + read_secret = _create_key(client, admin_headers, ["projects:read"]) + create_secret = _create_key(client, admin_headers, ["projects:create"]) + update_secret = _create_key(client, admin_headers, ["projects:update"]) + delete_secret = _create_key(client, admin_headers, ["projects:delete"]) + jobs_secret = _create_key(client, admin_headers, ["jobs:create", "jobs:cancel"]) + render_secret = _create_key( + client, + admin_headers, + ["projects:update", "jobs:create", "jobs:cancel"], + ) + read_headers = {"Authorization": f"Bearer {read_secret}"} + create_headers = {"Authorization": f"Bearer {create_secret}"} + update_headers = {"Authorization": f"Bearer {update_secret}"} + delete_headers = {"Authorization": f"Bearer {delete_secret}"} + jobs_headers = {"Authorization": f"Bearer {jobs_secret}"} + render_headers = {"Authorization": f"Bearer {render_secret}"} + + assert client.get("/v1/projects", headers=read_headers).status_code == 200 + assert ( + client.post("/v1/projects", headers=read_headers, json={"name": "No"}).status_code + == 403 + ) + created = client.post( + "/v1/projects", + headers=create_headers, + json={"name": " HTTP Project ", "metadata": {"source": "test"}}, + ) + assert created.status_code == 201 + project = created.json() + assert project["name"] == "HTTP Project" + assert "workspace_id" not in created.request.content.decode() + project_id = project["id"] + assert client.get(f"/v1/projects/{project_id}", headers=create_headers).status_code == 403 + assert client.get(f"/v1/projects/{project_id}", headers=read_headers).status_code == 200 + assert ( + client.patch( + f"/v1/projects/{project_id}", + headers=update_headers, + json={"description": "Changed"}, + ).status_code + == 200 + ) + empty_editor = { + "schemaVersion": 1, + "projectId": project_id, + "timeline": { + "timeUnit": "milliseconds", + "tracks": [], + "transitions": [], + "markers": [], + }, + "renderSettings": { + "format": "mp4", + "width": 1280, + "height": 720, + "frameRate": 30, + }, + } + saved_editor = client.put( + f"/v1/projects/{project_id}/editor", + headers=update_headers, + json={"expected_revision": 0, "schema_version": 1, "state": empty_editor}, + ) + assert saved_editor.status_code == 200 + assert saved_editor.json()["revision"] == 1 + assert ( + client.get(f"/v1/projects/{project_id}/editor", headers=read_headers).status_code == 200 + ) + render_payload = { + "editor_revision": 1, + "output_format": "mp4", + "width": 1280, + "height": 720, + "frame_rate": 30, + "quality": "standard", + "preset": "balanced", + } + render_path = f"/v1/projects/{project_id}/renders" + assert ( + client.post( + render_path, + headers={**update_headers, "Idempotency-Key": "missing-jobs-scope"}, + json=render_payload, + ).status_code + == 403 + ) + assert ( + client.post( + render_path, + headers={**jobs_headers, "Idempotency-Key": "missing-project-scope"}, + json=render_payload, + ).status_code + == 403 + ) + render_rejected = client.post( + render_path, + headers={**render_headers, "Idempotency-Key": "empty-editor"}, + json=render_payload, + ) + assert render_rejected.status_code == 422 + assert render_rejected.json()["error"]["code"] == "PROJECT_RENDER_INVALID" + cancel_path = f"{render_path}/{uuid4()}/cancel" + assert client.post(cancel_path, headers=update_headers).status_code == 403 + assert client.post(cancel_path, headers=render_headers).status_code == 404 + assert ( + client.delete(f"/v1/projects/{project_id}", headers=update_headers).status_code == 403 + ) + assert ( + client.delete(f"/v1/projects/{project_id}", headers=delete_headers).status_code == 204 + ) + archived = client.get(f"/v1/projects/{project_id}", headers=read_headers) + assert archived.status_code == 200 + assert archived.json()["status"] == "archived" + assert archived.json()["archived_at"] is not None + + invalid_name = client.post("/v1/projects", headers=admin_headers, json={"name": " "}) + assert invalid_name.status_code == 422 + extra_system_field = client.post( + "/v1/projects", + headers=admin_headers, + json={"name": "Unsafe", "workspace_id": str(uuid4())}, + ) + assert extra_system_field.status_code == 422 + oversized_metadata = client.post( + "/v1/projects", + headers=admin_headers, + json={"name": "Large", "metadata": {"value": "x" * 20_000}}, + ) + assert oversized_metadata.status_code == 422 + invalid_cursor = client.get("/v1/projects?cursor=not-a-cursor", headers=admin_headers) + assert invalid_cursor.status_code == 422 + assert invalid_cursor.json()["error"]["code"] == "PROJECT_INVALID_CURSOR" + + +def test_project_migration_is_additive_and_contains_security_guards() -> None: + migration = ( + ( + Path(__file__).resolve().parents[1] + / "app/projects/migrations/0001_projects_foundation.sql" + ) + .read_text(encoding="utf-8") + .casefold() + ) + for expected in ( + "create table if not exists projects", + "check (status in ('active', 'archived'))", + "ix_projects_workspace_status", + "enable row level security", + "force row level security", + "create policy projects_select", + "create policy projects_insert", + "create policy projects_update", + "create policy projects_delete", + "mediarouter_assert_project_ownership", + "audit_events", + ): + assert expected in migration + assert "drop table" not in migration + + resources = ( + (Path(__file__).resolve().parents[1] / "app/projects/migrations/0002_project_resources.sql") + .read_text(encoding="utf-8") + .casefold() + ) + for expected in ( + "alter table media_assets add column if not exists project_id", + "fk_media_assets_project", + "mediarouter_assert_media_asset_project_workspace", + "create table if not exists project_generation_jobs", + "uq_project_generation_job", + "mediarouter_assert_project_generation_job_workspace", + "enable row level security", + "force row level security", + "project_generation_jobs_select", + "project_generation_jobs_insert", + "project_generation_jobs_delete", + ): + assert expected in resources + assert "drop table" not in resources + + +def test_project_resource_scope_mapping_uses_project_update() -> None: + assert ScopePolicy._project_scope("/v1/projects", "POST") == "projects:create" + assert ScopePolicy._project_scope("/v1/projects/id", "DELETE") == "projects:delete" + assert ScopePolicy._project_scope("/v1/projects/id/assets", "GET") == "projects:read" + assert ScopePolicy._project_scope("/v1/projects/id/assets", "POST") == "projects:update" + assert ScopePolicy._project_scope("/v1/projects/id/assets/asset", "DELETE") == "projects:update" + assert ScopePolicy._project_scope("/v1/projects/id/jobs/job", "DELETE") == "projects:update" + assert ScopePolicy._project_scope("/v1/projects/id/editor", "PUT") == "projects:update" + assert ScopePolicy._project_scope("/v1/projects/id/renders", "POST") == "projects:update" diff --git a/tests/test_publishing_operations_phase9_static.py b/tests/test_publishing_operations_phase9_static.py new file mode 100644 index 0000000000000000000000000000000000000000..e55c9e7eeec7c37e3a27e453b238be17d52b5941 --- /dev/null +++ b/tests/test_publishing_operations_phase9_static.py @@ -0,0 +1,52 @@ +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] + + +def test_phase9_migration_is_additive_and_forces_batch_rls() -> None: + migration = ( + ROOT / "app/social/migrations/0009_publishing_operations.sql" + ).read_text() + normalized = migration.lower() + assert "drop table" not in normalized + assert "social_posts add column if not exists revision" in migration + assert "social_schedules add column if not exists revision" in migration + assert "create table if not exists social_publishing_batches" in migration + assert "create table if not exists social_publishing_batch_items" in migration + assert normalized.count("force row level security") >= 2 + assert "current_setting('app.workspace_id'" in migration + + +def test_phase9_uses_existing_posts_schedules_jobs_and_worker() -> None: + service = ( + ROOT / "app/social/services/publishing_operations_service.py" + ).read_text() + scheduler = (ROOT / "app/social/workers/scheduler.py").read_text() + assert "PostRepository" in service + assert "JobRepository" in service + assert "SchedulingService" in service + assert "process_batch_items" in scheduler + assert "SocialJob(" not in service + + +def test_phase9_transports_expose_narrow_typed_operations() -> None: + api = (ROOT / "app/api/social.py").read_text() + mcp = (ROOT / "app/mcp/tools/social.py").read_text() + for route in ( + '"/drafts"', + '"/posts/{post_id}/duplicate"', + '"/posts/{post_id}/reschedule"', + '"/calendar"', + '"/queue"', + '"/bulk"', + ): + assert route in api + for tool in ( + "publishing.list_queue", + "publishing.list_calendar", + "publishing.reschedule", + "publishing.duplicate", + "publishing.bulk", + ): + assert tool in mcp diff --git a/tests/test_python310_compat.py b/tests/test_python310_compat.py new file mode 100644 index 0000000000000000000000000000000000000000..af35423c77c11697c27ca24d170176ebc64d8b2c --- /dev/null +++ b/tests/test_python310_compat.py @@ -0,0 +1,23 @@ +"""Regression tests for the Python 3.10 production runtime contract.""" + +from app.core.enums import StrEnum +from app.social.domain.enums import JobStatus, Provider +from app.social.schemas.tiktok import TikTokPrivacyLevel +from app.social.schemas.youtube import YouTubePrivacyStatus + + +def test_string_enums_do_not_require_python_311_stdlib() -> None: + class Example(StrEnum): + VALUE = "value" + + assert Example.VALUE == "value" + assert str(Example.VALUE) == "value" + assert f"{Example.VALUE}" == "value" + assert Example.VALUE.value == "value" + + +def test_social_schema_enums_keep_wire_values() -> None: + assert str(Provider.TIKTOK) == "tiktok" + assert str(JobStatus.PUBLISHED) == "published" + assert str(TikTokPrivacyLevel.SELF_ONLY) == "SELF_ONLY" + assert str(YouTubePrivacyStatus.PRIVATE) == "private" diff --git a/tests/test_security_authorization_idor.py b/tests/test_security_authorization_idor.py new file mode 100644 index 0000000000000000000000000000000000000000..1c3b355e1be9abb99661259dcda648dfae0f90e2 --- /dev/null +++ b/tests/test_security_authorization_idor.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +from uuid import uuid4 +import unittest + +from fastapi.testclient import TestClient + +from app.container import build_container +from app.core.config import Settings +from app.security.schemas import APIKeyCreate +from app.security.service import APIKeyService +from main import create_app + + +def _security_settings(tmp_path): + return Settings( + _env_file=None, + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + auth_enabled=True, + auth_last_used_update_seconds=0, + cleanup_interval_seconds=3600, + whisper_model="tiny", + max_workers=1, + ) + + +async def _create_app(tmp_path): + settings = _security_settings(tmp_path) + container = build_container(settings) + await container.security_database.initialize() + application = create_app(settings) + return application, container + + +async def _project_context(container, name, scopes): + record, secret = await container.api_keys.create( + APIKeyCreate(name=name, environment="test", role=None, scopes=scopes), + created_by="tests", + ) + return record, secret, await container.api_keys.authenticate(secret) + + +class ApprovalAuthorizationTests(unittest.TestCase): + def test_list_approval_requests_requires_workflow_workspace_membership(self) -> None: + import asyncio + application, _ = asyncio.run(_create_app(self)) + with TestClient(application) as client: + unauthorized = client.get( + "/v1/projects/workspace/workflows/other-workflow/requests", + headers={"Authorization": "Bearer invalid"}, + ) + self.assertEqual(unauthorized.status_code, 401) + + def test_approve_and_reject_endpoints_authorize_by_request_workspace(self) -> None: + import asyncio + application, _ = asyncio.run(_create_app(self)) + with TestClient(application) as client: + missing = client.post( + "/v1/projects/workspace/requests/missing/approve", + headers={"Authorization": "Bearer invalid"}, + ) + self.assertEqual(missing.status_code, 404) + + def test_approval_workflow_requests_are_scoped_to_owner_workspace(self) -> None: + import asyncio + application, container = asyncio.run(_create_app(self)) + with TestClient(application) as client: + owner_key, _, actor_a = asyncio.run(_project_context(container, "Workspace A", [ + "projects:read", + "projects:create", + "approvals:create", + "approvals:read", + "approvals:review", + ])) + _, _, actor_b = asyncio.run(_project_context(container, "Workspace B", [ + "projects:read", + "projects:create", + "approvals:create", + "approvals:read", + "approvals:review", + ])) + + project = client.post( + "/v1/projects", + json={"name": "Approval Project"}, + headers={"Authorization": f"Bearer {owner_key}"}, + ).json() + + workflow = client.post( + "/v1/projects/workspace/workflows", + json={"project_id": project["id"], "name": "Review"}, + headers={"Authorization": f"Bearer {owner_key}"}, + ).json() + + approval_request = client.post( + f"/v1/projects/workspace/workflows/{workflow['id']}/requests", + json={"project_id": project["id"]}, + headers={"Authorization": f"Bearer {owner_key}"}, + ).json() + + self.assertEqual( + client.get( + f"/v1/projects/workspace/workflows/{workflow['id']}/requests", + headers={"Authorization": f"Bearer {actor_b.api_key_id}"}, + ).status_code, + 404, + ) + self.assertEqual( + client.post( + f"/v1/projects/workspace/requests/{approval_request['id']}/approve", + headers={"Authorization": f"Bearer {actor_b.api_key_id}"}, + ).status_code, + 404, + ) + self.assertEqual( + client.post( + f"/v1/projects/workspace/requests/{approval_request['id']}/reject", + headers={"Authorization": f"Bearer {actor_b.api_key_id}"}, + ).status_code, + 404, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_security_role_enforcement.py b/tests/test_security_role_enforcement.py new file mode 100644 index 0000000000000000000000000000000000000000..8cc6cabee341054aa1d60ac7592640196ccf8082 --- /dev/null +++ b/tests/test_security_role_enforcement.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import unittest + +from app.security.context import AuthContext +from app.security.errors import ForbiddenError +from app.security.service import APIKeyService + + +class SecurityRoleEnforcementTests(unittest.TestCase): + def test_viewer_scope_is_enforced_by_authoritative_membership_role(self) -> None: + context = AuthContext( + api_key_id="key", + key_name="Viewer Key", + key_prefix="mp_test_viewer", + environment="test", + role="viewer", + scopes=frozenset({"projects:read", "projects:write", "jobs:create"}), + requests_per_minute=60, + concurrent_jobs=1, + uploads_per_hour=1, + processing_bytes_per_day=1024, + ) + + APIKeyService.authorize(context, "projects:read") + APIKeyService.authorize(context, "jobs:read") + + with self.assertRaises(ForbiddenError): + APIKeyService.authorize(context, "projects:write") + with self.assertRaises(ForbiddenError): + APIKeyService.authorize(context, "jobs:create") + + def test_explicit_scopes_still_bypass_role_when_authorized(self) -> None: + context = AuthContext( + api_key_id="key", + key_name="Developer Key", + key_prefix="mp_test_dev", + environment="test", + role="developer", + scopes=frozenset({"projects:write", "jobs:create"}), + requests_per_minute=60, + concurrent_jobs=1, + uploads_per_hour=1, + processing_bytes_per_day=1024, + ) + + APIKeyService.authorize(context, "projects:write") + APIKeyService.authorize(context, "jobs:create") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_social_foundation.py b/tests/test_social_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..f129d27bdf8a9e083f8073bfc2e99aad23b064ce --- /dev/null +++ b/tests/test_social_foundation.py @@ -0,0 +1,526 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from pathlib import Path +from uuid import uuid4 + +import pytest +from pydantic import ValidationError +from sqlalchemy import select +from starlette.requests import Request + +from app.container import build_container +from app.core.config import Settings +from app.security.policy import ScopePolicy +from app.social.database import SocialDatabase +from app.social.domain.enums import JobStatus +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialIdempotencyConflictError, + SocialJobNotFoundError, + SocialMediaInvalidError, + SocialOAuthStateError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialReauthRequiredError, + SocialTransitionError, +) +from app.social.domain.retry import classify_retry +from app.social.domain.state_machine import validate_transition +from app.social.models import OAuthState, SocialAccount, SocialAccountToken, SocialJob, SocialMediaAsset +from app.social.schemas.posts import SocialPostCreate +from app.social.schemas.scheduling import SocialScheduleCreate + + +def social_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +@pytest.fixture +async def social_container(tmp_path: Path): + container = build_container(social_settings(tmp_path)) + await container.social.initialize() + try: + yield container + finally: + await container.social.close() + await container.security_database.close() + + +async def connected_account(container, workspace_id: str, provider: str = "youtube") -> SocialAccount: + return await container.social.accounts.repository.create( + SocialAccount( + workspace_id=workspace_id, + provider=provider, + account_type="channel", + external_account_id=f"external-{workspace_id}-{provider}", + display_name="Test channel", + status="connected", + ) + ) + + +async def registered_output_asset(container, workspace_id: str) -> str: + request_id = str(uuid4()) + output = container.settings.output_dir / request_id + output.mkdir(parents=True, exist_ok=True) + (output / "video.mp4").write_bytes(b"test") + record = await container.social.media_assets.repository.create( + SocialMediaAsset( + workspace_id=workspace_id, + request_id=request_id, + filename="video.mp4", + mime_type="video/mp4", + file_size=4, + ) + ) + return record.id + + +def post_payload(account_id: str, *, title: str = "Example") -> SocialPostCreate: + return SocialPostCreate.model_validate( + { + "media_asset_id": "asset-owned-by-workspace", + "publish_mode": "draft", + "targets": [ + { + "social_account_id": account_id, + "caption": {"title": title, "description": "Description"}, + "youtube": { + "title": title, + "description": "Description", + "privacy_status": "private", + "made_for_kids": False, + }, + } + ], + } + ) + + +async def test_provider_registry_and_capability_matrix(social_container) -> None: + providers = social_container.social.accounts.list_providers() + assert {item.provider.value for item in providers} == { + "youtube", + "facebook", + "instagram", + "tiktok", + "x", + "linkedin", + "telegram", + "whatsapp", + } + assert all(item.available is False for item in providers) + youtube = next(item for item in providers if item.provider.value == "youtube") + assert youtube.capabilities.implementation_status == "implemented" + assert youtube.capabilities.video_upload + assert youtube.capabilities.video_status + assert youtube.capabilities.channel_metadata + assert youtube.configured is False + assert next(item for item in providers if item.provider.value == "telegram").connection_strategy.value == "token_bot" + assert next(item for item in providers if item.provider.value == "whatsapp").connection_strategy.value == "business_api" + + +def test_job_state_machine_and_retry_classification() -> None: + assert validate_transition(JobStatus.DRAFT, JobStatus.SCHEDULED) == JobStatus.SCHEDULED + assert validate_transition(JobStatus.RETRYING, JobStatus.PROCESSING) == JobStatus.PROCESSING + with pytest.raises(SocialTransitionError): + validate_transition(JobStatus.PUBLISHED, JobStatus.QUEUED) + assert classify_retry(status_code=429, attempt=3).retryable + assert classify_retry(status_code=401, attempt=1).refresh_token_first + assert not classify_retry(status_code=403, attempt=1).retryable + assert not classify_retry(status_code=400, attempt=1).retryable + + +def test_only_the_exact_provider_callback_route_is_public() -> None: + policy = ScopePolicy() + + def request_for(path: str) -> Request: + return Request( + {"type": "http", "method": "GET", "path": path, "headers": []} + ) + + assert policy.is_public( + request_for("/v1/social/accounts/youtube/callback") + ) + assert not policy.is_public( + request_for("/v1/social/accounts/youtube/untrusted/callback") + ) + assert not policy.is_public( + Request( + { + "type": "http", + "method": "POST", + "path": "/v1/social/accounts/youtube/callback", + "headers": [], + } + ) + ) + + +async def test_social_auto_migrate_false_does_not_mutate_schema(tmp_path: Path) -> None: + settings = social_settings(tmp_path) + settings.social_auto_migrate = False + database = SocialDatabase(settings) + try: + await database.initialize() + assert not await database.schema_ready() + assert "social_accounts" in await database.missing_tables() + finally: + await database.close() + + +async def test_oauth_state_is_random_expiring_single_use_and_tenant_bound( + social_container, +) -> None: + state = await social_container.social.oauth.states.create( + provider="youtube", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri="https://api.example/v1/social/accounts/youtube/callback", + ) + assert len(state.state) >= 32 + consumed = await social_container.social.oauth.states.consume( + state=state.state, provider="youtube" + ) + assert consumed.workspace_id == "workspace-a" + with pytest.raises(SocialOAuthStateError): + await social_container.social.oauth.states.consume( + state=state.state, provider="youtube" + ) + + expired = OAuthState( + state="expired-state", + provider="youtube", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri="https://api.example/callback", + expires_at=datetime.now(timezone.utc) - timedelta(seconds=1), + ) + async with social_container.social.database.session() as session: + session.add(expired) + await session.commit() + with pytest.raises(SocialOAuthStateError): + await social_container.social.oauth.states.consume( + state="expired-state", provider="youtube" + ) + + wrong_provider = await social_container.social.oauth.states.create( + provider="youtube", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri="https://api.example/v1/social/accounts/youtube/callback", + ) + with pytest.raises(SocialOAuthStateError): + await social_container.social.oauth.states.consume( + state=wrong_provider.state, provider="linkedin" + ) + assert ( + await social_container.social.oauth.states.consume( + state=wrong_provider.state, provider="youtube" + ) + ).workspace_id == "workspace-a" + with pytest.raises(SocialOAuthStateError): + await social_container.social.oauth.states.consume( + state="not-a-valid-state", provider="youtube" + ) + + +async def test_oauth_redirect_uri_is_backend_owned(social_container) -> None: + oauth = social_container.social.oauth + with pytest.raises(SocialPermissionDeniedError): + oauth._redirect_uri( + "youtube", "https://attacker.example/v1/social/accounts/youtube/callback" + ) + social_container.settings.social_oauth_redirect_base_url = "https://api.example" + expected = "https://api.example/v1/social/accounts/youtube/callback" + assert oauth._redirect_uri("youtube", None) == expected + with pytest.raises(SocialPermissionDeniedError): + oauth._redirect_uri( + "youtube", "https://attacker.example/v1/social/accounts/youtube/callback" + ) + + +async def test_token_service_encrypts_and_never_returns_storage_metadata( + social_container, +) -> None: + account = await connected_account(social_container, "workspace-token") + secret = "provider-access-token-that-must-not-leak" + await social_container.social.accounts.tokens.store( + "workspace-token", + account.id, + {"access_token": secret, "refresh_token": "refresh-secret"}, + scopes=["upload"], + ) + async with social_container.social.database.session() as session: + row = await session.scalar( + select(SocialAccountToken).where( + SocialAccountToken.social_account_id == account.id + ) + ) + assert row is not None + assert secret not in (row.encrypted_payload or "") + assert await social_container.social.accounts.tokens.retrieve( + "workspace-token", account.id + ) == { + "access_token": secret, + "refresh_token": "refresh-secret", + } + view = await social_container.social.accounts.get("workspace-token", account.id) + assert "token" not in view.model_dump_json().lower() + with pytest.raises(SocialReauthRequiredError): + await social_container.social.accounts.tokens.retrieve("workspace-other", account.id) + + +async def test_workspace_ownership_idempotency_and_multi_target_foundation( + social_container, +) -> None: + account = await connected_account(social_container, "workspace-a") + payload = post_payload(account.id) + first = await social_container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="create-post-key", + ) + replay = await social_container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="create-post-key", + ) + assert replay.id == first.id + assert len(first.targets) == 1 + + with pytest.raises(SocialIdempotencyConflictError): + await social_container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=post_payload(account.id, title="Different"), + idempotency_key="create-post-key", + ) + with pytest.raises(SocialAccountNotFoundError): + await social_container.social.publishing.create( + workspace_id="workspace-b", + user_id="user-b", + payload=payload, + idempotency_key="cross-workspace-key", + ) + + +async def test_cross_workspace_asset_cannot_be_queued_for_publishing(social_container) -> None: + asset_id = await registered_output_asset(social_container, "workspace-a") + account_b = await connected_account(social_container, "workspace-b") + post = await social_container.social.publishing.create( + workspace_id="workspace-b", + user_id="user-b", + payload=SocialPostCreate.model_validate( + { + **post_payload(account_b.id).model_dump(mode="json"), + "media_asset_id": asset_id, + } + ), + idempotency_key="cross-workspace-asset-post", + ) + with pytest.raises(SocialMediaInvalidError): + await social_container.social.publishing.queue( + "workspace-b", post.id, idempotency_key="cross-workspace-asset-publish" + ) + + +async def test_cross_workspace_records_cannot_be_read_or_modified(social_container) -> None: + account_a = await connected_account(social_container, "workspace-a") + account_b = await connected_account(social_container, "workspace-b") + post_b = await social_container.social.publishing.create( + workspace_id="workspace-b", + user_id="user-b", + payload=post_payload(account_b.id), + idempotency_key="workspace-b-post", + ) + job_b = ( + await social_container.social.jobs.repository.create_many( + [ + SocialJob( + workspace_id="workspace-b", + social_post_id=post_b.id, + social_post_target_id=post_b.targets[0].id, + provider="youtube", + status="queued", + idempotency_key="workspace-b-job", + ) + ] + ) + )[0] + await social_container.social.accounts.tokens.store( + "workspace-b", account_b.id, {"access_token": "workspace-b-secret"} + ) + + with pytest.raises(SocialAccountNotFoundError): + await social_container.social.accounts.get("workspace-a", account_b.id) + with pytest.raises(SocialAccountNotFoundError): + await social_container.social.accounts.repository.set_status( + "workspace-a", account_b.id, "disconnected" + ) + with pytest.raises(SocialPostNotFoundError): + await social_container.social.publishing.get("workspace-a", post_b.id) + with pytest.raises(SocialPostNotFoundError): + await social_container.social.publishing.posts.set_status( + "workspace-a", post_b.id, "cancelled" + ) + with pytest.raises(SocialJobNotFoundError): + await social_container.social.jobs.get("workspace-a", job_b.id) + with pytest.raises(SocialJobNotFoundError): + await social_container.social.jobs.repository.transition( + "workspace-a", job_b.id, "cancelled" + ) + with pytest.raises(SocialReauthRequiredError): + await social_container.social.accounts.tokens.retrieve("workspace-a", account_b.id) + with pytest.raises(SocialAccountNotFoundError): + await social_container.social.analytics.account("workspace-a", account_b.id) + with pytest.raises(SocialPostNotFoundError): + await social_container.social.analytics.post("workspace-a", post_b.id) + assert account_a.id != account_b.id + + +async def test_token_401_is_refreshed_and_retried_once(social_container, monkeypatch) -> None: + account = await connected_account(social_container, "workspace-refresh") + await social_container.social.accounts.tokens.store( + "workspace-refresh", + account.id, + {"access_token": "expired-access", "refresh_token": "refresh-token"}, + ) + adapter = social_container.social.accounts.providers.get("youtube") + + async def refreshed(_: dict[str, object]) -> dict[str, object]: + return {"access_token": "fresh-access", "expires_in": 3600} + + monkeypatch.setattr(adapter, "refresh_token", refreshed) + received: list[str] = [] + + async def protected_call(token: dict[str, object]) -> str: + received.append(str(token["access_token"])) + if len(received) == 1: + raise SocialReauthRequiredError("first credential was rejected") + return "ok" + + assert await social_container.social.oauth.execute_with_reauth_retry( + workspace_id="workspace-refresh", + account_id=account.id, + operation=protected_call, + ) == "ok" + assert received == ["expired-access", "fresh-access"] + + +async def test_resumable_upload_state_is_encrypted_and_excluded_from_job_views( + social_container, +) -> None: + account = await connected_account(social_container, "workspace-upload-state") + post = await social_container.social.publishing.create( + workspace_id="workspace-upload-state", + user_id="user", + payload=post_payload(account.id), + idempotency_key="upload-state-post", + ) + job = ( + await social_container.social.jobs.repository.create_many( + [ + SocialJob( + workspace_id="workspace-upload-state", + social_post_id=post.id, + social_post_target_id=post.targets[0].id, + provider="youtube", + status="queued", + idempotency_key="upload-state-job", + ) + ] + ) + )[0] + session_url = "https://www.googleapis.com/upload/youtube/v3/videos?upload_id=bearer-like" + await social_container.social.jobs.repository.set_provider_state( + "workspace-upload-state", job.id, {"youtube_upload_session_url": session_url} + ) + async with social_container.social.database.session("workspace-upload-state") as session: + stored = await session.scalar(select(SocialJob).where(SocialJob.id == job.id)) + assert stored is not None and session_url not in (stored.provider_state_encrypted or "") + assert await social_container.social.jobs.repository.get_provider_state( + "workspace-upload-state", job.id + ) == {"youtube_upload_session_url": session_url} + view = await social_container.social.jobs.get("workspace-upload-state", job.id) + assert session_url not in view.model_dump_json() + + +async def test_scheduling_normalizes_to_utc_and_preserves_iana_timezone( + social_container, +) -> None: + account = await connected_account(social_container, "workspace-schedule") + asset_id = await registered_output_asset(social_container, "workspace-schedule") + post = await social_container.social.publishing.create( + workspace_id="workspace-schedule", + user_id="user", + payload=SocialPostCreate.model_validate( + { + **post_payload(account.id).model_dump(mode="json"), + "media_asset_id": asset_id, + } + ), + idempotency_key="schedule-create-key", + ) + payload = SocialScheduleCreate.model_validate( + { + "scheduled_at": "2030-02-01T14:00:00+01:00", + "timezone": "Africa/Lagos", + } + ) + schedule = await social_container.social.scheduling.schedule( + "workspace-schedule", post.id, payload + ) + assert schedule.timezone == "Africa/Lagos" + assert schedule.scheduled_at.astimezone(timezone.utc).hour == 13 + replacement = await social_container.social.scheduling.schedule( + "workspace-schedule", + post.id, + SocialScheduleCreate.model_validate( + { + "scheduled_at": "2030-02-01T15:00:00+01:00", + "timezone": "Africa/Lagos", + } + ), + ) + assert replacement.id == schedule.id + assert replacement.scheduled_at.astimezone(timezone.utc).hour == 14 + + +def test_schedule_rejects_past_naive_and_invalid_timezone_values() -> None: + with pytest.raises(ValidationError): + SocialScheduleCreate.model_validate( + {"scheduled_at": "2000-01-01T00:00:00+00:00", "timezone": "UTC"} + ) + with pytest.raises(ValidationError): + SocialScheduleCreate.model_validate( + {"scheduled_at": "2030-03-10T01:30:00", "timezone": "America/New_York"} + ) + with pytest.raises(ValidationError): + SocialScheduleCreate.model_validate( + { + "scheduled_at": "2030-03-10T01:30:00-05:00", + "timezone": "Not/A_Timezone", + } + ) + assert SocialScheduleCreate.model_validate( + { + "scheduled_at": "2030-03-10T01:30:00-05:00", + "timezone": "America/New_York", + } + ).timezone == "America/New_York" diff --git a/tests/test_template_marketplace.py b/tests/test_template_marketplace.py new file mode 100644 index 0000000000000000000000000000000000000000..5e4ac1a73b7e937fdbb2b8902b53b7695c8e6ffe --- /dev/null +++ b/tests/test_template_marketplace.py @@ -0,0 +1,82 @@ +from pathlib import Path + +import pytest +from pydantic import ValidationError + +from app.templates.marketplace_schemas import MarketplaceTemplateDefinition + + +def definition() -> dict[str, object]: + return { + "schema_version": 1, + "settings": {"duration_ms": 30_000, "aspect_ratio": "9:16", "frame_rate": 30}, + "slots": [ + { + "id": "hero-video", + "type": "video", + "label": "Hero video", + "required": True, + "accepted_media_types": ["video/*"], + } + ], + "tracks": [ + { + "id": "video-track", + "type": "video", + "name": "Video", + "clips": [ + { + "id": "hero-clip", + "track_id": "video-track", + "slot_id": "hero-video", + "label": "Hero", + "start_ms": 0, + "duration_ms": 30_000, + } + ], + } + ], + } + + +def test_marketplace_definition_rejects_unknown_slot_references() -> None: + payload = definition() + payload["tracks"][0]["clips"][0]["slot_id"] = "missing" # type: ignore[index] + with pytest.raises(ValidationError): + MarketplaceTemplateDefinition.model_validate(payload) + + +def test_marketplace_definition_rejects_executable_fields() -> None: + payload = definition() + payload["eval"] = "dangerous" + with pytest.raises(ValidationError): + MarketplaceTemplateDefinition.model_validate(payload) + + +def test_marketplace_definition_rejects_slot_track_type_mismatch() -> None: + payload = definition() + payload["slots"][0]["type"] = "audio" # type: ignore[index] + payload["slots"][0]["accepted_media_types"] = ["audio/*"] # type: ignore[index] + with pytest.raises(ValidationError): + MarketplaceTemplateDefinition.model_validate(payload) + + +def test_marketplace_migration_is_additive_versioned_and_rls_protected() -> None: + migration = ( + ( + Path(__file__).resolve().parents[1] + / "app/projects/migrations/0006_template_marketplace.sql" + ) + .read_text(encoding="utf-8") + .lower() + ) + for expected in ( + "create table if not exists marketplace_templates", + "create table if not exists marketplace_template_versions", + "create table if not exists marketplace_template_applications", + "published template versions are immutable", + "force row level security", + "uq_marketplace_template_application_idempotency", + ): + assert expected in migration + assert "drop table" not in migration diff --git a/tests/test_templates.py b/tests/test_templates.py new file mode 100644 index 0000000000000000000000000000000000000000..2eed62e7b181a9457ac1434c117b81f1661b7745 --- /dev/null +++ b/tests/test_templates.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +import base64 +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient + +from app.container import build_container +from app.core.exceptions import TemplateValidationError +from app.models.media import InputMedia, MediaSource, ResolvedRequest +from app.templates.executor import OPERATION_BINDINGS +from app.templates.loader import TemplateLoader +from app.templates.registry import TemplateRegistry +from app.templates.validator import TemplateValidator +from main import create_app + +EXPECTED_CATEGORIES = { + "branding", + "conversion", + "faceless", + "lyrics", + "motivation", + "podcast", + "social", + "subtitles", + "utility", + "youtube", +} + + +def test_builtin_templates_are_loaded_dynamically(settings) -> None: + registry = build_container(settings).template_registry + + assert registry.count == 71 + assert set(registry.categories()) == EXPECTED_CATEGORIES + assert registry.get("instagram_reel").version == 1 + assert registry.get("instagram_reel@latest").version == 1 + assert registry.get("instagram_reel@1").name == "Instagram Reel" + + +def test_parameter_substitution_preserves_declared_types(settings) -> None: + registry = build_container(settings).template_registry + + prepared = registry.prepare("youtube_shorts@1", {"crf": 19, "max_duration": 42.5}) + + trim_step = prepared.pipeline[0] + compress_step = prepared.pipeline[-1] + assert trim_step.parameters["duration"] == 42.5 + assert isinstance(trim_step.parameters["duration"], float) + assert compress_step.parameters["crf"] == 19 + assert isinstance(compress_step.parameters["crf"], int) + + +def test_template_registry_keeps_old_versions(tmp_path: Path) -> None: + root = tmp_path / "templates" + root.mkdir() + (root / "versions.yaml").write_text( + """ +templates: + - id: sample + name: Sample One + category: custom + description: First stable workflow. + author: Tests + version: 1 + tags: [test] + estimated_runtime: fast + supported_inputs: [video] + supported_outputs: [source] + parameters: {} + pipeline: [{operation: download}] + output: {format: source} + examples: [] + - id: sample + name: Sample Two + category: custom + description: Second stable workflow. + author: Tests + version: 2 + tags: [test] + estimated_runtime: fast + supported_inputs: [video] + supported_outputs: [source] + parameters: {} + pipeline: [{operation: download}] + output: {format: source} + examples: [] +""", + encoding="utf-8", + ) + validator = TemplateValidator(set(OPERATION_BINDINGS)) + registry = TemplateRegistry(TemplateLoader(root, validator), validator) + + assert registry.get("sample@1").name == "Sample One" + assert registry.get("sample@2").name == "Sample Two" + assert registry.get("sample@latest").version == 2 + assert registry.get("sample").version == 2 + + +def test_invalid_yaml_operation_is_never_registered(tmp_path: Path) -> None: + root = tmp_path / "templates" + root.mkdir() + (root / "invalid.yaml").write_text( + """ +id: invalid +name: Invalid +category: custom +description: Invalid operation must fail loading. +author: Tests +version: 1 +tags: [test] +estimated_runtime: fast +supported_inputs: [video] +supported_outputs: [mp4] +parameters: {} +pipeline: [{operation: shell_command}] +output: {format: mp4} +examples: [] +""", + encoding="utf-8", + ) + validator = TemplateValidator(set(OPERATION_BINDINGS)) + + with pytest.raises(TemplateValidationError, match="unsupported operation"): + TemplateRegistry(TemplateLoader(root, validator), validator) + + +def test_invalid_yaml_syntax_is_never_loaded(tmp_path: Path) -> None: + root = tmp_path / "templates" + root.mkdir() + (root / "broken.yaml").write_text("id: broken\npipeline: [\n", encoding="utf-8") + validator = TemplateValidator(set(OPERATION_BINDINGS)) + + with pytest.raises(TemplateValidationError, match="syntax"): + TemplateLoader(root, validator).load() + + +def test_required_and_typed_parameters_are_enforced(tmp_path: Path) -> None: + root = tmp_path / "templates" + root.mkdir() + (root / "required.yaml").write_text( + """ +id: required_sample +name: Required Sample +category: custom +description: Exercise strict runtime parameter validation. +author: Tests +version: 1 +tags: [test] +estimated_runtime: fast +supported_inputs: [video] +supported_outputs: [mp4] +parameters: + width: {type: integer, required: true, minimum: 2} +pipeline: [{operation: resize, width: "{{ width }}", height: 720}] +output: {format: mp4} +examples: [] +""", + encoding="utf-8", + ) + validator = TemplateValidator(set(OPERATION_BINDINGS)) + registry = TemplateRegistry(TemplateLoader(root, validator), validator) + + with pytest.raises(TemplateValidationError, match="Required"): + registry.prepare("required_sample", {}) + with pytest.raises(TemplateValidationError, match="must be integer"): + registry.prepare("required_sample", {"width": "1080"}) + assert ( + registry.prepare("required_sample", {"width": 1080}).pipeline[0].parameters["width"] == 1080 + ) + + +async def test_template_executor_calls_existing_operation(settings, tmp_path, monkeypatch) -> None: + container = build_container(settings) + source = tmp_path / "source.wav" + source.write_bytes(b"RIFF-test-audio") + + async def fake_probe(inputs): + return [ + { + "filename": media.filename, + "mime_type": media.mime_type, + "size": media.size, + } + for media in inputs + ] + + async def fake_ffmpeg(args, *, operation, timeout=None): + output = Path(args[-1]) + output.parent.mkdir(parents=True, exist_ok=True) + output.write_bytes(b"ID3-template-output") + + monkeypatch.setattr(container.processor, "probe_inputs", fake_probe) + monkeypatch.setattr(container.ffmpeg, "run", fake_ffmpeg) + resolved = ResolvedRequest( + request_id="5fdbe750-4cb7-4f87-aa5c-3df50c3a6629", + inputs=[ + InputMedia( + source=MediaSource.MULTIPART, + filename=source.name, + mime_type="audio/wav", + temp_path=source, + size=source.stat().st_size, + ) + ], + ) + + response = await container.template_executor.execute(resolved, "mp3@1", {}) + + assert response.success is True + assert response.download_url is not None + assert response.metadata["template"]["id"] == "mp3" + assert response.metadata["operations"] == ["convert_audio"] + published = container.cleanup.resolve_download( + resolved.request_id, Path(response.download_url).name + ) + assert published.read_bytes() == b"ID3-template-output" + + +def test_template_rest_endpoints_and_nested_input(settings, monkeypatch) -> None: + application = create_app(settings) + container = application.state.container + + async def fake_probe(inputs): + return [ + { + "filename": media.filename, + "mime_type": media.mime_type, + "size": media.size, + } + for media in inputs + ] + + async def fake_ffmpeg(args, *, operation, timeout=None): + output = Path(args[-1]) + output.parent.mkdir(parents=True, exist_ok=True) + output.write_bytes(b"ID3-rest-template") + + monkeypatch.setattr(container.processor, "probe_inputs", fake_probe) + monkeypatch.setattr(container.ffmpeg, "run", fake_ffmpeg) + + with TestClient(application) as client: + listing = client.get("/v1/templates") + categories = client.get("/v1/templates/categories") + details = client.get("/v1/templates/instagram_reel@1") + execution = client.post( + "/v1/templates/run", + json={ + "template": "mp3@latest", + "input": { + "base64": base64.b64encode(b"RIFF-rest-audio").decode(), + "filename": "audio.wav", + "mime_type": "audio/wav", + }, + "parameters": {}, + }, + ) + + assert listing.status_code == 200 + assert listing.json()["metadata"]["count"] == 71 + assert set(categories.json()["metadata"]["categories"]) == EXPECTED_CATEGORIES + assert details.json()["metadata"]["template"]["version"] == 1 + assert execution.status_code == 200 + assert execution.json()["metadata"]["template"]["id"] == "mp3" diff --git a/tests/test_tenant_foundation.py b/tests/test_tenant_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..f388fde9d3a6171ef9c29e4417b784a673d47fb0 --- /dev/null +++ b/tests/test_tenant_foundation.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +from pathlib import Path +from uuid import uuid4 + +import pytest + +from app.container import build_container +from app.core.config import Settings +from app.security.assets import CanonicalAssetNotFoundError +from app.security.schemas import APIKeyCreate +from app.social.models import SocialAccount + + +def foundation_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + auth_enabled=True, + social_oauth_encryption_key="test-only-encryption-material", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +async def _create_key(container, name: str, **kwargs: str): + return await container.api_keys.create( + APIKeyCreate( + name=name, + environment="test", + role=None, + scopes=["operations:execute", "operations:read"], + ), + created_by="tests", + **kwargs, + ) + + +@pytest.mark.asyncio +async def test_api_key_is_resolved_to_persisted_membership_and_rotation_preserves_it( + tmp_path: Path, +) -> None: + container = build_container(foundation_settings(tmp_path)) + await container.security_database.initialize() + try: + first, first_secret = await _create_key(container, "First") + first_context = await container.api_keys.authenticate(first_secret) + assert first_context.workspace_id and first_context.user_id + assert first_context.workspace_id != first.id + assert first_context.user_id != first.id + + sibling, sibling_secret = await _create_key( + container, + "Sibling", + workspace_id=first_context.workspace_id, + user_id=first_context.user_id, + ) + sibling_context = await container.api_keys.authenticate(sibling_secret) + assert sibling_context.workspace_id == first_context.workspace_id + assert sibling_context.user_id == first_context.user_id + + isolated, isolated_secret = await _create_key(container, "Isolated") + isolated_context = await container.api_keys.authenticate(isolated_secret) + assert isolated_context.workspace_id != first_context.workspace_id + + rotated, rotated_secret = await container.api_keys.rotate( + first.id, 0, created_by="tests" + ) + rotated_context = await container.api_keys.authenticate(rotated_secret) + assert rotated_context.workspace_id == first_context.workspace_id + assert rotated_context.user_id == first_context.user_id + finally: + await container.security_database.close() + + +@pytest.mark.asyncio +async def test_canonical_asset_cannot_be_claimed_by_another_workspace(tmp_path: Path) -> None: + container = build_container(foundation_settings(tmp_path)) + await container.security_database.initialize() + try: + _, secret_a = await _create_key(container, "A") + _, secret_b = await _create_key(container, "B") + context_a = await container.api_keys.authenticate(secret_a) + context_b = await container.api_keys.authenticate(secret_b) + request_id = str(uuid4()) + output = container.settings.output_dir / request_id + output.mkdir(parents=True) + path = output / "asset.mp4" + path.write_bytes(b"owned output") + + asset = await container.assets.register_output( + workspace_id=str(context_a.workspace_id), + user_id=context_a.user_id, + request_id=request_id, + path=path, + mime_type="video/mp4", + ) + assert asset.workspace_id == context_a.workspace_id + owned = await container.assets.get_owned( + workspace_id=str(context_a.workspace_id), + request_id=request_id, + filename=path.name, + ) + assert owned.id == asset.id + with pytest.raises(CanonicalAssetNotFoundError): + await container.assets.get_owned( + workspace_id=str(context_b.workspace_id), + request_id=request_id, + filename=path.name, + ) + path.write_bytes(b"tampered") + with pytest.raises(CanonicalAssetNotFoundError): + await container.assets.verify_file(asset, path) + finally: + await container.security_database.close() + + +@pytest.mark.asyncio +async def test_legacy_social_rows_are_adopted_without_reusing_api_key_tenant_id( + tmp_path: Path, +) -> None: + container = build_container(foundation_settings(tmp_path)) + await container.security_database.initialize() + await container.social.initialize() + try: + key, secret = await _create_key(container, "Legacy") + context = await container.api_keys.authenticate(secret) + legacy_account = await container.social.accounts.repository.create( + SocialAccount( + workspace_id=key.id, + provider="youtube", + account_type="channel", + external_account_id="legacy-channel", + status="connected", + ) + ) + await container.social.adopt_legacy_workspaces( + await container.tenants.list_principals() + ) + adopted = await container.social.accounts.repository.get( + str(context.workspace_id), legacy_account.id + ) + assert adopted.workspace_id == context.workspace_id + finally: + await container.social.close() + await container.security_database.close() diff --git a/tests/test_tiktok_foundation.py b/tests/test_tiktok_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..3a0b5b5770574848952de2c17db01b2ebbd0affb --- /dev/null +++ b/tests/test_tiktok_foundation.py @@ -0,0 +1,371 @@ +"""Phase 4A TikTok Login Kit foundation coverage. + +All provider traffic uses MockTransport. Normal CI never needs TikTok +credentials or an interactive browser consent flow. +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialOAuthStateError, + SocialPermissionDeniedError, + SocialReauthRequiredError, +) +from app.social.models import OAuthState, SocialAccountToken +from app.social.providers.tiktok import TikTokProvider +from app.social.schemas.accounts import SocialAccountConnectRequest + + +def tiktok_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="test-only-encryption-material", + social_oauth_redirect_base_url="https://api.example.com", + tiktok_client_key="tiktok-client-key", + tiktok_client_secret="tiktok-client-secret", + tiktok_redirect_uri=( + "https://api.example.com/v1/social/accounts/tiktok/callback" + ), + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +async def test_tiktok_web_authorization_uses_minimum_scope_and_no_unsupported_pkce( + tmp_path: Path, +) -> None: + provider = TikTokProvider(tiktok_settings(tmp_path)) + try: + url = await provider.get_authorization_url( + state="s" * 43, + redirect_uri="https://api.example.com/v1/social/accounts/tiktok/callback", + code_challenge="challenge-that-web-login-kit-does-not-support", + ) + finally: + await provider.close() + + parsed = urlparse(url) + query = parse_qs(parsed.query) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == ( + "https://www.tiktok.com/v2/auth/authorize/" + ) + assert query["client_key"] == ["tiktok-client-key"] + assert query["response_type"] == ["code"] + assert query["scope"] == ["user.info.basic"] + assert query["state"] == ["s" * 43] + assert "code_challenge" not in query + assert "code_challenge_method" not in query + + +async def test_tiktok_exchange_refresh_discovery_and_revoke_use_official_v2_endpoints( + tmp_path: Path, +) -> None: + calls: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + if request.url.path == "/v2/oauth/token/": + form = parse_qs(request.content.decode()) + assert form["client_key"] == ["tiktok-client-key"] + assert form["client_secret"] == ["tiktok-client-secret"] + if form["grant_type"] == ["authorization_code"]: + assert form["code"] == ["authorization-code"] + assert form["redirect_uri"] == [ + "https://api.example.com/v1/social/accounts/tiktok/callback" + ] + assert "code_verifier" not in form + else: + assert form["grant_type"] == ["refresh_token"] + assert form["refresh_token"] == ["refresh-token"] + return httpx.Response( + 200, + json={ + "access_token": "access-token", + "refresh_token": "rotated-refresh-token", + "expires_in": 86400, + "refresh_expires_in": 31536000, + "open_id": "open-id", + "scope": "user.info.basic", + "token_type": "Bearer", + }, + ) + if request.url.path == "/v2/user/info/": + assert request.headers["authorization"] == "Bearer access-token" + assert parse_qs(request.url.query.decode())["fields"] == [ + "open_id,union_id,avatar_url,display_name" + ] + return httpx.Response( + 200, + json={ + "data": { + "user": { + "open_id": "open-id", + "union_id": "union-id", + "display_name": "TikTok Creator", + "avatar_url": "https://example.com/avatar.jpg", + } + }, + "error": {"code": "ok", "message": ""}, + }, + ) + assert request.url.path == "/v2/oauth/revoke/" + form = parse_qs(request.content.decode()) + assert form == { + "client_key": ["tiktok-client-key"], + "client_secret": ["tiktok-client-secret"], + "token": ["access-token"], + } + return httpx.Response(200) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(tiktok_settings(tmp_path), http_client=client) + try: + token = await provider.exchange_code( + code="authorization-code", + redirect_uri="https://api.example.com/v1/social/accounts/tiktok/callback", + code_verifier="unused-web-verifier", + ) + account = await provider.get_account(token) + refreshed = await provider.refresh_token( + {"access_token": "old-access", "refresh_token": "refresh-token"} + ) + await provider.revoke_token({"access_token": "access-token"}) + finally: + await client.aclose() + + assert account == { + "external_account_id": "open-id", + "account_type": "creator", + "username": None, + "display_name": "TikTok Creator", + "avatar_url": "https://example.com/avatar.jpg", + "metadata": { + "tiktok_open_id": "open-id", + "tiktok_union_id": "union-id", + }, + } + assert refreshed["refresh_token"] == "rotated-refresh-token" + assert calls == [ + "/v2/oauth/token/", + "/v2/user/info/", + "/v2/oauth/token/", + "/v2/oauth/revoke/", + ] + + +async def test_tiktok_invalid_code_is_normalized_without_provider_secret( + tmp_path: Path, +) -> None: + secret = "authorization-code-that-must-not-leak" + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + json={ + "error": "invalid_grant", + "error_description": f"bad code {secret}", + "log_id": "provider-log-id", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(tiktok_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialReauthRequiredError) as raised: + await provider.exchange_code( + code=secret, + redirect_uri="https://api.example.com/v1/social/accounts/tiktok/callback", + ) + finally: + await client.aclose() + assert secret not in str(raised.value) + assert "provider-log-id" not in str(raised.value) + + +async def test_tiktok_oauth_callback_is_single_use_duplicate_safe_and_workspace_bound( + tmp_path: Path, +) -> None: + settings = tiktok_settings(tmp_path) + container = build_container(settings) + await container.social.initialize() + adapter = container.social.accounts.providers.get("tiktok") + assert isinstance(adapter, TikTokProvider) + await adapter._client.aclose() + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v2/oauth/token/": + return httpx.Response( + 200, + json={ + "access_token": "token-that-must-stay-encrypted", + "refresh_token": "refresh-that-must-stay-encrypted", + "expires_in": 86400, + "scope": "user.info.basic", + "token_type": "Bearer", + }, + ) + return httpx.Response( + 200, + json={ + "data": { + "user": { + "open_id": "stable-open-id", + "display_name": "Workspace Creator", + } + }, + "error": {"code": "ok", "message": ""}, + }, + ) + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + try: + first_connect = await container.social.oauth.connect( + provider="tiktok", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(), + ) + first_state = parse_qs(urlparse(first_connect.authorization_url or "").query)[ + "state" + ][0] + assert "code_challenge" not in parse_qs( + urlparse(first_connect.authorization_url or "").query + ) + first = await container.social.oauth.callback( + provider="tiktok", state=first_state, code="first-code" + ) + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.callback( + provider="tiktok", state=first_state, code="replayed-code" + ) + + second_connect = await container.social.oauth.connect( + provider="tiktok", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(), + ) + second_state = parse_qs( + urlparse(second_connect.authorization_url or "").query + )["state"][0] + second = await container.social.oauth.callback( + provider="tiktok", state=second_state, code="second-code" + ) + + assert first.id == second.id + accounts = await container.social.accounts.list("workspace-a") + assert [account.id for account in accounts] == [first.id] + assert "token-that-must-stay-encrypted" not in first.model_dump_json() + with pytest.raises(SocialAccountNotFoundError): + await container.social.accounts.get("workspace-b", first.id) + + async with container.social.database.session("workspace-a") as session: + stored = await session.scalar( + select(SocialAccountToken).where( + SocialAccountToken.social_account_id == first.id + ) + ) + assert stored is not None + assert stored.encrypted_payload + assert "token-that-must-stay-encrypted" not in stored.encrypted_payload + finally: + await container.social.close() + await container.security_database.close() + + +async def test_tiktok_state_expiry_provider_binding_and_redirect_validation( + tmp_path: Path, +) -> None: + container = build_container(tiktok_settings(tmp_path)) + await container.social.initialize() + try: + assert container.social.oauth._redirect_uri("tiktok", None) == ( + "https://api.example.com/v1/social/accounts/tiktok/callback" + ) + with pytest.raises(SocialPermissionDeniedError): + container.social.oauth._redirect_uri( + "tiktok", + "https://attacker.example/v1/social/accounts/tiktok/callback", + ) + + state = await container.social.oauth.states.create( + provider="tiktok", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=settings_redirect(container.settings), + ) + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=state.state, provider="youtube" + ) + consumed = await container.social.oauth.states.consume( + state=state.state, provider="tiktok" + ) + assert consumed.workspace_id == "workspace-a" + assert consumed.user_id == "user-a" + + expired = OAuthState( + state="expired-tiktok-state-value-that-is-long-enough", + provider="tiktok", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=settings_redirect(container.settings), + expires_at=datetime.now(timezone.utc) - timedelta(seconds=1), + ) + async with container.social.database.session("workspace-a") as session: + session.add(expired) + await session.commit() + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=expired.state, provider="tiktok" + ) + finally: + await container.social.close() + await container.security_database.close() + + +def settings_redirect(settings: Settings) -> str: + return settings.tiktok_redirect_uri + + +async def test_tiktok_capability_discovery_does_not_advertise_publishing( + tmp_path: Path, +) -> None: + container = build_container(tiktok_settings(tmp_path)) + try: + tiktok = container.social.accounts.get_provider("tiktok") + assert tiktok.available + assert tiktok.configured + assert tiktok.capabilities.implementation_status == "implemented" + assert tiktok.capabilities.required_scopes == ["user.info.basic"] + assert tiktok.capabilities.account_types == ["creator"] + assert not tiktok.capabilities.video + assert not tiktok.capabilities.video_upload + assert not tiktok.capabilities.direct_publish + assert not tiktok.capabilities.draft_upload + assert not tiktok.capabilities.scheduled_publish + assert not tiktok.capabilities.delete_post + assert tiktok.capabilities.analytics + assert tiktok.capabilities.analytics_required_scopes == ["video.list"] + finally: + await container.social.close() + await container.security_database.close() diff --git a/tests/test_tiktok_production.py b/tests/test_tiktok_production.py new file mode 100644 index 0000000000000000000000000000000000000000..cce5ca507495f900c258fb8ebba5155369f84bbe --- /dev/null +++ b/tests/test_tiktok_production.py @@ -0,0 +1,597 @@ +"""Phase 4C TikTok analytics, security, tenancy, and certification coverage. + +Normal CI uses only SQLite and mocked official TikTok endpoints. Live provider +traffic is opt-in and requires a dedicated test creator plus explicit consent +to create a SELF_ONLY post. +""" + +from __future__ import annotations + +import json +import logging +import os +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.core.logger import JsonFormatter +from app.mcp.registry import MCPRegistry +from app.mcp.server import create_mcp_server +from app.security.context import AuthContext, auth_context, http_auth_applied +from app.services.ffprobe_service import FFprobeService +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialJobNotFoundError, + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialPublishFailedError, + SocialRateLimitedError, + SocialReauthRequiredError, +) +from app.social.domain.retry import classify_retry +from app.social.models import ( + SocialAccount, + SocialAuditEvent, + SocialJob, + SocialMediaAsset, + SocialPost, + SocialPostMetric, + SocialPostTarget, +) +from app.social.providers.tiktok import TikTokProvider +from app.social.schemas.accounts import SocialAccountConnectRequest, SocialAccountView +from app.social.schemas.jobs import SocialJobView +from app.social.schemas.tiktok import TikTokPostMetadata + + +def phase4c_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "phase-4c-test-encryption-material", + "tiktok_client_key": "tiktok-client-key", + "tiktok_client_secret": "tiktok-client-secret", + "tiktok_redirect_uri": ( + "https://api.example.com/v1/social/accounts/tiktok/callback" + ), + "tiktok_direct_post_enabled": True, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +@pytest.fixture +async def phase4c_container(tmp_path: Path): + container = build_container(phase4c_settings(tmp_path)) + await container.social.initialize() + try: + yield container + finally: + await container.social.close() + await container.security_database.close() + + +async def test_tiktok_video_query_analytics_normalizes_only_official_metrics( + tmp_path: Path, +) -> None: + secret = "analytics-access-token-that-must-not-leak" + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v2/video/query/" + assert request.headers["authorization"] == f"Bearer {secret}" + assert secret not in str(request.url) + fields = parse_qs(request.url.query.decode())["fields"][0].split(",") + assert {"view_count", "like_count", "comment_count", "share_count"} <= set( + fields + ) + assert json.loads(request.content) == { + "filters": {"video_ids": ["public-video-id"]} + } + return httpx.Response( + 200, + json={ + "data": { + "videos": [ + { + "id": "public-video-id", + "create_time": 1_785_456_000, + "share_url": "https://www.tiktok.com/@creator/video/public-video-id", + "view_count": 101, + "like_count": 22, + "comment_count": 3, + "share_count": 4, + "title": "Provider-returned title", + "access_token": secret, + } + ] + }, + "error": {"code": "ok", "message": "", "log_id": "safe-log-id"}, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(phase4c_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics( + {"access_token": secret}, "public-video-id" + ) + finally: + await client.aclose() + + assert result["status"] == "available" + assert result["views"] == 101 + assert result["likes"] == 22 + assert result["comments"] == 3 + assert result["shares"] == 4 + assert result["published_at"] == 1_785_456_000 + assert result["raw_metrics"]["title"] == "Provider-returned title" + assert secret not in json.dumps(result) + + +async def test_tiktok_analytics_scope_is_explicit_and_persists_public_video_metrics( + phase4c_container, +) -> None: + social = phase4c_container.social + provider = social.accounts.providers.get("tiktok") + assert provider.capabilities.analytics + assert provider.capabilities.analytics_required_scopes == ["video.list"] + + normal = await social.oauth.connect( + provider="tiktok", + workspace_id="workspace-tiktok", + user_id="user-tiktok", + payload=SocialAccountConnectRequest(), + ) + elevated = await social.oauth.connect( + provider="tiktok", + workspace_id="workspace-tiktok", + user_id="user-tiktok", + payload=SocialAccountConnectRequest(authorization_purpose="analytics"), + ) + assert "video.list" not in parse_qs( + urlparse(normal.authorization_url or "").query + )["scope"][0].split(",") + assert "video.list" in parse_qs( + urlparse(elevated.authorization_url or "").query + )["scope"][0].split(",") + + account = await social.accounts.repository.create( + SocialAccount( + workspace_id="workspace-tiktok", + provider="tiktok", + account_type="creator", + external_account_id="open-id", + status="connected", + ) + ) + await social.accounts.tokens.store( + "workspace-tiktok", + account.id, + {"access_token": "encrypted-analytics-token"}, + scopes=["user.info.basic", "video.publish"], + ) + readiness = await social.analytics.account("workspace-tiktok", account.id) + assert readiness == { + "account_id": account.id, + "metrics": [], + "status": "unavailable", + "reason": "TIKTOK_ANALYTICS_ADDITIONAL_AUTHORIZATION_REQUIRED", + "required_scopes": ["video.list"], + } + + await social.accounts.tokens.store( + "workspace-tiktok", + account.id, + {"access_token": "encrypted-analytics-token"}, + scopes=["user.info.basic", "video.publish", "video.list"], + ) + post, targets = await social.publishing.posts.create( + SocialPost( + workspace_id="workspace-tiktok", + media_asset_id="workspace-owned-asset", + ), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="tiktok", + status="published", + external_post_id="private-publish-id", + platform_metadata={ + "provider": {"public_post_ids": ["public-video-id"]} + }, + ) + ], + ) + requested_ids: list[str] = [] + + async def metrics(_: dict[str, object], video_id: str) -> dict[str, object]: + requested_ids.append(video_id) + return { + "status": "available", + "views": 12, + "likes": 3, + "comments": 2, + "shares": 1, + "raw_metrics": { + "view_count": 12, + "authorization": "Bearer secret-that-must-not-persist", + }, + } + + provider.get_metrics = metrics # type: ignore[method-assign] + result = await social.analytics.post("workspace-tiktok", post.id) + assert requested_ids == ["public-video-id"] + assert result["metrics"][0]["views"] == 12 + assert result["metrics"][0]["raw_metrics"] == {"view_count": 12} + assert "secret-that-must-not-persist" not in str(result) + + async with social.database.session("workspace-tiktok") as session: + record = await session.scalar( + select(SocialPostMetric).where( + SocialPostMetric.social_post_target_id == targets[0].id + ) + ) + assert record is not None + assert record.raw_metrics == {"view_count": 12} + + +async def test_tiktok_cross_workspace_accounts_targets_jobs_assets_and_analytics_fail( + phase4c_container, +) -> None: + social = phase4c_container.social + account_b = await social.accounts.repository.create( + SocialAccount( + workspace_id="workspace-b", + provider="tiktok", + account_type="creator", + external_account_id="workspace-b-open-id", + status="connected", + ) + ) + await social.accounts.tokens.store( + "workspace-b", + account_b.id, + {"access_token": "workspace-b-token"}, + scopes=["user.info.basic", "video.list"], + ) + asset_b = await social.media_assets.repository.create( + SocialMediaAsset( + workspace_id="workspace-b", + request_id="00000000-0000-0000-0000-00000000000b", + filename="video.mp4", + mime_type="video/mp4", + file_size=10, + ) + ) + post_b, targets_b = await social.publishing.posts.create( + SocialPost(workspace_id="workspace-b", media_asset_id=asset_b.id), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account_b.id, + provider="tiktok", + status="published", + external_post_id="workspace-b-publish-id", + platform_metadata={ + "provider": {"public_post_ids": ["workspace-b-video-id"]} + }, + ) + ], + ) + job_b = ( + await social.jobs.repository.create_many( + [ + SocialJob( + workspace_id="workspace-b", + social_post_id=post_b.id, + social_post_target_id=targets_b[0].id, + provider="tiktok", + status="queued", + idempotency_key="workspace-b-job-key", + ) + ] + ) + )[0] + + with pytest.raises(SocialAccountNotFoundError): + await social.accounts.get("workspace-a", account_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.get("workspace-a", post_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.posts.set_target_status( + "workspace-a", targets_b[0].id, "failed" + ) + with pytest.raises(SocialJobNotFoundError): + await social.jobs.get("workspace-a", job_b.id) + with pytest.raises(SocialMediaInvalidError): + await social.media_assets.repository.get("workspace-a", asset_b.id) + with pytest.raises(SocialAccountNotFoundError): + await social.analytics.account("workspace-a", account_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.analytics.post("workspace-a", post_b.id) + + +async def test_tiktok_tokens_are_redacted_from_views_logs_and_audit_records( + phase4c_container, +) -> None: + secret = "phase-4c-secret-token" + account = SocialAccount( + workspace_id="workspace-a", + provider="tiktok", + account_type="creator", + external_account_id="open-id", + status="connected", + metadata_json={ + "display": "Creator", + "access_token": secret, + "provider_message": f"Authorization: Bearer {secret}", + }, + ) + job = SocialJob( + workspace_id="workspace-a", + social_post_id="post-a", + provider="tiktok", + status="queued", + payload_json={"access_token": secret, "message": f"Bearer {secret}"}, + provider_state_encrypted=f"encrypted:{secret}", + ) + assert secret not in SocialAccountView.from_record(account).model_dump_json() + assert secret not in SocialJobView.from_record(job).model_dump_json() + + record = logging.LogRecord( + "security-test", + logging.ERROR, + __file__, + 1, + f"provider failed Authorization: Bearer {secret}", + (), + None, + ) + record.provider_payload = { + "refresh_token": secret, + "message": f"access_token={secret}", + } + rendered = JsonFormatter().format(record) + assert secret not in rendered + assert "[REDACTED]" in rendered + + await phase4c_container.social.audit.record( + workspace_id="workspace-a", + event_type="SOCIAL_TIKTOK_SECURITY_TEST", + provider="tiktok", + metadata={ + "client_secret": secret, + "message": f"Authorization: Bearer {secret}", + }, + ) + async with phase4c_container.social.database.session("workspace-a") as session: + audit = await session.scalar( + select(SocialAuditEvent).where( + SocialAuditEvent.event_type == "SOCIAL_TIKTOK_SECURITY_TEST" + ) + ) + assert audit is not None + assert secret not in json.dumps(audit.metadata_json) + + +@pytest.mark.parametrize( + ("status_code", "code", "error_type", "retryable"), + [ + (429, "rate_limit_exceeded", SocialRateLimitedError, True), + (500, "internal_error", SocialProviderUnavailableError, True), + (502, "server_error", SocialProviderUnavailableError, True), + (503, "server_error", SocialProviderUnavailableError, True), + (504, "server_error", SocialProviderUnavailableError, True), + (401, "access_token_expired", SocialReauthRequiredError, True), + (403, "scope_not_authorized", SocialPermissionDeniedError, False), + (400, "invalid_param", SocialPublishFailedError, False), + ], +) +def test_tiktok_retry_matrix_is_bounded_and_permanent_errors_fail( + status_code: int, + code: str, + error_type: type[Exception], + retryable: bool, +) -> None: + response = httpx.Response(status_code, json={"error": {"code": code}}) + with pytest.raises(error_type) as raised: + TikTokProvider._raise_tiktok_error( + response, response.json(), operation="production audit" + ) + decision = classify_retry( + status_code=getattr(raised.value, "status_code", status_code), attempt=1 + ) + assert decision.retryable is retryable + if status_code == 401: + assert decision.refresh_token_first + assert not classify_retry(status_code=401, attempt=2).retryable + + +async def test_tiktok_network_timeout_is_safe_and_retryable(tmp_path: Path) -> None: + secret = "timeout-token-that-must-not-leak" + + async def handler(request: httpx.Request) -> httpx.Response: + raise httpx.ReadTimeout("Authorization: Bearer " + secret, request=request) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(phase4c_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError) as raised: + await provider.get_metrics({"access_token": secret}, "video-id") + finally: + await client.aclose() + assert secret not in str(raised.value) + assert classify_retry(status_code=raised.value.status_code, attempt=1).retryable + + +async def test_mcp_registers_social_contract_and_enforces_analytics_scope( + phase4c_container, +) -> None: + server = create_mcp_server(phase4c_container) + tools = {tool.name for tool in await server.list_tools()} + assert { + "social.list_providers", + "social.get_capabilities", + "social.list_accounts", + "social.create_post", + "social.publish_post", + "social.schedule_post", + "social.get_job", + "social.get_analytics", + } <= tools + + context = AuthContext( + api_key_id="workspace-a", + key_name="phase-4c", + key_prefix="mp_test", + environment="test", + role="viewer", + scopes=frozenset({"social:accounts:read"}), + requests_per_minute=100, + concurrent_jobs=2, + uploads_per_hour=10, + processing_bytes_per_day=1_000_000, + expires_at=None, + ) + auth_token = auth_context.set(context) + http_token = http_auth_applied.set(True) + called = False + + async def forbidden_action() -> dict[str, object]: + nonlocal called + called = True + return {"metrics": []} + + try: + result = await MCPRegistry(phase4c_container).run_metadata_tool( + "social.get_analytics", + forbidden_action, + required_scope="social:analytics:read", + ) + finally: + http_auth_applied.reset(http_token) + auth_context.reset(auth_token) + assert result["success"] is False + assert result["error"]["code"] == "FORBIDDEN" + assert not called + assert "token" not in json.dumps(result).lower() + + +@pytest.mark.skipif( + os.getenv("RUN_TIKTOK_INTEGRATION_TESTS", "").lower() != "true", + reason="Set RUN_TIKTOK_INTEGRATION_TESTS=true for a dedicated TikTok test creator.", +) +async def test_live_tiktok_self_only_publish_status_and_analytics() -> None: + """Optional destructive live smoke test guarded by two explicit opt-ins. + + Required secrets are read only from the test process environment. TikTok + currently provides no official delete endpoint, so the test insists on + SELF_ONLY privacy and documents that the created post remains in the + dedicated test account. + """ + + if os.getenv("TIKTOK_TEST_ALLOW_PUBLISH", "").lower() != "true": + pytest.skip("Set TIKTOK_TEST_ALLOW_PUBLISH=true to create a SELF_ONLY post.") + token = os.getenv("TIKTOK_TEST_ACCESS_TOKEN") + media_value = os.getenv("TIKTOK_TEST_VIDEO_PATH") + client_key = os.getenv("TIKTOK_CLIENT_KEY") + client_secret = os.getenv("TIKTOK_CLIENT_SECRET") + redirect_uri = os.getenv("TIKTOK_REDIRECT_URI") + if not all((token, media_value, client_key, client_secret, redirect_uri)): + pytest.skip("Dedicated TikTok credentials, token, and test video are not configured.") + media_path = Path(str(media_value)).resolve() + if not media_path.is_file(): + pytest.skip("TIKTOK_TEST_VIDEO_PATH is not a readable file.") + + settings = Settings( + _env_file=None, + tiktok_client_key=str(client_key), + tiktok_client_secret=str(client_secret), + tiktok_redirect_uri=str(redirect_uri), + tiktok_direct_post_enabled=True, + max_upload_size=max(media_path.stat().st_size, 1_048_576), + whisper_model="tiny", + ) + provider = TikTokProvider(settings) + try: + account = await provider.get_account({"access_token": str(token)}) + assert account["external_account_id"] + creator = await provider.get_publish_options({"access_token": str(token)}) + if "SELF_ONLY" not in creator["privacy_level_options"]: + pytest.skip("Dedicated TikTok creator does not currently allow SELF_ONLY posts.") + probe = await FFprobeService(settings).probe(media_path) + metadata = TikTokPostMetadata.model_validate( + { + "title": "MediaRouter Phase 4C integration verification", + "privacy_level": "SELF_ONLY", + "disable_comment": True, + "disable_duet": True, + "disable_stitch": True, + "brand_content_toggle": False, + "brand_organic_toggle": False, + "is_aigc": False, + "music_usage_confirmed": True, + } + ) + state: dict[str, object] = {} + + async def persist(value: dict[str, object]) -> None: + state.clear() + state.update(value) + + uploaded = await provider.upload_media( + {"access_token": str(token)}, + { + "path": media_path, + "mime_type": "video/mp4", + "file_size": media_path.stat().st_size, + "probe": probe, + "tiktok_post_info": metadata.to_post_info(), + "provider_state": state, + "persist_provider_state": persist, + }, + ) + publish_id = str(uploaded["id"]) + terminal: dict[str, object] | None = None + for _ in range(60): + status = await provider.get_publish_status( + {"access_token": str(token)}, publish_id + ) + if status["status"] in {"published", "failed", "unavailable"}: + terminal = status + break + import asyncio + + await asyncio.sleep(10) + assert terminal is not None + assert terminal["status"] == "published" + public_ids = terminal.get("metadata", {}).get("public_post_ids", []) + granted = { + value + for value in os.getenv("TIKTOK_TEST_GRANTED_SCOPES", "").replace(",", " ").split() + if value + } + if "video.list" in granted and public_ids: + metrics = await provider.get_metrics( + {"access_token": str(token)}, str(public_ids[0]) + ) + assert metrics["status"] in {"available", "unavailable"} + assert not provider.capabilities.delete_post + finally: + await provider.close() diff --git a/tests/test_tiktok_publishing.py b/tests/test_tiktok_publishing.py new file mode 100644 index 0000000000000000000000000000000000000000..dda8a4631b2bab10b0ee213637417fee180be64b --- /dev/null +++ b/tests/test_tiktok_publishing.py @@ -0,0 +1,553 @@ +"""Phase 4B TikTok Direct Post coverage using only mocked official endpoints.""" + +from __future__ import annotations + +import json +from pathlib import Path +from unittest.mock import AsyncMock +from urllib.parse import parse_qs, urlparse +from uuid import uuid4 + +import httpx +import pytest +from pydantic import ValidationError + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialCapabilityUnsupportedError, + SocialIdempotencyConflictError, + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialPublishFailedError, + SocialRateLimitedError, + SocialReauthRequiredError, +) +from app.social.models import SocialAccount, SocialMediaAsset +from app.social.providers.tiktok import TikTokProvider +from app.social.schemas.posts import SocialPostCreate +from app.social.schemas.tiktok import TikTokPostMetadata +from app.social.workers.publisher import SocialPublisher + + +def publishing_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "test-only-encryption-material", + "tiktok_client_key": "tiktok-client-key", + "tiktok_client_secret": "tiktok-client-secret", + "tiktok_redirect_uri": ( + "https://api.example.com/v1/social/accounts/tiktok/callback" + ), + "tiktok_direct_post_enabled": True, + "tiktok_upload_chunk_bytes": 5_000_000, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +def valid_probe(*, duration: float = 15.0) -> dict[str, object]: + return { + "container": "mov,mp4,m4a,3gp,3g2,mj2", + "duration": duration, + "fps": 30.0, + "resolution": {"width": 1080, "height": 1920}, + "video_streams": [{"codec": "h264"}], + "audio_streams": [{"codec": "aac"}], + } + + +def valid_metadata(**overrides: object) -> dict[str, object]: + values: dict[str, object] = { + "title": "A production-safe TikTok post", + "privacy_level": "SELF_ONLY", + "disable_comment": False, + "disable_duet": False, + "disable_stitch": False, + "brand_content_toggle": False, + "brand_organic_toggle": False, + "is_aigc": False, + "music_usage_confirmed": True, + } + values.update(overrides) + return values + + +async def test_direct_post_capabilities_are_fail_closed_and_approval_gated( + tmp_path: Path, +) -> None: + disabled = TikTokProvider( + publishing_settings(tmp_path, tiktok_direct_post_enabled=False) + ) + enabled = TikTokProvider(publishing_settings(tmp_path)) + try: + assert not disabled.capabilities.direct_publish + assert not disabled.capabilities.video_upload + assert disabled.capabilities.publishing_required_scopes == [] + assert enabled.capabilities.direct_publish + assert enabled.capabilities.video_upload + assert enabled.capabilities.video_status + assert enabled.capabilities.scheduled_publish + assert not enabled.capabilities.native_scheduling + assert not enabled.capabilities.delete_post + assert enabled.capabilities.publishing_required_scopes == ["video.publish"] + with pytest.raises(SocialCapabilityUnsupportedError): + await enabled.delete_post({"access_token": "access-token"}, "post-id") + finally: + await disabled.close() + await enabled.close() + + +async def test_publishing_oauth_scope_is_requested_only_by_explicit_elevation( + tmp_path: Path, +) -> None: + provider = TikTokProvider(publishing_settings(tmp_path)) + try: + connection_url = await provider.get_authorization_url( + state="s" * 43, + redirect_uri=provider.redirect_uri, + ) + publishing_url = await provider.get_authorization_url( + state="s" * 43, + redirect_uri=provider.redirect_uri, + additional_scopes=["video.publish"], + ) + finally: + await provider.close() + assert parse_qs(urlparse(connection_url).query)["scope"] == ["user.info.basic"] + assert parse_qs(urlparse(publishing_url).query)["scope"] == [ + "user.info.basic,video.publish" + ] + + +async def test_direct_post_queries_creator_initializes_streams_and_reconciles( + tmp_path: Path, +) -> None: + video = tmp_path / "video.mp4" + video.write_bytes(b"streamed-tiktok-video") + calls: list[str] = [] + persisted: list[dict[str, object]] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(f"{request.method} {request.url.path}") + if request.url.path.endswith("/creator_info/query/"): + assert request.headers["authorization"] == "Bearer access-token" + return httpx.Response(200, json={ + "data": { + "privacy_level_options": ["SELF_ONLY", "PUBLIC_TO_EVERYONE"], + "comment_disabled": False, + "duet_disabled": False, + "stitch_disabled": False, + "max_video_post_duration_sec": 300, + }, + "error": {"code": "ok", "message": ""}, + }) + if request.url.path.endswith("/video/init/"): + payload = json.loads(request.content) + assert payload["source_info"] == { + "source": "FILE_UPLOAD", + "video_size": video.stat().st_size, + "chunk_size": video.stat().st_size, + "total_chunk_count": 1, + } + assert payload["post_info"]["privacy_level"] == "SELF_ONLY" + assert "music_usage_confirmed" not in payload["post_info"] + return httpx.Response(200, json={ + "data": { + "publish_id": "publish-id", + "upload_url": "https://open-upload.tiktokapis.com/video/session", + }, + "error": {"code": "ok", "message": ""}, + }) + if request.method == "PUT": + assert request.headers["content-range"] == ( + f"bytes 0-{video.stat().st_size - 1}/{video.stat().st_size}" + ) + assert request.content == video.read_bytes() + return httpx.Response(201) + if request.url.path.endswith("/status/fetch/"): + return httpx.Response(200, json={ + "data": { + "status": "PUBLISH_COMPLETE", + "publicaly_available_post_id": ["public-video-id"], + "uploaded_bytes": video.stat().st_size, + }, + "error": {"code": "ok", "message": ""}, + }) + raise AssertionError(f"Unexpected request {request.method} {request.url}") + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(publishing_settings(tmp_path), http_client=client) + + async def persist(value: dict[str, object]) -> None: + persisted.append(dict(value)) + + try: + uploaded = await provider.upload_media( + {"access_token": "access-token"}, + { + "path": video, + "mime_type": "video/mp4", + "file_size": video.stat().st_size, + "probe": valid_probe(), + "tiktok_post_info": TikTokPostMetadata.model_validate( + valid_metadata() + ).to_post_info(), + "persist_provider_state": persist, + }, + ) + published = await provider.publish( + {"access_token": "access-token"}, {"upload": uploaded} + ) + status = await provider.get_publish_status( + {"access_token": "access-token"}, "publish-id" + ) + finally: + await client.aclose() + + assert uploaded == {"id": "publish-id"} + assert published == uploaded + assert status["status"] == "published" + assert status["metadata"]["public_post_ids"] == ["public-video-id"] + assert persisted[0] == { + "tiktok_init_started": True, + "tiktok_video_size": video.stat().st_size, + } + assert persisted[-1]["tiktok_uploaded_bytes"] == video.stat().st_size + assert calls == [ + "POST /v2/post/publish/creator_info/query/", + "POST /v2/post/publish/video/init/", + "PUT /video/session", + "POST /v2/post/publish/status/fetch/", + ] + + +async def test_media_validation_rejects_incompatible_video_before_provider_call( + tmp_path: Path, +) -> None: + video = tmp_path / "video.avi" + video.write_bytes(b"invalid") + client = httpx.AsyncClient( + transport=httpx.MockTransport( + lambda request: pytest.fail(f"Unexpected provider call {request.url}") + ) + ) + provider = TikTokProvider(publishing_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialMediaInvalidError): + await provider.validate_media({ + "path": video, + "mime_type": "video/x-msvideo", + "file_size": video.stat().st_size, + "probe": { + **valid_probe(), + "container": "avi", + "video_streams": [{"codec": "mpeg4"}], + }, + }) + finally: + await client.aclose() + + +@pytest.mark.parametrize( + ("provider_status", "expected"), + [ + ("PROCESSING_UPLOAD", "processing"), + ("PROCESSING_DOWNLOAD", "processing"), + ("SEND_TO_USER_INBOX", "processing"), + ("PUBLISH_COMPLETE", "published"), + ("FAILED", "failed"), + ("UNKNOWN_PROVIDER_STATE", "unavailable"), + ], +) +async def test_tiktok_status_reconciliation_normalizes_official_states( + tmp_path: Path, provider_status: str, expected: str +) -> None: + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={ + "data": {"status": provider_status, "fail_reason": "internal"}, + "error": {"code": "ok"}, + }) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(publishing_settings(tmp_path), http_client=client) + try: + result = await provider.get_publish_status( + {"access_token": "access-token"}, "publish-id" + ) + finally: + await client.aclose() + assert result["status"] == expected + + +async def test_tiktok_metadata_and_chunk_planning_enforce_current_contract( + tmp_path: Path, +) -> None: + with pytest.raises(ValidationError): + TikTokPostMetadata.model_validate( + valid_metadata(music_usage_confirmed=False) + ) + with pytest.raises(ValidationError): + TikTokPostMetadata.model_validate( + valid_metadata(title="\U0001f600" * 1101) + ) + provider = TikTokProvider(publishing_settings(tmp_path)) + try: + assert provider._chunk_plan(4_000_000) == (4_000_000, 1) + assert provider._chunk_plan(70_000_000) == (5_000_000, 14) + finally: + await provider.close() + + +async def test_unknown_init_outcome_never_creates_a_second_tiktok_post( + tmp_path: Path, +) -> None: + video = tmp_path / "video.mp4" + video.write_bytes(b"video") + init_calls = 0 + durable_state: dict[str, object] = {} + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal init_calls + if request.url.path.endswith("/creator_info/query/"): + return httpx.Response(200, json={ + "data": { + "privacy_level_options": ["SELF_ONLY"], + "comment_disabled": False, + "duet_disabled": False, + "stitch_disabled": False, + "max_video_post_duration_sec": 300, + }, + "error": {"code": "ok"}, + }) + if request.url.path.endswith("/video/init/"): + init_calls += 1 + return httpx.Response(200, json={ + "data": { + "publish_id": "accepted-but-not-durable", + "upload_url": "https://open-upload.tiktokapis.com/video/session", + }, + "error": {"code": "ok"}, + }) + raise AssertionError("No upload is safe after provider-state persistence fails") + + async def fail_after_marker(value: dict[str, object]) -> None: + if "tiktok_publish_id" in value: + raise RuntimeError("simulated database outage") + durable_state.clear() + durable_state.update(value) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = TikTokProvider(publishing_settings(tmp_path), http_client=client) + media = { + "path": video, + "mime_type": "video/mp4", + "file_size": video.stat().st_size, + "probe": valid_probe(), + "tiktok_post_info": TikTokPostMetadata.model_validate( + valid_metadata() + ).to_post_info(), + "persist_provider_state": fail_after_marker, + } + try: + with pytest.raises(RuntimeError): + await provider.upload_media( + {"access_token": "access-token"}, media + ) + with pytest.raises(SocialPublishFailedError) as raised: + await provider.upload_media( + {"access_token": "access-token"}, + {**media, "provider_state": durable_state}, + ) + finally: + await client.aclose() + assert init_calls == 1 + assert "duplicate publishing was prevented" in str(raised.value) + + +@pytest.mark.parametrize( + ("status_code", "error_code", "exception_type"), + [ + (401, "access_token_expired", SocialReauthRequiredError), + (403, "scope_not_authorized", SocialPermissionDeniedError), + (429, "rate_limit_exceeded", SocialRateLimitedError), + (500, "internal_error", SocialProviderUnavailableError), + (400, "invalid_file_upload", SocialMediaInvalidError), + ], +) +async def test_tiktok_error_normalization_is_safe_and_retry_classifiable( + tmp_path: Path, + status_code: int, + error_code: str, + exception_type: type[Exception], +) -> None: + secret = "token-that-must-not-leak" + client = httpx.AsyncClient(transport=httpx.MockTransport( + lambda request: httpx.Response( + status_code, + json={"error": {"code": error_code, "message": secret}}, + ) + )) + provider = TikTokProvider(publishing_settings(tmp_path), http_client=client) + try: + with pytest.raises(exception_type) as raised: + await provider.get_publish_options({"access_token": secret}) + finally: + await client.aclose() + assert secret not in str(raised.value) + + +async def test_tiktok_worker_lifecycle_idempotency_scope_and_workspace_isolation( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + workspace = "workspace-tiktok" + other_workspace = "workspace-other" + request_id = str(uuid4()) + output = container.settings.output_dir / request_id + output.mkdir(parents=True, exist_ok=True) + video = output / "video.mp4" + video.write_bytes(b"video") + account = await container.social.accounts.repository.create(SocialAccount( + workspace_id=workspace, + provider="tiktok", + account_type="creator", + external_account_id="creator-open-id", + display_name="Creator", + status="connected", + )) + asset = await container.social.media_assets.repository.create(SocialMediaAsset( + workspace_id=workspace, + request_id=request_id, + filename=video.name, + mime_type="video/mp4", + file_size=video.stat().st_size, + metadata_json=valid_probe(), + )) + await container.social.accounts.tokens.store( + workspace, + account.id, + {"access_token": "encrypted-token", "refresh_token": "encrypted-refresh"}, + scopes=["user.info.basic", "video.publish"], + ) + adapter = container.social.accounts.providers.get("tiktok") + statuses = iter([ + {"id": "publish-id", "status": "processing", "metadata": {"provider_status": "PROCESSING_UPLOAD"}}, + {"id": "publish-id", "status": "published", "metadata": {"provider_status": "PUBLISH_COMPLETE"}}, + ]) + monkeypatch.setattr(adapter, "validate_media", AsyncMock(return_value=None)) + monkeypatch.setattr(adapter, "upload_media", AsyncMock(return_value={"id": "publish-id"})) + monkeypatch.setattr(adapter, "publish", AsyncMock(return_value={"id": "publish-id"})) + monkeypatch.setattr(adapter, "get_publish_status", AsyncMock(side_effect=lambda *_: next(statuses))) + + async def resolve(*_: object, **__: object) -> dict[str, object]: + return { + "path": video, + "mime_type": "video/mp4", + "file_size": video.stat().st_size, + "probe": valid_probe(), + } + + monkeypatch.setattr(container.social.media_assets, "resolve_for_publish", resolve) + payload = SocialPostCreate.model_validate({ + "media_asset_id": asset.id, + "publish_mode": "now", + "targets": [{ + "social_account_id": account.id, + "caption": {"caption": "TikTok caption"}, + "tiktok": valid_metadata(), + }], + }) + try: + post = await container.social.publishing.create( + workspace_id=workspace, + user_id="user", + payload=payload, + idempotency_key="one-logical-publish", + ) + replay = await container.social.publishing.create( + workspace_id=workspace, + user_id="user", + payload=payload, + idempotency_key="one-logical-publish", + ) + assert replay.id == post.id + with pytest.raises(SocialIdempotencyConflictError): + await container.social.publishing.create( + workspace_id=workspace, + user_id="user", + payload=SocialPostCreate.model_validate({ + **payload.model_dump(mode="json"), + "targets": [{ + **payload.targets[0].model_dump(mode="json"), + "tiktok": valid_metadata(title="different"), + }], + }), + idempotency_key="one-logical-publish", + ) + job = (await container.social.jobs.repository.list_for_post( + workspace, post.id + ))[0] + worker = SocialPublisher(container.social) + await worker.process(workspace, job.id) + processing = await container.social.jobs.get(workspace, job.id) + assert processing.status == "publishing" + await worker.process(workspace, job.id) + published = await container.social.jobs.get(workspace, job.id) + assert published.status == "published" + assert adapter.upload_media.await_count == 1 + with pytest.raises(SocialPostNotFoundError): + await container.social.publishing.get(other_workspace, post.id) + with pytest.raises(SocialAccountNotFoundError): + await container.social.publishing.publish_options( + other_workspace, account.id + ) + with pytest.raises(SocialMediaInvalidError): + await container.social.media_assets.repository.get( + other_workspace, asset.id + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_tiktok_publish_requires_explicit_video_publish_scope( + tmp_path: Path, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + account = await container.social.accounts.repository.create(SocialAccount( + workspace_id="workspace", + provider="tiktok", + account_type="creator", + external_account_id="open-id", + status="connected", + )) + await container.social.accounts.tokens.store( + "workspace", + account.id, + {"access_token": "foundation-only"}, + scopes=["user.info.basic"], + ) + try: + with pytest.raises(SocialPermissionDeniedError): + await container.social.publishing.publish_options( + "workspace", account.id + ) + finally: + await container.social.close() + await container.security_database.close() diff --git a/tests/test_unified_publishing.py b/tests/test_unified_publishing.py new file mode 100644 index 0000000000000000000000000000000000000000..ad71e4326f082b879ab8549a391b1b9e550c620e --- /dev/null +++ b/tests/test_unified_publishing.py @@ -0,0 +1,75 @@ +from pathlib import Path + +import pytest +from pydantic import ValidationError + +from app.copilot.actions import ACTION_DEFINITIONS +from app.social.schemas.posts import SocialPostCreate, SocialPostValidation + + +def test_canonical_copy_and_project_provenance_are_strict() -> None: + payload = SocialPostCreate.model_validate({ + "project_id": "123e4567-e89b-12d3-a456-426614174000", + "caption": "Release update", + "hashtags": ["#release", "release", "media"], + "targets": [{ + "social_account_id": "x-account", + "caption": {"text": "X override"}, + "x": {"text": "X override"}, + }], + }) + assert payload.caption == "Release update" + assert payload.hashtags == ["release", "media"] + with pytest.raises(ValidationError): + SocialPostCreate.model_validate({ + "targets": [{ + "social_account_id": "x-account", + "caption": {}, + "x": {"text": "valid"}, + "unknown_provider_payload": {}, + }], + }) + + +def test_structured_validation_is_per_target() -> None: + result = SocialPostValidation.model_validate({ + "post_id": "post-1", + "valid": False, + "targets": [{ + "target_id": "target-1", + "provider": "youtube", + "account_id": "account-1", + "valid": False, + "errors": [{"code": "SOCIAL_MEDIA_INVALID", "message": "Invalid media."}], + "warnings": [], + }], + }) + assert not result.valid + assert result.targets[0].errors[0].code == "SOCIAL_MEDIA_INVALID" + + +def test_copilot_external_publishing_actions_require_confirmation() -> None: + definitions = { + item.type: item for item in ACTION_DEFINITIONS + if item.type.startswith("publishing.") + } + assert set(definitions) == { + "publishing.validate", + "publishing.create_post", + "publishing.schedule", + "publishing.publish", + "publishing.cancel", + } + for name in ("publishing.schedule", "publishing.publish", "publishing.cancel"): + assert definitions[name].external_side_effect + assert definitions[name].requires_confirmation + + +def test_additive_migration_extends_existing_social_tables() -> None: + sql = Path("app/social/migrations/0008_unified_publishing.sql").read_text() + lowered = sql.lower() + assert "alter table social_posts" in lowered + assert "foreign key (project_id) references projects(id)" in lowered + assert "force row level security" in lowered + assert "create table social_posts" not in lowered + assert "create table social_jobs" not in lowered diff --git a/tests/test_whisper_service.py b/tests/test_whisper_service.py new file mode 100644 index 0000000000000000000000000000000000000000..9ced8dca763fd8b1f2c95162f939f11dc012a323 --- /dev/null +++ b/tests/test_whisper_service.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from types import SimpleNamespace + +from app.services.whisper_service import WhisperService + + +class FakeWhisperModel: + def transcribe(self, path, **kwargs): + segments = iter( + [ + SimpleNamespace(id=0, start=0.0, end=1.25, text=" Hello"), + SimpleNamespace(id=1, start=1.25, end=2.0, text=" world"), + ] + ) + info = SimpleNamespace(language="en", language_probability=0.99, duration=2.0) + return segments, info + + +async def test_whisper_writes_srt_without_loading_real_model(settings, tmp_path) -> None: + service = WhisperService(settings) + service._models["tiny"] = FakeWhisperModel() + media = tmp_path / "audio.wav" + media.write_bytes(b"test") + result = await service.transcribe( + media, tmp_path / "out", model_name="tiny", output_format="srt" + ) + assert result.path is not None + content = result.path.read_text() + assert "00:00:00,000 --> 00:00:01,250" in content + assert "Hello" in content + assert result.metadata["language"] == "en" diff --git a/tests/test_x_foundation.py b/tests/test_x_foundation.py new file mode 100644 index 0000000000000000000000000000000000000000..fd003fad7b66e80503086c4491207dcb2b0ffb7b --- /dev/null +++ b/tests/test_x_foundation.py @@ -0,0 +1,499 @@ +"""Phase 5A X API v2 OAuth and account-discovery coverage. + +All provider traffic is mocked. Normal CI never needs X credentials, API +credits, or an interactive browser authorization flow. +""" + +from __future__ import annotations + +import base64 +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from pydantic import ValidationError +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialOAuthStateError, + SocialPermissionDeniedError, + SocialProviderUnavailableError, + SocialReauthRequiredError, +) +from app.social.models import OAuthState, SocialAccountToken +from app.social.providers.x import XProvider +from app.social.schemas.accounts import SocialAccountConnectRequest + +_REDIRECT_URI = "https://api.example.com/v1/social/accounts/x/callback" + + +def x_settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + database_url=f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + social_database_url=f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + social_auto_migrate=True, + social_worker_enabled=False, + social_oauth_encryption_key="phase-5a-test-encryption-material", + social_oauth_redirect_base_url="https://api.example.com", + x_client_id="x-client-id", + x_client_secret="x-client-secret", + x_redirect_uri=_REDIRECT_URI, + x_publishing_enabled=True, + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + cleanup_interval_seconds=3600, + whisper_model="tiny", + ) + + +def assert_confidential_client(request: httpx.Request) -> None: + scheme, encoded = request.headers["authorization"].split(" ", 1) + assert scheme == "Basic" + assert base64.b64decode(encoded).decode() == "x-client-id:x-client-secret" + + +async def test_x_authorization_uses_official_url_minimum_scopes_and_s256_pkce( + tmp_path: Path, +) -> None: + provider = XProvider(x_settings(tmp_path)) + try: + url = await provider.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + code_challenge="s256-code-challenge", + ) + with pytest.raises(SocialPermissionDeniedError): + await provider.get_authorization_url( + state="s" * 43, + redirect_uri=_REDIRECT_URI, + code_challenge=None, + ) + finally: + await provider.close() + + parsed = urlparse(url) + query = parse_qs(parsed.query) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == ( + "https://x.com/i/oauth2/authorize" + ) + assert query["client_id"] == ["x-client-id"] + assert query["redirect_uri"] == [_REDIRECT_URI] + assert query["response_type"] == ["code"] + assert query["scope"] == ["tweet.read users.read offline.access"] + assert query["state"] == ["s" * 43] + assert query["code_challenge"] == ["s256-code-challenge"] + assert query["code_challenge_method"] == ["S256"] + assert "tweet.write" not in query["scope"][0] + assert "media.write" not in query["scope"][0] + + +async def test_x_exchange_refresh_discovery_and_revoke_use_official_v2_endpoints( + tmp_path: Path, +) -> None: + calls: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + assert request.url.host == "api.x.com" + if request.url.path == "/2/oauth2/token": + assert_confidential_client(request) + form = parse_qs(request.content.decode()) + assert "client_secret" not in form + assert "client_id" not in form + if form["grant_type"] == ["authorization_code"]: + assert form == { + "code": ["authorization-code"], + "grant_type": ["authorization_code"], + "redirect_uri": [_REDIRECT_URI], + "code_verifier": ["pkce-verifier"], + } + else: + assert form == { + "refresh_token": ["refresh-token"], + "grant_type": ["refresh_token"], + } + return httpx.Response( + 200, + json={ + "access_token": "x-access-token", + "refresh_token": "x-rotated-refresh-token", + "expires_in": 7200, + "scope": "tweet.read users.read offline.access", + "token_type": "bearer", + }, + ) + if request.url.path == "/2/users/me": + assert request.headers["authorization"] == "Bearer x-access-token" + assert parse_qs(request.url.query.decode()) == { + "user.fields": [ + "created_at,description,profile_image_url,protected,verified" + ] + } + return httpx.Response( + 200, + json={ + "data": { + "id": "2244994945", + "username": "XDevelopers", + "name": "X Developers", + "profile_image_url": "https://pbs.twimg.com/profile.jpg", + "created_at": "2013-12-14T04:35:55.000Z", + "description": "Official developer account", + "protected": False, + "verified": True, + } + }, + ) + assert request.url.path == "/2/oauth2/revoke" + assert_confidential_client(request) + assert parse_qs(request.content.decode()) == { + "token": ["x-rotated-refresh-token"] + } + return httpx.Response(200) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(x_settings(tmp_path), http_client=client) + try: + token = await provider.exchange_code( + code="authorization-code", + redirect_uri=_REDIRECT_URI, + code_verifier="pkce-verifier", + ) + account = await provider.get_account(token) + refreshed = await provider.refresh_token( + {"access_token": "old-token", "refresh_token": "refresh-token"} + ) + await provider.revoke_token(refreshed) + finally: + await client.aclose() + + assert account == { + "external_account_id": "2244994945", + "account_type": "user", + "username": "XDevelopers", + "display_name": "X Developers", + "avatar_url": "https://pbs.twimg.com/profile.jpg", + "metadata": { + "x_user_id": "2244994945", + "created_at": "2013-12-14T04:35:55.000Z", + "verified": True, + "protected": False, + "description": "Official developer account", + }, + } + assert refreshed["refresh_token"] == "x-rotated-refresh-token" + assert calls == [ + "/2/oauth2/token", + "/2/users/me", + "/2/oauth2/token", + "/2/oauth2/revoke", + ] + + +async def test_x_invalid_code_and_pkce_failure_are_normalized_without_secrets( + tmp_path: Path, +) -> None: + secret_code = "x-code-that-must-not-leak" + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + json={ + "error": "invalid_grant", + "error_description": f"invalid code {secret_code}", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(x_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialPermissionDeniedError): + await provider.exchange_code( + code=secret_code, + redirect_uri=_REDIRECT_URI, + code_verifier=None, + ) + with pytest.raises(SocialReauthRequiredError) as raised: + await provider.exchange_code( + code=secret_code, + redirect_uri=_REDIRECT_URI, + code_verifier="incorrect-verifier", + ) + finally: + await client.aclose() + assert secret_code not in str(raised.value) + + +async def test_x_invalid_client_is_configuration_failure_not_consent_loop( + tmp_path: Path, +) -> None: + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 401, + json={ + "error": "invalid_client", + "error_description": "client secret is not accepted", + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(x_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError) as raised: + await provider.exchange_code( + code="authorization-code", + redirect_uri=_REDIRECT_URI, + code_verifier="pkce-verifier", + ) + finally: + await client.aclose() + + assert "client secret is not accepted" not in str(raised.value) + + +async def test_x_account_discovery_rejects_non_ascii_or_oversized_user_ids( + tmp_path: Path, +) -> None: + invalid_ids = ["٢٢٤٤٩٩٤٩٤٥", "12345678901234567890"] + for user_id in invalid_ids: + + async def handler(_: httpx.Request, value: str = user_id) -> httpx.Response: + return httpx.Response( + 200, + json={"data": {"id": value, "username": "invalid"}}, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(x_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError): + await provider.get_account({"access_token": "x-access-token"}) + finally: + await client.aclose() + + +async def test_x_callback_is_single_use_duplicate_safe_and_workspace_bound( + tmp_path: Path, +) -> None: + container = build_container(x_settings(tmp_path)) + await container.social.initialize() + adapter = container.social.accounts.providers.get("x") + assert isinstance(adapter, XProvider) + await adapter._client.aclose() + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/2/oauth2/token": + form = parse_qs(request.content.decode()) + assert form.get("code_verifier", [""])[0] + return httpx.Response( + 200, + json={ + "access_token": "x-token-that-must-stay-encrypted", + "refresh_token": "x-refresh-that-must-stay-encrypted", + "expires_in": 7200, + "scope": "tweet.read users.read offline.access", + "token_type": "bearer", + }, + ) + return httpx.Response( + 200, + json={ + "data": { + "id": "2244994945", + "username": "workspace_user", + "name": "Workspace User", + } + }, + ) + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + try: + first_connect = await container.social.oauth.connect( + provider="x", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(), + ) + first_query = parse_qs(urlparse(first_connect.authorization_url or "").query) + first_state = first_query["state"][0] + assert first_query["code_challenge_method"] == ["S256"] + assert first_query["code_challenge"][0] + + first = await container.social.oauth.callback( + provider="x", + state=first_state, + code="first-code", + ) + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.callback( + provider="x", + state=first_state, + code="replayed-code", + ) + + second_connect = await container.social.oauth.connect( + provider="x", + workspace_id="workspace-a", + user_id="user-a", + payload=SocialAccountConnectRequest(), + ) + second_state = parse_qs( + urlparse(second_connect.authorization_url or "").query + )["state"][0] + second = await container.social.oauth.callback( + provider="x", + state=second_state, + code="second-code", + ) + + assert first.id == second.id + accounts = await container.social.accounts.list("workspace-a") + assert [account.id for account in accounts if account.provider.value == "x"] == [ + first.id + ] + assert "x-token-that-must-stay-encrypted" not in first.model_dump_json() + with pytest.raises(SocialAccountNotFoundError): + await container.social.accounts.get("workspace-b", first.id) + + async with container.social.database.session("workspace-a") as session: + stored = await session.scalar( + select(SocialAccountToken).where( + SocialAccountToken.social_account_id == first.id + ) + ) + assert stored is not None + assert stored.expires_at is not None + assert stored.encrypted_payload + assert "x-token-that-must-stay-encrypted" not in stored.encrypted_payload + finally: + await container.social.close() + await container.security_database.close() + + +async def test_x_state_redirect_provider_and_expiry_validation(tmp_path: Path) -> None: + container = build_container(x_settings(tmp_path)) + await container.social.initialize() + try: + assert container.social.oauth._redirect_uri("x", None) == _REDIRECT_URI + with pytest.raises(SocialPermissionDeniedError): + container.social.oauth._redirect_uri( + "x", + "https://attacker.example/v1/social/accounts/x/callback", + ) + + state = await container.social.oauth.states.create( + provider="x", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=_REDIRECT_URI, + ) + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=state.state, + provider="linkedin", + ) + consumed = await container.social.oauth.states.consume( + state=state.state, + provider="x", + ) + assert consumed.workspace_id == "workspace-a" + assert consumed.user_id == "user-a" + + expired = OAuthState( + state="expired-x-state-value-that-is-long-enough", + provider="x", + workspace_id="workspace-a", + user_id="user-a", + redirect_uri=_REDIRECT_URI, + expires_at=datetime.now(timezone.utc) - timedelta(seconds=1), + ) + async with container.social.database.session("workspace-a") as session: + session.add(expired) + await session.commit() + with pytest.raises(SocialOAuthStateError): + await container.social.oauth.states.consume( + state=expired.state, + provider="x", + ) + finally: + await container.social.close() + await container.security_database.close() + + +def test_x_redirect_configuration_is_fail_closed() -> None: + invalid_redirects = [ + "https://attacker.example/not-the-x-callback", + "ftp://localhost/v1/social/accounts/x/callback", + "http://api.example.com/v1/social/accounts/x/callback", + "https://api.example.com/v1/social/accounts/x/callback?next=attacker", + ] + for redirect in invalid_redirects: + with pytest.raises(ValidationError): + Settings(_env_file=None, x_redirect_uri=redirect) + + settings = Settings( + _env_file=None, + x_redirect_uri="http://localhost/v1/social/accounts/x/callback", + ) + assert settings.x_redirect_uri.startswith("http://localhost/") + + +async def test_x_capability_discovery_advertises_implemented_publishing( + tmp_path: Path, +) -> None: + container = build_container(x_settings(tmp_path)) + try: + provider = container.social.accounts.get_provider("x") + assert provider.available + assert provider.configured + assert provider.capabilities.implementation_status == "implemented" + assert provider.capabilities.account_types == ["user"] + assert provider.capabilities.required_scopes == [ + "tweet.read", + "users.read", + "offline.access", + ] + assert provider.capabilities.video + assert provider.capabilities.video_upload + assert provider.capabilities.video_status + assert provider.capabilities.image + assert provider.capabilities.direct_publish + assert not provider.capabilities.draft_upload + assert provider.capabilities.scheduled_publish + assert not provider.capabilities.native_scheduling + assert provider.capabilities.delete_post + assert provider.capabilities.publishing_required_scopes == [ + "tweet.write", + "media.write", + ] + assert provider.capabilities.analytics + assert provider.capabilities.analytics_required_scopes == ["tweet.read"] + finally: + await container.social.close() + await container.security_database.close() + + +async def test_x_publishing_capabilities_are_fail_closed_without_operator_gate( + tmp_path: Path, +) -> None: + settings = x_settings(tmp_path).model_copy( + update={"x_publishing_enabled": False} + ) + provider = XProvider(settings) + try: + assert provider.configuration_ready + assert not provider.publishing_ready + assert not provider.capabilities.direct_publish + assert not provider.capabilities.video_upload + assert not provider.capabilities.delete_post + assert provider.capabilities.publishing_required_scopes == [] + finally: + await provider.close() diff --git a/tests/test_x_live.py b/tests/test_x_live.py new file mode 100644 index 0000000000000000000000000000000000000000..0d4aa3f7132c8140a0be3b2f5f296273b205c30d --- /dev/null +++ b/tests/test_x_live.py @@ -0,0 +1,177 @@ +"""Opt-in, destructive X API v2 integration verification. + +Normal CI always skips this module. Run only against a dedicated X account; +the test requires explicit permission to publish and delete its own posts. +Provider credentials are read from the process environment and never logged. +""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import os +import secrets +from pathlib import Path + +import pytest + + +pytestmark = pytest.mark.skipif( + os.getenv("RUN_X_INTEGRATION_TESTS", "").lower() != "true", + reason=( + "X live integration is NOT VERIFIED; set RUN_X_INTEGRATION_TESTS=true " + "with dedicated test-account credentials." + ), +) + + +def _require_live_configuration() -> None: + required = ( + "X_CLIENT_ID", + "X_CLIENT_SECRET", + "X_REDIRECT_URI", + "X_LIVE_TEST_ACCESS_TOKEN", + ) + missing = [name for name in required if not os.getenv(name)] + if missing: + pytest.skip(f"X live integration is NOT VERIFIED; missing {', '.join(missing)}") + if os.getenv("X_LIVE_TEST_ALLOW_PUBLISH", "").lower() != "true": + pytest.skip("Set X_LIVE_TEST_ALLOW_PUBLISH=true to create dedicated test posts.") + if os.getenv("X_LIVE_TEST_DELETE", "").lower() != "true": + pytest.skip("Set X_LIVE_TEST_DELETE=true to require cleanup of test posts.") + + +def test_x_live_configuration_requires_explicit_destructive_consent() -> None: + assert os.getenv("RUN_X_INTEGRATION_TESTS", "").lower() == "true" + _require_live_configuration() + + +async def test_x_live_oauth_discovery_publish_status_analytics_and_delete() -> None: + """Exercise official endpoints with a dedicated, approved X project. + + Browser consent remains an operator action. If a fresh authorization code + and its PKCE verifier are supplied, this test also performs the live code + exchange. Otherwise OAuth exchange remains NOT VERIFIED even though the + authorization URL/PKCE contract and bearer-token lifecycle are exercised. + """ + + _require_live_configuration() + from app.core.config import Settings + from app.services.ffprobe_service import FFprobeService + from app.services.validator import MediaValidator + from app.social.providers.x import XProvider + + settings = Settings( + _env_file=None, + auth_enabled=False, + x_client_id=os.environ["X_CLIENT_ID"], + x_client_secret=os.environ["X_CLIENT_SECRET"], + x_redirect_uri=os.environ["X_REDIRECT_URI"], + x_publishing_enabled=True, + whisper_model="tiny", + ) + provider = XProvider(settings) + token: dict[str, object] = { + "access_token": os.environ["X_LIVE_TEST_ACCESS_TOKEN"] + } + if refresh := os.getenv("X_LIVE_TEST_REFRESH_TOKEN"): + token["refresh_token"] = refresh + token = {**token, **await provider.refresh_token(token)} + if code := os.getenv("X_LIVE_TEST_AUTHORIZATION_CODE"): + verifier = os.getenv("X_LIVE_TEST_PKCE_VERIFIER") + if not verifier: + pytest.skip( + "X_LIVE_TEST_PKCE_VERIFIER is required with a live authorization code." + ) + token = await provider.exchange_code( + code=code, + redirect_uri=os.environ["X_REDIRECT_URI"], + code_verifier=verifier, + ) + + created_ids: list[str] = [] + + async def persist(_: dict[str, object]) -> None: + return None + + try: + verifier = secrets.token_urlsafe(64) + challenge = base64.urlsafe_b64encode( + hashlib.sha256(verifier.encode("ascii")).digest() + ).rstrip(b"=").decode("ascii") + authorization_url = await provider.get_authorization_url( + state="phase5c-live-state-value-that-is-long-enough", + redirect_uri=os.environ["X_REDIRECT_URI"], + code_challenge=challenge, + ) + assert authorization_url.startswith("https://x.com/i/oauth2/authorize?") + account = await provider.get_account(token) + assert account["external_account_id"] + + text = await provider.publish( + token, + { + "x_post_metadata": { + "text": "MediaRouter Phase 5C live integration verification" + }, + "upload": {"identity_type": "none"}, + "provider_state": {"x_post_submission_attempted": False}, + "persist_provider_state": persist, + }, + ) + created_ids.append(str(text["id"])) + + media_value = os.getenv("X_LIVE_TEST_MEDIA_PATH") + if media_value: + media_path = Path(media_value).expanduser().resolve() + if not media_path.is_file(): + pytest.skip("X_LIVE_TEST_MEDIA_PATH is not a readable test asset.") + probe = await FFprobeService(settings).probe(media_path) + mime_type = MediaValidator(settings).infer_mime(media_path) + media_state: dict[str, object] = {} + + async def persist_media(value: dict[str, object]) -> None: + media_state.clear() + media_state.update(value) + + media = { + "path": media_path, + "mime_type": mime_type, + "file_size": media_path.stat().st_size, + "probe": probe, + "provider_state": media_state, + "persist_provider_state": persist_media, + } + await provider.validate_media(media) + uploaded = await provider.upload_media(token, media) + published = await provider.publish( + token, + { + "x_post_metadata": { + "text": "MediaRouter Phase 5C media verification" + }, + "upload": uploaded, + "provider_state": { + **media_state, + "x_post_submission_attempted": False, + }, + "persist_provider_state": persist_media, + }, + ) + created_ids.append(str(published["id"])) + + for external_id in created_ids: + status: dict[str, object] | None = None + for _ in range(12): + status = await provider.get_publish_status(token, external_id) + if status.get("status") == "published": + break + await asyncio.sleep(5) + assert status is not None and status.get("status") == "published" + metrics = await provider.get_metrics(token, external_id) + assert metrics["status"] in {"available", "unavailable"} + finally: + for external_id in reversed(created_ids): + await provider.delete_post(token, external_id) + await provider.close() diff --git a/tests/test_x_production.py b/tests/test_x_production.py new file mode 100644 index 0000000000000000000000000000000000000000..95c339a4ad1f8f3676b82237d093238ed9cb56d3 --- /dev/null +++ b/tests/test_x_production.py @@ -0,0 +1,711 @@ +"""Phase 5C X analytics, security, tenancy, and certification coverage. + +Normal CI uses SQLite and mocked official X API v2 traffic. Destructive live +provider traffic is isolated in test_x_live.py and requires explicit opt-in. +""" + +from __future__ import annotations + +import json +import logging +from datetime import datetime, timedelta, timezone +from pathlib import Path +from types import SimpleNamespace +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from sqlalchemy import select + +from app.container import build_container +from app.core.config import Settings +from app.core.logger import JsonFormatter +from app.mcp.registry import MCPRegistry +from app.mcp.server import create_mcp_server +from app.security.context import AuthContext, auth_context, http_auth_applied +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialJobNotFoundError, + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialPublishFailedError, + SocialRateLimitedError, + SocialReauthRequiredError, +) +from app.social.domain.retry import classify_retry +from app.social.models import ( + SocialAccount, + SocialAuditEvent, + SocialJob, + SocialMediaAsset, + SocialPost, + SocialPostMetric, + SocialPostTarget, +) +from app.social.providers.x import XProvider +from app.social.schemas.accounts import SocialAccountConnectRequest, SocialAccountView +from app.social.schemas.jobs import SocialJobView +from app.social.workers.publisher import SocialPublisher + + +def phase5c_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "phase-5c-test-encryption-material", + "social_oauth_redirect_base_url": "https://api.example.com", + "x_client_id": "x-client-id", + "x_client_secret": "x-client-secret", + "x_redirect_uri": "https://api.example.com/v1/social/accounts/x/callback", + "x_publishing_enabled": True, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +@pytest.fixture +async def phase5c_container(tmp_path: Path): + container = build_container(phase5c_settings(tmp_path)) + await container.social.initialize() + try: + yield container + finally: + await container.social.close() + await container.security_database.close() + + +async def _connected_account(container: object, workspace_id: str) -> SocialAccount: + social = container.social # type: ignore[attr-defined] + account = await social.accounts.repository.create( + SocialAccount( + workspace_id=workspace_id, + provider="x", + account_type="user", + external_account_id="2244994945", + username="production_user", + display_name="Production User", + status="connected", + ) + ) + await social.accounts.tokens.store( + workspace_id, + account.id, + {"access_token": "encrypted-x-token", "refresh_token": "encrypted-refresh"}, + expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + scopes=[ + "tweet.read", + "users.read", + "offline.access", + "tweet.write", + "media.write", + ], + token_type="bearer", + ) + return account + + +async def test_x_public_metrics_use_official_post_lookup_and_normalize_without_guessing( + tmp_path: Path, +) -> None: + secret = "x-analytics-token-that-must-not-leak" + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + assert request.url.path == "/2/tweets/1900000000000000001" + assert request.headers["authorization"] == f"Bearer {secret}" + assert secret not in str(request.url) + query = parse_qs(request.url.query.decode()) + assert query == { + "tweet.fields": ["author_id,created_at,public_metrics"], + "expansions": ["attachments.media_keys"], + "media.fields": ["media_key,type,public_metrics"], + } + return httpx.Response( + 200, + json={ + "data": { + "id": "1900000000000000001", + "author_id": "2244994945", + "created_at": "2026-08-01T01:02:03.000Z", + "public_metrics": { + "bookmark_count": 7, + "impression_count": 101, + "like_count": 22, + "quote_count": 5, + "reply_count": 3, + "retweet_count": 4, + }, + }, + "includes": { + "media": [{ + "media_key": "7_1900000000000000000", + "type": "video", + "public_metrics": {"view_count": 88}, + "access_token": secret, + }] + }, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(phase5c_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics( + {"access_token": secret}, "1900000000000000001" + ) + finally: + await client.aclose() + + assert result["status"] == "available" + assert result["views"] == 88 + assert result["impressions"] == 101 + assert result["likes"] == 22 + assert result["comments"] == 3 + assert result["shares"] == 4 + assert result["published_at"] == "2026-08-01T01:02:03.000Z" + assert result["raw_metrics"]["post"]["quote_count"] == 5 + assert result["raw_metrics"]["post"]["bookmark_count"] == 7 + assert secret not in json.dumps(result) + + +@pytest.mark.parametrize( + ("response", "reason"), + [ + (httpx.Response(404, json={"title": "Not Found"}), "X_POST_NOT_AVAILABLE_TO_AUTHORIZED_USER"), + ( + httpx.Response(200, json={"data": {"id": "1900000000000000002"}}), + "X_PUBLIC_METRICS_UNAVAILABLE", + ), + ], +) +async def test_x_missing_analytics_are_unavailable_not_fabricated( + tmp_path: Path, response: httpx.Response, reason: str +) -> None: + client = httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: response) + ) + provider = XProvider(phase5c_settings(tmp_path), http_client=client) + try: + result = await provider.get_metrics( + {"access_token": "provider-token"}, "1900000000000000002" + ) + with pytest.raises(SocialPublishFailedError): + await provider.get_metrics({"access_token": "provider-token"}, "invalid") + finally: + await client.aclose() + assert result == {"status": "unavailable", "reason": reason} + assert not any(name in result for name in ("views", "impressions", "likes")) + + +async def test_x_analytics_capability_scope_and_snapshot_persistence( + phase5c_container, +) -> None: + social = phase5c_container.social + provider = social.accounts.providers.get("x") + assert provider.capabilities.analytics + assert provider.capabilities.analytics_required_scopes == ["tweet.read"] + + normal = await social.oauth.connect( + provider="x", + workspace_id="workspace-x", + user_id="user-x", + payload=SocialAccountConnectRequest(), + ) + analytics = await social.oauth.connect( + provider="x", + workspace_id="workspace-x", + user_id="user-x", + payload=SocialAccountConnectRequest(authorization_purpose="analytics"), + ) + normal_scopes = set( + parse_qs(urlparse(normal.authorization_url or "").query)["scope"][0].split() + ) + analytics_scopes = set( + parse_qs(urlparse(analytics.authorization_url or "").query)["scope"][0].split() + ) + assert normal_scopes == analytics_scopes + assert "tweet.read" in normal_scopes + + account = await _connected_account(phase5c_container, "workspace-x") + post, targets = await social.publishing.posts.create( + SocialPost(workspace_id="workspace-x", media_asset_id=None), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="x", + status="published", + external_post_id="1900000000000000003", + ) + ], + ) + + async def metrics(_: dict[str, object], post_id: str) -> dict[str, object]: + assert post_id == "1900000000000000003" + return { + "status": "available", + "impressions": 40, + "likes": 8, + "comments": 2, + "shares": 3, + "raw_metrics": { + "post": {"quote_count": 1}, + "authorization": "Bearer secret-that-must-not-persist", + }, + } + + provider.get_metrics = metrics # type: ignore[method-assign] + result = await social.analytics.post("workspace-x", post.id) + assert result["metrics"][0]["impressions"] == 40 + assert result["metrics"][0]["raw_metrics"] == {"post": {"quote_count": 1}} + assert "secret-that-must-not-persist" not in str(result) + + async with social.database.session("workspace-x") as session: + record = await session.scalar( + select(SocialPostMetric).where( + SocialPostMetric.social_post_target_id == targets[0].id + ) + ) + assert record is not None + assert record.provider == "x" + assert record.raw_metrics == {"post": {"quote_count": 1}} + + +async def test_x_analytics_capability_is_fail_closed_without_oauth_configuration( + tmp_path: Path, +) -> None: + provider = XProvider( + phase5c_settings(tmp_path, x_client_id="", x_client_secret=None) + ) + try: + assert not provider.configuration_ready + assert not provider.capabilities.analytics + assert provider.capabilities.analytics_required_scopes == [] + finally: + await provider.close() + + +async def test_x_cross_workspace_accounts_posts_targets_jobs_assets_and_analytics_fail( + phase5c_container, +) -> None: + social = phase5c_container.social + account_b = await _connected_account(phase5c_container, "workspace-b") + asset_b = await social.media_assets.repository.create( + SocialMediaAsset( + workspace_id="workspace-b", + request_id="00000000-0000-0000-0000-00000000000b", + filename="post.mp4", + mime_type="video/mp4", + file_size=10, + ) + ) + post_b, targets_b = await social.publishing.posts.create( + SocialPost(workspace_id="workspace-b", media_asset_id=asset_b.id), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account_b.id, + provider="x", + status="published", + external_post_id="1900000000000000004", + ) + ], + ) + job_b = ( + await social.jobs.repository.create_many( + [ + SocialJob( + workspace_id="workspace-b", + social_post_id=post_b.id, + social_post_target_id=targets_b[0].id, + provider="x", + status="queued", + idempotency_key="workspace-b-x-job", + ) + ] + ) + )[0] + + with pytest.raises(SocialAccountNotFoundError): + await social.accounts.get("workspace-a", account_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.get("workspace-a", post_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.publishing.posts.set_target_status( + "workspace-a", targets_b[0].id, "failed" + ) + with pytest.raises(SocialJobNotFoundError): + await social.jobs.get("workspace-a", job_b.id) + with pytest.raises(SocialMediaInvalidError): + await social.media_assets.repository.get("workspace-a", asset_b.id) + with pytest.raises(SocialAccountNotFoundError): + await social.analytics.account("workspace-a", account_b.id) + with pytest.raises(SocialPostNotFoundError): + await social.analytics.post("workspace-a", post_b.id) + + migrations = Path("app/social/migrations") + assets_sql = (migrations / "0004_youtube_media_assets.sql").read_text() + metrics_sql = (migrations / "0002_social_rls.sql").read_text() + assert "alter table social_media_assets enable row level security" in assets_sql + assert "create policy social_workspace_isolation on social_media_assets" in assets_sql + assert "'social_post_metrics'" in metrics_sql + + +async def test_x_tokens_are_redacted_and_revoked_credentials_require_reauthorization( + phase5c_container, +) -> None: + secret = "phase-5c-secret-token" + social = phase5c_container.social + account = await social.accounts.repository.create( + SocialAccount( + workspace_id="workspace-security", + provider="x", + account_type="user", + external_account_id="2244994946", + status="connected", + metadata_json={ + "display": "X user", + "access_token": secret, + "message": f"Authorization: Bearer {secret}", + }, + ) + ) + await social.accounts.tokens.store( + "workspace-security", + account.id, + {"access_token": secret, "refresh_token": secret}, + scopes=["tweet.read"], + ) + post, targets = await social.publishing.posts.create( + SocialPost(workspace_id="workspace-security", media_asset_id=None), + [ + SocialPostTarget( + social_post_id="", + social_account_id=account.id, + provider="x", + status="published", + external_post_id="1900000000000000005", + ) + ], + ) + job = ( + await social.jobs.repository.create_many( + [ + SocialJob( + workspace_id="workspace-security", + social_post_id=post.id, + social_post_target_id=targets[0].id, + provider="x", + status="queued", + payload_json={"access_token": secret, "message": f"Bearer {secret}"}, + provider_state_encrypted=f"encrypted:{secret}", + ) + ] + ) + )[0] + assert secret not in SocialAccountView.from_record(account).model_dump_json() + assert secret not in SocialJobView.from_record(job).model_dump_json() + + record = logging.LogRecord( + "x-security-test", + logging.ERROR, + __file__, + 1, + f"provider failed Authorization: Bearer {secret}", + (), + None, + ) + record.provider_payload = { + "refresh_token": secret, + "message": f"access_token={secret}", + } + rendered = JsonFormatter().format(record) + assert secret not in rendered + assert "[REDACTED]" in rendered + + await social.audit.record( + workspace_id="workspace-security", + event_type="SOCIAL_X_SECURITY_TEST", + provider="x", + metadata={ + "client_secret": secret, + "message": f"Authorization: Bearer {secret}", + }, + ) + async with social.database.session("workspace-security") as session: + audit = await session.scalar( + select(SocialAuditEvent).where( + SocialAuditEvent.event_type == "SOCIAL_X_SECURITY_TEST" + ) + ) + assert audit is not None + assert secret not in json.dumps(audit.metadata_json) + + await social.accounts.tokens.revoke("workspace-security", account.id) + with pytest.raises(SocialReauthRequiredError): + await social.accounts.tokens.retrieve("workspace-security", account.id) + with pytest.raises(SocialReauthRequiredError): + await social.analytics.post("workspace-security", post.id) + + +@pytest.mark.parametrize( + ("status_code", "error_type", "retryable"), + [ + (429, SocialRateLimitedError, True), + (500, SocialProviderUnavailableError, True), + (502, SocialProviderUnavailableError, True), + (503, SocialProviderUnavailableError, True), + (504, SocialProviderUnavailableError, True), + (401, SocialReauthRequiredError, True), + (403, SocialPermissionDeniedError, False), + (400, SocialPublishFailedError, False), + ], +) +def test_x_retry_matrix_is_bounded_and_permanent_errors_fail( + status_code: int, + error_type: type[Exception], + retryable: bool, +) -> None: + response = httpx.Response(status_code, json={"title": "provider detail"}) + with pytest.raises(error_type) as raised: + XProvider._raise_x_error(response, operation="production audit") + decision = classify_retry( + status_code=getattr(raised.value, "status_code", status_code), attempt=1 + ) + assert decision.retryable is retryable + if status_code == 401: + assert decision.refresh_token_first + assert not classify_retry(status_code=401, attempt=2).retryable + if retryable: + assert not classify_retry(status_code=status_code, attempt=10).refresh_token_first + + +async def test_x_transient_retry_stops_at_the_job_attempt_ceiling() -> None: + transitions: list[str] = [] + + class Jobs: + async def complete_attempt(self, *_: object, **__: object) -> None: + return None + + async def transition( + self, _: str, __: str, status: str, **___: object + ) -> SocialJob: + transitions.append(status) + return job + + class Audit: + async def record(self, **_: object) -> None: + return None + + job = SocialJob( + id="job-at-limit", + workspace_id="workspace-x", + social_post_id="post-at-limit", + social_post_target_id=None, + provider="x", + status="publishing", + attempt_count=5, + max_attempts=5, + ) + social = SimpleNamespace( + jobs=SimpleNamespace(repository=Jobs()), + audit=Audit(), + ) + publisher = SocialPublisher(social) # type: ignore[arg-type] + + await publisher._handle_failure( + "workspace-x", + job, + "attempt-at-limit", + SocialProviderUnavailableError("temporary X failure"), + ) + + assert transitions == ["failed"] + + +async def test_x_timeout_is_secret_safe_and_retryable(tmp_path: Path) -> None: + secret = "x-timeout-token-that-must-not-leak" + + async def handler(request: httpx.Request) -> httpx.Response: + raise httpx.ReadTimeout(f"Authorization: Bearer {secret}", request=request) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(phase5c_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialProviderUnavailableError) as raised: + await provider.get_metrics( + {"access_token": secret}, "1900000000000000006" + ) + finally: + await client.aclose() + assert secret not in str(raised.value) + assert classify_retry(status_code=raised.value.status_code, attempt=1).retryable + + +async def test_x_status_reconciliation_normalizes_published_deleted_and_unavailable( + tmp_path: Path, +) -> None: + secret = "status-token-that-must-not-leak" + + async def handler(request: httpx.Request) -> httpx.Response: + external_id = request.url.path.rsplit("/", 1)[-1] + if external_id.endswith("1"): + return httpx.Response( + 200, + json={"data": {"id": external_id, "text": "published", "token": secret}}, + ) + if external_id.endswith("2"): + return httpx.Response(404, json={"title": "Not Found"}) + return httpx.Response(200, json={"meta": {"result_count": 0}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(phase5c_settings(tmp_path), http_client=client) + try: + published = await provider.get_publish_status( + {"access_token": secret}, "1900000000000000011" + ) + deleted = await provider.get_publish_status( + {"access_token": secret}, "1900000000000000012" + ) + unavailable = await provider.get_publish_status( + {"access_token": secret}, "1900000000000000013" + ) + finally: + await client.aclose() + assert published["status"] == "published" + assert deleted == { + "id": "1900000000000000012", + "status": "deleted", + "metadata": {"reason": "not_found"}, + } + assert unavailable["status"] == "unavailable" + assert secret not in json.dumps(published) + + +async def test_x_provider_success_then_backend_loss_reconciles_without_second_post( + tmp_path: Path, +) -> None: + create_calls = 0 + state: dict[str, object] = {"x_post_submission_attempted": False} + + async def persist(value: dict[str, object]) -> None: + state.clear() + state.update(value) + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal create_calls + if request.method == "POST": + create_calls += 1 + return httpx.Response( + 201, + json={ + "data": { + "id": "1900000000000000014", + "text": "Recovered after local persistence loss", + } + }, + ) + assert request.method == "GET" + assert request.url.path == "/2/users/2244994945/tweets" + return httpx.Response( + 200, + json={ + "data": [{ + "id": "1900000000000000014", + "text": "Recovered after local persistence loss", + "created_at": datetime.now(timezone.utc).isoformat(), + }], + "meta": {"result_count": 1}, + }, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(phase5c_settings(tmp_path), http_client=client) + payload = { + "provider_account_id": "2244994945", + "provider_state": state, + "persist_provider_state": persist, + "x_post_metadata": {"text": "Recovered after local persistence loss"}, + "upload": {"identity_type": "none"}, + } + try: + accepted = await provider.publish( + {"access_token": "provider-token"}, payload + ) + assert accepted["id"] == "1900000000000000014" + # Simulate loss of the accepted result before target persistence. The + # encrypted provider marker survives and is the restart boundary. + recovered = await provider.reconcile_pending_publish( + {"access_token": "provider-token"}, + {**payload, "provider_state": state}, + ) + finally: + await client.aclose() + + assert recovered is not None + assert recovered["id"] == "1900000000000000014" + assert recovered["metadata"] == {"reconciled": True} + assert create_calls == 1 + + +async def test_mcp_registers_phase5_contract_and_enforces_analytics_scope( + phase5c_container, +) -> None: + server = create_mcp_server(phase5c_container) + tools = {tool.name for tool in await server.list_tools()} + assert { + "social.list_providers", + "social.get_capabilities", + "social.list_accounts", + "social.create_post", + "social.publish_post", + "social.schedule_post", + "social.get_job", + "social.get_analytics", + } <= tools + + context = AuthContext( + api_key_id="workspace-x", + key_name="phase-5c", + key_prefix="mp_test", + environment="test", + role="viewer", + scopes=frozenset({"social:accounts:read"}), + requests_per_minute=100, + concurrent_jobs=2, + uploads_per_hour=10, + processing_bytes_per_day=1_000_000, + expires_at=None, + ) + auth_token = auth_context.set(context) + http_token = http_auth_applied.set(True) + called = False + + async def forbidden_action() -> dict[str, object]: + nonlocal called + called = True + return {"metrics": [], "access_token": "must-not-appear"} + + try: + result = await MCPRegistry(phase5c_container).run_metadata_tool( + "social.get_analytics", + forbidden_action, + required_scope="social:analytics:read", + ) + finally: + http_auth_applied.reset(http_token) + auth_context.reset(auth_token) + assert result["success"] is False + assert result["error"]["code"] == "FORBIDDEN" + assert not called + assert "must-not-appear" not in json.dumps(result) diff --git a/tests/test_x_publishing.py b/tests/test_x_publishing.py new file mode 100644 index 0000000000000000000000000000000000000000..c0e7dcc540beb9c50384bb0f7c84c7b45fe6ffd6 --- /dev/null +++ b/tests/test_x_publishing.py @@ -0,0 +1,749 @@ +"""Phase 5B X publishing coverage using only mocked official API traffic.""" + +from __future__ import annotations + +import base64 +import json +from datetime import datetime, timedelta, timezone +from pathlib import Path +from urllib.parse import parse_qs + +import httpx +import pytest +from pydantic import ValidationError + +from app.container import build_container +from app.core.config import Settings +from app.social.domain.errors import ( + SocialAccountNotFoundError, + SocialCapabilityUnsupportedError, + SocialIdempotencyConflictError, + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialPostNotFoundError, + SocialProviderUnavailableError, + SocialPublishFailedError, + SocialRateLimitedError, + SocialReauthRequiredError, +) +from app.social.models import SocialAccount +from app.social.providers.x import XProvider +from app.social.schemas.posts import SocialPostCreate +from app.social.schemas.x import XPostMetadata +from app.social.workers.publisher import SocialPublisher + +_REDIRECT_URI = "https://api.example.com/v1/social/accounts/x/callback" + + +def publishing_settings(tmp_path: Path, **overrides: object) -> Settings: + values: dict[str, object] = { + "_env_file": None, + "auth_enabled": False, + "database_url": f"sqlite+aiosqlite:///{tmp_path / 'security.db'}", + "social_database_url": f"sqlite+aiosqlite:///{tmp_path / 'social.db'}", + "social_auto_migrate": True, + "social_worker_enabled": False, + "social_oauth_encryption_key": "phase-5b-test-encryption-material", + "social_oauth_redirect_base_url": "https://api.example.com", + "x_client_id": "x-client-id", + "x_client_secret": "x-client-secret", + "x_redirect_uri": _REDIRECT_URI, + "x_publishing_enabled": True, + "x_upload_chunk_bytes": 1_048_576, + "x_media_processing_poll_seconds": 1, + "temp_dir": tmp_path / "temp", + "output_dir": tmp_path / "outputs", + "cleanup_interval_seconds": 3600, + "whisper_model": "tiny", + } + values.update(overrides) + return Settings(**values) + + +def video_probe(**overrides: object) -> dict[str, object]: + values: dict[str, object] = { + "container": "mov,mp4,m4a,3gp,3g2,mj2", + "duration": 15.0, + "fps": 30.0, + "resolution": {"width": 1280, "height": 720}, + "video_streams": [{ + "codec": "h264", + "profile": "High", + "pixel_format": "yuv420p", + "field_order": "progressive", + "sample_aspect_ratio": "1:1", + }], + "audio_streams": [{"codec": "aac", "profile": "LC", "channels": 2}], + } + values.update(overrides) + return values + + +def image_probe( + *, codec: str = "png", container: str = "png_pipe" +) -> dict[str, object]: + return { + "container": container, + "duration": None, + "fps": 25.0, + "resolution": {"width": 1200, "height": 675}, + "video_streams": [{"codec": codec, "frame_count": 1}], + "audio_streams": [], + } + + +def media(path: Path, mime_type: str, probe: dict[str, object]) -> dict[str, object]: + return { + "path": path, + "mime_type": mime_type, + "file_size": path.stat().st_size, + "probe": probe, + } + + +def x_post_payload( + account_id: str, + *, + text: str = "A production-safe X update", + publish_mode: str = "draft", + scheduled_at: datetime | None = None, +) -> SocialPostCreate: + value: dict[str, object] = { + "publish_mode": publish_mode, + "targets": [{ + "social_account_id": account_id, + "caption": {"text": text}, + "x": {"text": text}, + }], + } + if scheduled_at is not None: + value.update({"scheduled_at": scheduled_at, "timezone": "Asia/Tokyo"}) + return SocialPostCreate.model_validate(value) + + +async def connected_x_account( + container: object, + workspace_id: str, + *, + scopes: list[str] | None = None, +) -> SocialAccount: + social = container.social # type: ignore[attr-defined] + account = await social.accounts.repository.create( + SocialAccount( + workspace_id=workspace_id, + provider="x", + account_type="user", + external_account_id="2244994945", + username="x_user", + display_name="X User", + status="connected", + ) + ) + await social.accounts.tokens.store( + workspace_id, + account.id, + {"access_token": "provider-token", "refresh_token": "refresh-token"}, + expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + scopes=scopes + or ["tweet.read", "users.read", "offline.access", "tweet.write", "media.write"], + token_type="bearer", + ) + return account + + +def test_x_metadata_is_typed_and_text_only_posts_are_supported() -> None: + metadata = XPostMetadata.model_validate({ + "text": "A reply", + "reply": { + "in_reply_to_tweet_id": "1890123456789012345", + "auto_populate_reply_metadata": True, + }, + }) + assert metadata.to_post_payload() == { + "text": "A reply", + "reply": { + "in_reply_to_tweet_id": "1890123456789012345", + "auto_populate_reply_metadata": True, + }, + } + with pytest.raises(ValidationError): + XPostMetadata.model_validate({"text": "post", "quote_tweet_id": "1"}) + with pytest.raises(ValidationError): + XPostMetadata.model_validate({ + "text": "post", + "reply": {"in_reply_to_tweet_id": "١٢٣"}, + }) + with pytest.raises(ValidationError): + XPostMetadata.model_validate({ + "text": "post", + "reply": { + "in_reply_to_tweet_id": "123", + "exclude_reply_controls": True, + }, + }) + + post = SocialPostCreate.model_validate({ + "targets": [{"social_account_id": "x-account", "x": {"text": "post"}}], + }) + assert post.media_asset_id is None + with pytest.raises(ValidationError): + SocialPostCreate.model_validate({ + "targets": [{"social_account_id": "x-account", "x": {}}], + }) + + +async def test_x_text_and_reply_creation_uses_official_create_post_endpoint( + tmp_path: Path, +) -> None: + requests: list[httpx.Request] = [] + persisted: list[dict[str, object]] = [] + + async def persist(value: dict[str, object]) -> None: + persisted.append(dict(value)) + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.url == httpx.URL("https://api.x.com/2/tweets") + assert request.headers["authorization"] == "Bearer provider-token" + assert json.loads(request.content) == { + "text": "A production update", + "reply": { + "in_reply_to_tweet_id": "1890123456789012345", + "auto_populate_reply_metadata": True, + }, + } + return httpx.Response( + 201, + json={"data": {"id": "1901234567890123456", "text": "A production update"}}, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + result = await provider.publish( + {"access_token": "provider-token"}, + { + "x_post_metadata": { + "text": "A production update", + "reply": { + "in_reply_to_tweet_id": "1890123456789012345", + "auto_populate_reply_metadata": True, + }, + }, + "upload": {"identity_type": "none"}, + "provider_state": {"x_post_submission_attempted": False}, + "persist_provider_state": persist, + }, + ) + finally: + await client.aclose() + + assert result["id"] == "1901234567890123456" + assert len(requests) == 1 + assert persisted[-1]["x_post_submission_attempted"] is True + + +async def test_x_image_upload_and_media_post_use_server_generated_media_id( + tmp_path: Path, +) -> None: + image = tmp_path / "post.png" + image.write_bytes(b"production-image") + persisted: list[dict[str, object]] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/2/media/upload": + body = json.loads(request.content) + assert body == { + "media": base64.b64encode(image.read_bytes()).decode("ascii"), + "media_category": "tweet_image", + "media_type": "image/png", + "shared": False, + } + return httpx.Response( + 200, + json={"data": {"id": "1890000000000000001", "media_key": "3_1890000000000000001"}}, + ) + assert request.url.path == "/2/tweets" + assert json.loads(request.content) == { + "media": {"media_ids": ["1890000000000000001"]}, + } + return httpx.Response( + 201, + json={"data": {"id": "1900000000000000001", "text": ""}}, + ) + + async def persist(value: dict[str, object]) -> None: + persisted.append(dict(value)) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + uploaded = await provider.upload_media( + {"access_token": "provider-token"}, + { + **media(image, "image/png", image_probe()), + "persist_provider_state": persist, + }, + ) + published = await provider.publish( + {"access_token": "provider-token"}, + { + "upload": uploaded, + "x_post_metadata": {}, + "provider_state": persisted[-1], + "persist_provider_state": persist, + }, + ) + finally: + await client.aclose() + + assert uploaded["identity_type"] == "media" + assert published["id"] == "1900000000000000001" + assert persisted[-1]["x_media_id"] == "1890000000000000001" + assert persisted[-1]["x_post_submission_attempted"] is True + + +async def test_x_video_upload_streams_chunks_finalizes_and_checks_status( + tmp_path: Path, +) -> None: + chunk_size = 1_048_576 + video = tmp_path / "post.mp4" + video.write_bytes(b"a" * chunk_size + b"tail") + appended: list[tuple[int, bytes]] = [] + heartbeats = 0 + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/initialize"): + assert json.loads(request.content) == { + "media_category": "tweet_video", + "media_type": "video/mp4", + "shared": False, + "total_bytes": video.stat().st_size, + } + return httpx.Response( + 200, + json={"data": {"id": "1890000000000000002", "media_key": "7_1890000000000000002"}}, + ) + if request.url.path.endswith("/append"): + body = json.loads(request.content) + appended.append((body["segment_index"], base64.b64decode(body["media"]))) + return httpx.Response(200, json={"data": {}}) + if request.url.path.endswith("/finalize"): + return httpx.Response(200, json={"data": {"id": "1890000000000000002"}}) + assert request.method == "GET" + assert request.url.path == "/2/media/upload" + assert parse_qs(request.url.query.decode()) == { + "command": ["STATUS"], + "media_id": ["1890000000000000002"], + } + return httpx.Response( + 200, + json={"data": {"processing_info": {"state": "succeeded"}}}, + ) + + async def heartbeat() -> None: + nonlocal heartbeats + heartbeats += 1 + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + result = await provider.upload_media( + {"access_token": "provider-token"}, + { + **media(video, "video/mp4", video_probe()), + "persist_provider_state": lambda _value: _async_none(), + "heartbeat": heartbeat, + }, + ) + finally: + await client.aclose() + + assert result["id"] == "1890000000000000002" + assert appended == [(0, b"a" * chunk_size), (1, b"tail")] + assert heartbeats == 2 + + +async def _async_none() -> None: + return None + + +async def test_x_animated_gif_uses_chunked_tweet_gif_workflow( + tmp_path: Path, +) -> None: + gif = tmp_path / "post.gif" + gif.write_bytes(b"gif-data") + paths: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + paths.append(request.url.path) + if request.url.path.endswith("/initialize"): + assert json.loads(request.content)["media_category"] == "tweet_gif" + return httpx.Response( + 200, + json={ + "data": { + "id": "1890000000000000003", + "media_key": "16_1890000000000000003", + } + }, + ) + if request.url.path.endswith("/append"): + return httpx.Response(200, json={"data": {}}) + if request.url.path.endswith("/finalize"): + return httpx.Response(200, json={"data": {}}) + return httpx.Response(200, json={"data": {}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + result = await provider.upload_media( + {"access_token": "provider-token"}, + media(gif, "image/gif", image_probe(codec="gif", container="gif")), + ) + finally: + await client.aclose() + + assert result["id"] == "1890000000000000003" + assert "/2/media/upload/initialize" in paths + assert paths.count("/2/media/upload") == 1 + + +@pytest.mark.parametrize( + ("probe", "message"), + [ + (video_probe(duration=141), "duration"), + (video_probe(fps=None), "frame rate"), + (video_probe(fps=61), "frame rate"), + (video_probe(resolution={"width": 1280, "height": 300}), "aspect ratio"), + (video_probe(video_streams=[{"codec": "hevc"}]), "H.264"), + ( + video_probe( + video_streams=[{ + "codec": "h264", + "pixel_format": None, + "field_order": "progressive", + "sample_aspect_ratio": "1:1", + }] + ), + "4:2:0", + ), + (video_probe(audio_streams=[{"codec": "opus"}]), "AAC"), + ], +) +async def test_x_media_validation_rejects_incompatible_video( + tmp_path: Path, + probe: dict[str, object], + message: str, +) -> None: + video = tmp_path / "invalid.mp4" + video.write_bytes(b"invalid") + provider = XProvider(publishing_settings(tmp_path)) + try: + with pytest.raises(SocialMediaInvalidError, match=message): + await provider.validate_media(media(video, "video/mp4", probe)) + finally: + await provider.close() + + +@pytest.mark.parametrize( + ("status", "expected"), + [ + (401, SocialReauthRequiredError), + (403, SocialPermissionDeniedError), + (404, SocialCapabilityUnsupportedError), + (429, SocialRateLimitedError), + (500, SocialProviderUnavailableError), + (502, SocialProviderUnavailableError), + (503, SocialProviderUnavailableError), + (504, SocialProviderUnavailableError), + (400, SocialPublishFailedError), + (422, SocialPublishFailedError), + ], +) +async def test_x_post_error_normalization( + tmp_path: Path, + status: int, + expected: type[Exception], +) -> None: + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(status, json={"title": "provider detail must stay private"}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + persisted: list[dict[str, object]] = [] + + async def persist(value: dict[str, object]) -> None: + persisted.append(dict(value)) + + try: + with pytest.raises(expected) as raised: + await provider.publish( + {"access_token": "provider-token"}, + { + "x_post_metadata": {"text": "post"}, + "upload": {}, + "provider_state": {"x_post_submission_attempted": False}, + "persist_provider_state": persist, + }, + ) + finally: + await client.aclose() + assert "provider-token" not in str(raised.value) + assert "provider detail" not in str(raised.value) + assert persisted + if status in {400, 401, 403, 404, 422, 429}: + assert persisted[-1]["x_post_submission_attempted"] is False + assert "x_post_submission_started_at" not in persisted[-1] + else: + assert persisted[-1]["x_post_submission_attempted"] is True + + +async def test_x_media_access_tier_404_is_a_capability_error(tmp_path: Path) -> None: + image = tmp_path / "post.png" + image.write_bytes(b"image") + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(404, json={"title": "Not Found"}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialCapabilityUnsupportedError): + await provider.upload_media( + {"access_token": "provider-token"}, + media(image, "image/png", image_probe()), + ) + finally: + await client.aclose() + + +async def test_x_status_delete_and_pending_publish_reconciliation( + tmp_path: Path, +) -> None: + started = datetime.now(timezone.utc) - timedelta(seconds=1) + calls: list[str] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + calls.append(f"{request.method} {request.url.path}") + if request.method == "GET" and request.url.path.endswith("/tweets"): + query = parse_qs(request.url.query.decode()) + assert query["max_results"] == ["100"] + assert "start_time" in query + return httpx.Response(200, json={ + "data": [{ + "id": "1900000000000000004", + "text": "Recovered post", + "created_at": datetime.now(timezone.utc).isoformat(), + }], + "meta": {"result_count": 1}, + }) + if request.method == "GET": + return httpx.Response(200, json={ + "data": {"id": "1900000000000000004", "text": "Recovered post"}, + }) + assert request.method == "DELETE" + return httpx.Response(200, json={"data": {"deleted": True}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + try: + recovered = await provider.reconcile_pending_publish( + {"access_token": "provider-token"}, + { + "provider_account_id": "2244994945", + "provider_state": { + "x_post_submission_attempted": True, + "x_post_submission_started_at": started.isoformat(), + }, + "x_post_metadata": {"text": "Recovered post"}, + "upload": {"identity_type": "none"}, + }, + ) + status = await provider.get_publish_status( + {"access_token": "provider-token"}, "1900000000000000004" + ) + await provider.delete_post( + {"access_token": "provider-token"}, "1900000000000000004" + ) + finally: + await client.aclose() + + assert recovered and recovered["id"] == "1900000000000000004" + assert status["status"] == "published" + assert calls[-1] == "DELETE /2/tweets/1900000000000000004" + + +async def test_x_uncertain_create_post_outcome_never_resubmits( + tmp_path: Path, +) -> None: + create_calls = 0 + provider_state: dict[str, object] = { + "x_post_submission_attempted": False, + "x_publish_started_at": datetime.now(timezone.utc).isoformat(), + } + + async def persist(value: dict[str, object]) -> None: + provider_state.clear() + provider_state.update(value) + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal create_calls + if request.method == "POST": + create_calls += 1 + raise httpx.ReadTimeout("response lost after X accepted the post", request=request) + assert request.method == "GET" + return httpx.Response(200, json={"meta": {"result_count": 0}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = XProvider(publishing_settings(tmp_path), http_client=client) + payload = { + "provider_account_id": "2244994945", + "provider_state": provider_state, + "persist_provider_state": persist, + "x_post_metadata": {"text": "Uncertain post"}, + "upload": {"identity_type": "none"}, + } + try: + with pytest.raises(SocialProviderUnavailableError): + await provider.publish( + {"access_token": "provider-token"}, + payload, + ) + assert provider_state["x_post_submission_attempted"] is True + + with pytest.raises(SocialProviderUnavailableError, match="reconciled"): + await provider.reconcile_pending_publish( + {"access_token": "provider-token"}, + {**payload, "provider_state": provider_state}, + ) + + with pytest.raises(SocialProviderUnavailableError, match="uncertain"): + await provider.publish( + {"access_token": "provider-token"}, + {**payload, "provider_state": provider_state}, + ) + finally: + await client.aclose() + + assert create_calls == 1 + + +async def test_x_idempotency_scheduling_authorization_and_workspace_isolation( + tmp_path: Path, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + try: + account = await connected_x_account(container, "workspace-a") + payload = x_post_payload(account.id) + first = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="x-create-key", + ) + duplicate = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=payload, + idempotency_key="x-create-key", + ) + assert duplicate.id == first.id + with pytest.raises(SocialIdempotencyConflictError): + await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=x_post_payload(account.id, text="Different payload"), + idempotency_key="x-create-key", + ) + scheduled = await container.social.publishing.create( + workspace_id="workspace-a", + user_id="user-a", + payload=x_post_payload( + account.id, + publish_mode="schedule", + scheduled_at=datetime.now(timezone.utc) + timedelta(hours=1), + ), + idempotency_key="x-schedule-key", + ) + assert scheduled.status.value == "scheduled" + with pytest.raises(SocialAccountNotFoundError): + await container.social.publishing.create( + workspace_id="workspace-b", + user_id="user-b", + payload=x_post_payload(account.id), + idempotency_key="workspace-b-key", + ) + with pytest.raises(SocialPostNotFoundError): + await container.social.publishing.delete("workspace-b", first.id) + + read_only = await connected_x_account( + container, + "workspace-read-only", + scopes=["tweet.read", "users.read", "offline.access"], + ) + with pytest.raises(SocialPermissionDeniedError): + await container.social.publishing.create( + workspace_id="workspace-read-only", + user_id="user-read-only", + payload=x_post_payload(read_only.id, publish_mode="now"), + idempotency_key="missing-write-scopes", + ) + finally: + await container.social.close() + await container.security_database.close() + + +async def test_x_worker_publishes_text_once_and_persists_final_post_identity( + tmp_path: Path, +) -> None: + container = build_container(publishing_settings(tmp_path)) + await container.social.initialize() + adapter = container.social.accounts.providers.get("x") + assert isinstance(adapter, XProvider) + await adapter._client.aclose() + create_calls = 0 + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal create_calls + if request.method == "GET" and request.url.path.endswith("/users/2244994945/tweets"): + return httpx.Response(200, json={"meta": {"result_count": 0}}) + if request.method == "POST" and request.url.path == "/2/tweets": + create_calls += 1 + return httpx.Response(201, json={ + "data": {"id": "1900000000000000005", "text": "Worker post"}, + }) + if request.method == "GET" and request.url.path.endswith("1900000000000000005"): + return httpx.Response(200, json={ + "data": {"id": "1900000000000000005", "text": "Worker post"}, + }) + raise AssertionError(f"Unexpected request: {request.method} {request.url}") + + adapter._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + adapter._owns_client = True + try: + account = await connected_x_account(container, "workspace-worker") + post = await container.social.publishing.create( + workspace_id="workspace-worker", + user_id="worker-user", + payload=x_post_payload( + account.id, text="Worker post", publish_mode="now" + ), + idempotency_key="worker-create-key", + ) + jobs = await container.social.jobs.list("workspace-worker") + assert len(jobs) == 1 + await SocialPublisher(container.social).process("workspace-worker", jobs[0].id) + + stored = await container.social.publishing.get("workspace-worker", post.id) + stored_job = await container.social.jobs.get("workspace-worker", jobs[0].id) + assert stored.status.value == "published" + assert stored.targets[0].external_post_id == "1900000000000000005" + assert stored_job.status.value == "published" + assert create_calls == 1 + assert "provider-token" not in stored.model_dump_json() + assert "provider-token" not in stored_job.model_dump_json() + finally: + await container.social.close() + await container.security_database.close() diff --git a/tests/test_youtube_live.py b/tests/test_youtube_live.py new file mode 100644 index 0000000000000000000000000000000000000000..d1cef25dc483022deaee16349841099117114dd5 --- /dev/null +++ b/tests/test_youtube_live.py @@ -0,0 +1,109 @@ +"""Opt-in, destructive YouTube Data API integration test. + +OAuth browser consent remains an operator action. The access token supplied +to this test must therefore come from the staging channel after that consent +flow. CI never runs this module implicitly and the test refuses to upload +unless explicit cleanup has been requested. +""" + +from __future__ import annotations + +import os + +import pytest + + +pytestmark = pytest.mark.skipif( + os.getenv("RUN_YOUTUBE_INTEGRATION_TESTS") != "true", + reason="YouTube live integration is NOT VERIFIED; set RUN_YOUTUBE_INTEGRATION_TESTS=true with staging credentials.", +) + + +def _require_live_configuration() -> None: + required = ( + "GOOGLE_CLIENT_ID", + "GOOGLE_CLIENT_SECRET", + "YOUTUBE_LIVE_TEST_ACCESS_TOKEN", + "YOUTUBE_LIVE_TEST_MEDIA_PATH", + ) + missing = [name for name in required if not os.getenv(name)] + if missing: + pytest.skip(f"YouTube live integration is NOT VERIFIED; missing {', '.join(missing)}") + if os.getenv("YOUTUBE_LIVE_TEST_DELETE") != "true": + pytest.skip("Set YOUTUBE_LIVE_TEST_DELETE=true to permit cleanup of the staged test video.") + + +def test_youtube_live_configuration_is_explicit() -> None: + """Protect the opt-in switch from accidentally becoming an implicit test.""" + assert os.getenv("RUN_YOUTUBE_INTEGRATION_TESTS") == "true" + _require_live_configuration() + + +@pytest.mark.asyncio +async def test_youtube_live_channel_upload_status_metrics_and_delete() -> None: + """Exercise the official API against an operator-provisioned staging grant. + + The access token is never printed or returned. The test discovers the + channel, streams a real local test asset through the resumable endpoint, + reconciles the normalized status and statistics, then deletes the video. + """ + # Imports stay inside the opt-in test so a normal collection on a minimal + # machine does not need the backend's dependency set. + _require_live_configuration() + from pathlib import Path + + from app.core.config import Settings + from app.services.ffprobe_service import FFprobeService + from app.services.validator import MediaValidator + from app.social.providers.youtube import YouTubeProvider + from app.social.schemas.youtube import YouTubePostMetadata + + path = Path(os.environ["YOUTUBE_LIVE_TEST_MEDIA_PATH"]).expanduser().resolve() + if not path.is_file(): + pytest.skip("YOUTUBE_LIVE_TEST_MEDIA_PATH does not point to a readable staged video.") + settings = Settings( + _env_file=None, + auth_enabled=False, + google_client_id=os.environ["GOOGLE_CLIENT_ID"], + google_client_secret=os.environ["GOOGLE_CLIENT_SECRET"], + ) + provider = YouTubeProvider(settings) + token: dict[str, object] = {"access_token": os.environ["YOUTUBE_LIVE_TEST_ACCESS_TOKEN"]} + if refresh_token := os.getenv("YOUTUBE_LIVE_TEST_REFRESH_TOKEN"): + token["refresh_token"] = refresh_token + refreshed = await provider.refresh_token(token) + token = {**token, **refreshed} + + uploaded_id: str | None = None + try: + account = await provider.get_account(token) + assert account["external_account_id"] + probe = await FFprobeService(settings).probe(path) + mime_type = MediaValidator(settings).infer_mime(path) + metadata = YouTubePostMetadata( + title="MediaRouter YouTube integration test", + description="Automatically deleted staging verification video.", + privacy_status="private", + made_for_kids=False, + notify_subscribers=False, + ) + uploaded = await provider.upload_media( + token, + { + "path": path, + "mime_type": mime_type, + "file_size": path.stat().st_size, + "probe": probe, + "youtube_resource": metadata.to_youtube_resource(), + "notify_subscribers": False, + }, + ) + uploaded_id = str(uploaded["id"]) + status = await provider.get_publish_status(token, uploaded_id) + assert status["status"] in {"processing", "published"} + metrics = await provider.get_metrics(token, uploaded_id) + assert metrics["status"] in {"available", "unavailable"} + finally: + if uploaded_id: + await provider.delete_post(token, uploaded_id) + await provider.close() diff --git a/tests/test_youtube_provider.py b/tests/test_youtube_provider.py new file mode 100644 index 0000000000000000000000000000000000000000..3b3642b4f427b3488aa7354549f53a3d34d2038b --- /dev/null +++ b/tests/test_youtube_provider.py @@ -0,0 +1,292 @@ +from __future__ import annotations + +from pathlib import Path +from urllib.parse import parse_qs, urlparse + +import httpx +import pytest +from pydantic import ValidationError + +from app.core.config import Settings +from app.social.domain.errors import ( + SocialMediaInvalidError, + SocialPermissionDeniedError, + SocialReauthRequiredError, +) +from app.social.providers.youtube import YouTubeProvider +from app.social.schemas.youtube import YouTubePostMetadata + + +def settings(tmp_path: Path) -> Settings: + return Settings( + _env_file=None, + auth_enabled=False, + google_client_id="google-client-id", + google_client_secret="google-client-secret", + social_oauth_encryption_key="test-only-encryption-material", + temp_dir=tmp_path / "temp", + output_dir=tmp_path / "outputs", + youtube_upload_chunk_bytes=262_144, + whisper_model="tiny", + ) + + +def metadata() -> YouTubePostMetadata: + return YouTubePostMetadata( + title="MediaRouter test video", + description="A test upload", + tags=["mediarouter", "test"], + privacy_status="private", + made_for_kids=False, + ) + + +async def test_youtube_authorization_url_requests_minimum_scope_and_pkce(tmp_path: Path) -> None: + client = httpx.AsyncClient(transport=httpx.MockTransport(lambda request: httpx.Response(500))) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + url = await provider.get_authorization_url( + state="s" * 32, + redirect_uri="https://api.example/v1/social/accounts/youtube/callback", + code_challenge="challenge", + ) + finally: + await client.aclose() + parsed = parse_qs(urlparse(url).query) + assert parsed["scope"] == ["https://www.googleapis.com/auth/youtube.upload"] + assert parsed["code_challenge"] == ["challenge"] + assert parsed["code_challenge_method"] == ["S256"] + assert parsed["access_type"] == ["offline"] + + +async def test_youtube_exchange_refresh_and_channel_discovery(tmp_path: Path) -> None: + received_verifier = False + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal received_verifier + if request.url.path == "/token": + form = request.content.decode() + if "grant_type=authorization_code" in form: + received_verifier = "code_verifier=verifier" in form + return httpx.Response(200, json={"access_token": "access", "refresh_token": "refresh", "expires_in": 3600, "scope": "https://www.googleapis.com/auth/youtube.upload", "token_type": "Bearer"}) + return httpx.Response(200, json={"access_token": "refreshed", "expires_in": 3600, "token_type": "Bearer"}) + if request.url.path.endswith("/channels"): + return httpx.Response(200, json={"items": [{"id": "UC-stable-channel", "snippet": {"title": "Test Channel", "customUrl": "@test", "thumbnails": {"high": {"url": "https://img.example/avatar.jpg"}}}}]}) + raise AssertionError(f"Unexpected request {request.method} {request.url}") + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + exchanged = await provider.exchange_code(code="code", redirect_uri="https://api.example/callback", code_verifier="verifier") + refreshed = await provider.refresh_token(exchanged) + account = await provider.get_account(exchanged) + finally: + await client.aclose() + assert exchanged["access_token"] == "access" + assert received_verifier + assert refreshed["access_token"] == "refreshed" + assert account["external_account_id"] == "UC-stable-channel" + assert account["username"] == "@test" + assert "email" not in account + + +async def test_youtube_resumable_upload_streams_file_and_returns_video_id(tmp_path: Path) -> None: + media = tmp_path / "video.mp4" + media.write_bytes(b"video-bytes") + requests: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.method == "POST" and request.url.path == "/upload/youtube/v3/videos": + return httpx.Response(200, headers={"Location": "https://www.googleapis.com/upload/youtube/v3/videos?upload_id=session"}) + if request.method == "PUT": + assert request.headers["Content-Range"] == f"bytes 0-{media.stat().st_size - 1}/{media.stat().st_size}" + return httpx.Response(200, json={"id": "yt-video-id"}) + raise AssertionError(f"Unexpected request {request.method} {request.url}") + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + sessions: list[str | None] = [] + try: + result = await provider.upload_media( + {"access_token": "access"}, + { + "path": media, + "mime_type": "video/mp4", + "file_size": media.stat().st_size, + "probe": { + "container": "mov,mp4,m4a,3gp,3g2,mj2", + "duration": 1.0, + "resolution": {"width": 1280, "height": 720}, + "video_streams": [{"codec": "h264"}], + }, + "youtube_resource": metadata().to_youtube_resource(), + "persist_upload_session": sessions.append, + }, + ) + finally: + await client.aclose() + assert result == {"id": "yt-video-id", "url": "https://www.youtube.com/watch?v=yt-video-id"} + assert sessions == ["https://www.googleapis.com/upload/youtube/v3/videos?upload_id=session"] + assert len(requests) == 2 + + +async def test_youtube_resumable_upload_reconciles_a_retry_without_restarting(tmp_path: Path) -> None: + media = tmp_path / "video.mp4" + media.write_bytes(b"video-bytes") + chunk_attempts = 0 + session_queries = 0 + + async def handler(request: httpx.Request) -> httpx.Response: + nonlocal chunk_attempts, session_queries + if request.method == "POST": + return httpx.Response(200, headers={"Location": "https://www.googleapis.com/upload/youtube/v3/videos?upload_id=session"}) + content_range = request.headers.get("Content-Range") + if content_range == f"bytes */{media.stat().st_size}": + session_queries += 1 + return httpx.Response(200, json={"id": "yt-reconciled-video"}) + chunk_attempts += 1 + return httpx.Response(503, json={"error": {"errors": [{"reason": "backendError"}]}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + result = await provider.upload_media( + {"access_token": "access"}, + { + "path": media, + "mime_type": "video/mp4", + "file_size": media.stat().st_size, + "probe": { + "container": "mov,mp4,m4a,3gp,3g2,mj2", + "duration": 1.0, + "resolution": {"width": 1280, "height": 720}, + "video_streams": [{"codec": "h264"}], + }, + "youtube_resource": metadata().to_youtube_resource(), + }, + ) + finally: + await client.aclose() + assert result["id"] == "yt-reconciled-video" + assert chunk_attempts == 1 + assert session_queries == 1 + + +async def test_youtube_rejects_invalid_media_before_creating_an_upload_session(tmp_path: Path) -> None: + media = tmp_path / "audio.mp3" + media.write_bytes(b"not-a-video") + client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _: httpx.Response(500))) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialMediaInvalidError): + await provider.validate_media( + { + "path": media, + "mime_type": "audio/mpeg", + "file_size": media.stat().st_size, + "probe": {}, + } + ) + finally: + await client.aclose() + + +async def test_youtube_status_deletion_and_public_video_metrics(tmp_path: Path) -> None: + async def handler(request: httpx.Request) -> httpx.Response: + if request.method == "DELETE": + return httpx.Response(204) + if request.url.params.get("part") == "snippet,status,processingDetails": + return httpx.Response(200, json={"items": [{"id": "video", "snippet": {"publishedAt": "2030-01-01T00:00:00Z"}, "status": {"uploadStatus": "processed", "privacyStatus": "unlisted"}, "processingDetails": {"processingStatus": "succeeded"}}]}) + if request.url.params.get("part") == "statistics,snippet,status": + return httpx.Response(200, json={"items": [{"id": "video", "snippet": {"publishedAt": "2030-01-01T00:00:00Z"}, "statistics": {"viewCount": "7", "likeCount": "2", "commentCount": "1"}}]}) + raise AssertionError(f"Unexpected request {request.method} {request.url}") + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + status = await provider.get_publish_status({"access_token": "access"}, "video") + metrics = await provider.get_metrics({"access_token": "access"}, "video") + await provider.delete_post({"access_token": "access"}, "video") + finally: + await client.aclose() + assert status["status"] == "published" + assert metrics["views"] == 7 + assert metrics["comments"] == 1 + + +async def test_youtube_provider_normalizes_auth_and_permission_errors(tmp_path: Path) -> None: + async def unauthorized(_: httpx.Request) -> httpx.Response: + return httpx.Response(401, json={"error": {"message": "do not leak"}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(unauthorized)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialReauthRequiredError): + await provider.get_publish_status({"access_token": "access"}, "video") + finally: + await client.aclose() + + async def forbidden(_: httpx.Request) -> httpx.Response: + return httpx.Response(403, json={"error": {"errors": [{"reason": "forbidden"}]}}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(forbidden)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialPermissionDeniedError): + await provider.delete_post({"access_token": "access"}, "video") + finally: + await client.aclose() + + +async def test_youtube_rejects_a_pkce_mismatch_without_exposing_google_details(tmp_path: Path) -> None: + async def rejected(_: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + json={"error": {"message": "PKCE verifier did not match", "errors": []}}, + ) + + client = httpx.AsyncClient(transport=httpx.MockTransport(rejected)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + with pytest.raises(SocialReauthRequiredError) as raised: + await provider.exchange_code( + code="code", + redirect_uri="https://api.example/callback", + code_verifier="wrong-verifier", + ) + finally: + await client.aclose() + assert "PKCE verifier" not in str(raised.value) + + +async def test_youtube_status_reconciliation_reports_processing_and_failed(tmp_path: Path) -> None: + responses = iter( + [ + {"items": [{"id": "video", "status": {"uploadStatus": "uploaded"}, "processingDetails": {"processingStatus": "processing"}}]}, + {"items": [{"id": "video", "status": {"uploadStatus": "failed"}, "processingDetails": {"processingStatus": "failed", "processingFailureReason": "transcodeFailed"}}]}, + ] + ) + + async def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=next(responses)) + + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + provider = YouTubeProvider(settings(tmp_path), http_client=client) + try: + processing = await provider.get_publish_status({"access_token": "access"}, "video") + failed = await provider.get_publish_status({"access_token": "access"}, "video") + finally: + await client.aclose() + assert processing["status"] == "processing" + assert failed["status"] == "failed" + assert failed["metadata"]["failure_reason"] == "transcodeFailed" + + +def test_youtube_metadata_requires_explicit_policy_declaration() -> None: + with pytest.raises(ValidationError): + YouTubePostMetadata.model_validate({"title": "No audience declaration"}) + with pytest.raises(ValidationError): + YouTubePostMetadata.model_validate( + {"title": "Invalid scheduled status", "made_for_kids": False, "privacy_status": "public", "scheduled_publish_at": "2030-01-01T00:00:00Z"} + ) diff --git a/workers/__init__.py b/workers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e1a071c3d48104549283fb8acebac70d5c9fe906 --- /dev/null +++ b/workers/__init__.py @@ -0,0 +1,3 @@ +"""Compatibility exports for the canonical :mod:`app.workers` package.""" + +from app.workers.cleanup_worker import * # noqa diff --git a/workers/cleanup_worker.py b/workers/cleanup_worker.py new file mode 100644 index 0000000000000000000000000000000000000000..e18f1bc506c49ec237927eaa8d76b2b4a3be0426 --- /dev/null +++ b/workers/cleanup_worker.py @@ -0,0 +1 @@ +from app.workers.cleanup_worker import * # noqa