from __future__ import annotations from app.ai.schemas import ( AiCapabilities, AiGenerateImageRequest, AiGenerateVideoRequest, AiGenerationRequest, AiHistory, AiJob, AiModel, AiOutput, AiProvider, AiTool, ) from app.generation.domain.enums import ( GenerationModality, GenerationRequestStatus, ) from app.generation.domain.errors import ( GenerationCapabilityUnsupportedError, GenerationValidationError, ) from app.generation.model_registry import GenerationModelView from app.generation.schemas.requests import ( FluxGenerationOptions, GenerationRequestCreate, GenerationRequestView, WanGenerationOptions, ) from app.generation.services.generation_service import GenerationService from app.projects.schemas import ProjectStatus from app.projects.services.project_service import ProjectService from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService from app.security.audit import AuditService _OPERATION_MODALITY = { "generate_image": GenerationModality.IMAGE, "generate_video": GenerationModality.VIDEO, } class AiStudioService: """Capability-driven facade over the existing durable generation domain.""" def __init__( self, generation: GenerationService, projects: ProjectService, assets: CanonicalAssetService, audit: AuditService, ) -> None: self.generation = generation self.projects = projects self.assets = assets self.audit = audit def capabilities(self) -> AiCapabilities: models = self.generation.models.list() tools = [ self._tool( operation="generate_image", name="Generate Image", description="Create or transform a canonical image with an available image model.", models=[item for item in models if item.model.modality is GenerationModality.IMAGE], input_types=["text", "image"], output_types=["image"], ), self._tool( operation="generate_video", name="Generate Video", description="Create a video from a canonical project image.", models=[item for item in models if item.model.modality is GenerationModality.VIDEO], input_types=["text", "image"], output_types=["video"], ), ] providers = [] for adapter in self.generation.providers.list(): provider_models = [item for item in models if item.provider_id == adapter.provider] operations = sorted({self._operation_for_model(item) for item in provider_models}) providers.append( AiProvider( id=adapter.provider, display_name=adapter.capabilities.name, available=bool( self.generation.ready and adapter.available and any(item.available for item in provider_models) ), operations=operations, supports_cancellation=adapter.capabilities.supports_cancellation, ) ) return AiCapabilities( available=any(tool.available for tool in tools), categories=["generate"], tools=tools, providers=providers, permissions=[ "ai:read", "ai:generate", "ai:transform", "ai:analyze", "ai:create", ], ) async def create( self, *, workspace_id: str, user_id: str, api_key_id: str | None, request_id: str, payload: AiGenerationRequest, idempotency_key: str, ) -> AiJob: project_id = str(payload.project_id) if payload.project_id else None source_asset_id = str(payload.source_asset_ids[0]) if payload.source_asset_ids else None if payload.project_id is not None: project = await self.projects.get( workspace_id=workspace_id, user_id=user_id, project_id=project_id or "", ) if project.status is not ProjectStatus.ACTIVE: raise GenerationValidationError("AI jobs require an active project.") if source_asset_id is not None: try: asset = await self.assets.get_owned_by_id( workspace_id=workspace_id, user_id=user_id, asset_id=source_asset_id, ) except CanonicalAssetNotFoundError as exc: raise GenerationValidationError( "AI source asset was not found in this workspace." ) from exc if project_id is not None and asset.project_id != project_id: raise GenerationValidationError( "AI source asset must belong to the selected project." ) model = self._select_model(payload.operation, payload.model) request = self._generation_request(payload, model) created = await self.generation.create( workspace_id=workspace_id, user_id=user_id, payload=request, idempotency_key=idempotency_key, project_id=(project_id if payload.output_preferences.attach_to_project else None), product_surface="ai_studio", ) await self.audit.record_event( workspace_id=workspace_id, user_id=user_id, api_key_id=api_key_id, request_id=request_id, event_type="ai.generation_requested", entity_type="generation_job", entity_id=created.job.id, metadata={ "operation": payload.operation, "model": created.model_id, "project_id": project_id, }, ) return self._job(created) async def history( self, *, workspace_id: str, user_id: str, offset: int, limit: int, ) -> AiHistory: items = await self.generation.list_requests( workspace_id, user_id, offset=offset, limit=limit, product_surface="ai_studio", ) return AiHistory( items=[self._job(item) for item in items], offset=offset, limit=limit, ) async def get(self, *, workspace_id: str, user_id: str, generation_id: str) -> AiJob: request = await self.generation.get_request(workspace_id, user_id, generation_id) if request.product_surface != "ai_studio": raise GenerationValidationError("AI generation was not found.") return self._job(request) async def cancel( self, *, workspace_id: str, user_id: str, api_key_id: str | None, request_id: str, generation_id: str, ) -> AiJob: request = await self.generation.get_request(workspace_id, user_id, generation_id) if request.product_surface != "ai_studio": raise GenerationValidationError("AI generation was not found.") await self.generation.cancel(workspace_id, user_id, request.job.id) updated = await self.generation.get_request(workspace_id, user_id, generation_id) if updated.status is GenerationRequestStatus.CANCELLED: await self.audit.record_event( workspace_id=workspace_id, user_id=user_id, api_key_id=api_key_id, request_id=request_id, event_type="ai.generation_cancelled", entity_type="generation_job", entity_id=request.job.id, metadata={"operation": self._operation_for_request(request)}, ) return self._job(updated) def _select_model(self, operation: str, requested_model: str | None) -> GenerationModelView: modality = _OPERATION_MODALITY[operation] matches = [ model for model in self.generation.models.list() if model.model.modality is modality and model.available and (requested_model is None or model.model.id == requested_model) ] if not matches: raise GenerationCapabilityUnsupportedError( "No ready model supports the requested AI operation." ) if requested_model is not None and len(matches) != 1: raise GenerationCapabilityUnsupportedError("The requested AI model is not available.") return sorted(matches, key=lambda item: (item.provider_id, item.model.id))[0] @staticmethod def _generation_request( payload: AiGenerationRequest, model: GenerationModelView ) -> GenerationRequestCreate: source_asset_id = str(payload.source_asset_ids[0]) if payload.source_asset_ids else None if isinstance(payload, AiGenerateImageRequest): options = payload.parameters return GenerationRequestCreate( provider=model.provider_id, model_id=model.model.id, modality=GenerationModality.IMAGE, prompt=payload.prompt, input_asset_id=source_asset_id, flux=FluxGenerationOptions( mode_choice=( "Base (50 steps)" if options.mode == "quality" else "Distilled (4 steps)" ), seed=options.seed, randomize_seed=options.randomize_seed, width=options.width, height=options.height, num_inference_steps=options.steps, guidance_scale=options.guidance, prompt_upsampling=options.enhance_prompt, ), ) if isinstance(payload, AiGenerateVideoRequest): options = payload.parameters return GenerationRequestCreate( provider=model.provider_id, model_id=model.model.id, modality=GenerationModality.VIDEO, prompt=payload.prompt, input_asset_id=source_asset_id, wan=WanGenerationOptions( negative_prompt=options.negative_prompt, duration_seconds=options.duration_seconds, steps=options.steps, guidance_scale=options.guidance, guidance_scale_2=options.secondary_guidance, seed=options.seed, randomize_seed=options.randomize_seed, ), ) raise GenerationCapabilityUnsupportedError("AI operation is unsupported.") def _tool( self, *, operation: str, name: str, description: str, models: list[GenerationModelView], input_types: list[str], output_types: list[str], ) -> AiTool: public_models = [self._model(item, operation) for item in models] return AiTool( id=operation.replace("_", "-"), category="generate", operation=operation, name=name, description=description, input_types=input_types, output_types=output_types, required_permission="ai:generate", supports_asset_input=any(item.model.input_asset_supported for item in models), available=any(item.available and self.generation.ready for item in models), models=public_models, ) def _model(self, item: GenerationModelView, operation: str) -> AiModel: provider = self.generation.providers.get(item.provider_id) return AiModel( id=item.model.id, display_name=item.model.name, operation=operation, input_types=(["text", "image"] if item.model.input_asset_supported else ["text"]), output_types=[item.model.modality.value], input_schema=dict(item.model.input_schema), available=bool(self.generation.ready and item.available), provider_display_name=provider.capabilities.name, ) @staticmethod def _operation_for_model(item: GenerationModelView) -> str: return ( "generate_image" if item.model.modality is GenerationModality.IMAGE else "generate_video" ) @staticmethod def _operation_for_request(item: GenerationRequestView) -> str: return "generate_image" if item.modality is GenerationModality.IMAGE else "generate_video" def _job(self, item: GenerationRequestView) -> AiJob: operation = self._operation_for_request(item) status_map = { GenerationRequestStatus.QUEUED: "queued", GenerationRequestStatus.SUBMITTING: "processing", GenerationRequestStatus.RUNNING: "processing", GenerationRequestStatus.RETRYING: "retrying", GenerationRequestStatus.SUCCEEDED: "completed", GenerationRequestStatus.FAILED: "failed", GenerationRequestStatus.CANCEL_REQUESTED: "cancelling", GenerationRequestStatus.CANCELLED: "cancelled", } provider = self.generation.providers.get(item.provider) output = None if item.job.output_asset_id and item.status is GenerationRequestStatus.SUCCEEDED: output = AiOutput( asset_id=item.job.output_asset_id, media_type=("image" if item.modality is GenerationModality.IMAGE else "video"), ) retryable_codes = { "GENERATION_MODEL_UNAVAILABLE", "GENERATION_SUBMISSION_RETRYING", "GENERATION_PROVIDER_STATUS_FAILED", } return AiJob( id=item.job.id, generation_id=item.id, operation=operation, status=status_map[item.status], prompt=item.prompt, project_id=item.project_id, source_asset_id=item.input_asset_id, output=output, model=item.model_id, provider_display_name=provider.capabilities.name, error_code=item.job.error_code, error_message=item.job.error_message, retryable=item.job.error_code in retryable_codes, created_at=item.created_at, updated_at=item.updated_at, completed_at=item.completed_at, )