Spaces:
Running
Running
| 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] | |
| 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, | |
| ) | |
| def _operation_for_model(item: GenerationModelView) -> str: | |
| return ( | |
| "generate_image" | |
| if item.model.modality is GenerationModality.IMAGE | |
| else "generate_video" | |
| ) | |
| 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, | |
| ) | |