MediaRouter / app /ai /service.py
basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
14.7 kB
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,
)