6efe524974
Extract CRUD helpers into services/config/crud_service.py. Move instance-related config logic to services/tool/instance_service.py. Quality gates: py_compile pass
356 lines
12 KiB
Python
356 lines
12 KiB
Python
"""Config profile CRUD service functions."""
|
|
|
|
import uuid
|
|
from typing import Any
|
|
|
|
from fastapi import HTTPException, status
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy.orm import selectinload
|
|
|
|
from src.models import ConfigProfile, ConfigProfileInclude, ToolType, UserConfig
|
|
from src.models.project import Project
|
|
|
|
MAX_PROFILE_SIZE_MB = 10
|
|
MAX_PROFILE_SIZE_BYTES = MAX_PROFILE_SIZE_MB * 1024 * 1024
|
|
|
|
|
|
def calculate_profile_size(data: dict) -> int:
|
|
"""Calculate approximate serialized size of profile data."""
|
|
total = 0
|
|
for key, value in data.get("env_vars", {}).items():
|
|
total += len(key.encode("utf-8")) + len(str(value).encode("utf-8"))
|
|
for key, value in data.get("runtime_hints", {}).items():
|
|
total += len(key.encode("utf-8")) + len(str(value).encode("utf-8"))
|
|
for mount in data.get("mounts", []):
|
|
total += len(str(mount.get("target", "")).encode("utf-8"))
|
|
total += len(str(mount.get("mode", "")).encode("utf-8"))
|
|
for path, content in mount.get("files", {}).items():
|
|
total += len(path.encode("utf-8")) + len(content.encode("utf-8"))
|
|
for path, content in data.get("files", {}).items():
|
|
total += len(path.encode("utf-8")) + len(content.encode("utf-8"))
|
|
return total
|
|
|
|
|
|
async def get_profile_with_includes(
|
|
session: AsyncSession, profile_id: uuid.UUID
|
|
) -> ConfigProfile | None:
|
|
"""Fetch a profile with includes eagerly loaded."""
|
|
result = await session.execute(
|
|
select(ConfigProfile)
|
|
.where(ConfigProfile.id == profile_id)
|
|
.options(selectinload(ConfigProfile.includes))
|
|
)
|
|
return result.scalar_one_or_none()
|
|
|
|
|
|
async def check_access(
|
|
session: AsyncSession,
|
|
user_id: uuid.UUID,
|
|
project_id: uuid.UUID | None = None,
|
|
tool_type_id: uuid.UUID | None = None,
|
|
) -> None:
|
|
"""Verify user has access to referenced project and tool type."""
|
|
if project_id is not None:
|
|
project = await session.get(Project, project_id)
|
|
if project is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND, detail="Project not found"
|
|
)
|
|
if tool_type_id is not None:
|
|
tool_type = await session.get(ToolType, tool_type_id)
|
|
if tool_type is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND, detail="Tool type not found"
|
|
)
|
|
|
|
|
|
async def validate_git_mounts(
|
|
session: AsyncSession,
|
|
user_id: uuid.UUID,
|
|
git_mounts: list[Any],
|
|
project_id: uuid.UUID | None = None,
|
|
) -> None:
|
|
"""Validate git mount URLs."""
|
|
for mount in git_mounts:
|
|
remote_url = mount.get("remote_url")
|
|
if not remote_url:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail="Git mount missing remote_url",
|
|
)
|
|
if not remote_url.startswith(("http://", "https://", "git@", "ssh://")):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=f"Invalid git URL: {remote_url}",
|
|
)
|
|
|
|
|
|
def profile_to_response(
|
|
profile: ConfigProfile, includes: list[ConfigProfileInclude] | None = None
|
|
) -> dict:
|
|
return {
|
|
"id": str(profile.id),
|
|
"user_id": str(profile.user_id),
|
|
"name": profile.name,
|
|
"description": profile.description,
|
|
"project_id": str(profile.project_id) if profile.project_id else None,
|
|
"tool_type_id": str(profile.tool_type_id) if profile.tool_type_id else None,
|
|
"env_vars": profile.env_vars or {},
|
|
"runtime_hints": profile.runtime_hints or {},
|
|
"mounts": profile.mounts or [],
|
|
"git_mounts": profile.git_mounts or [],
|
|
"files": profile.files or {},
|
|
"is_default": profile.is_default,
|
|
"includes": [
|
|
{
|
|
"id": str(inc.id),
|
|
"included_profile_id": str(inc.included_profile_id),
|
|
"order_index": inc.order_index,
|
|
}
|
|
for inc in (includes or profile.includes)
|
|
],
|
|
"created_at": profile.created_at.isoformat() if profile.created_at else None,
|
|
"updated_at": profile.updated_at.isoformat() if profile.updated_at else None,
|
|
}
|
|
|
|
|
|
async def get_or_create_user_config(
|
|
session: AsyncSession,
|
|
user_id: uuid.UUID,
|
|
) -> UserConfig:
|
|
"""Get existing user config or create a new one."""
|
|
result = await session.execute(
|
|
select(UserConfig).where(UserConfig.user_id == user_id)
|
|
)
|
|
user_config = result.scalar_one_or_none()
|
|
if user_config is None:
|
|
user_config = UserConfig(user_id=user_id, config={})
|
|
session.add(user_config)
|
|
return user_config
|
|
|
|
|
|
async def validate_default_profiles(
|
|
session: AsyncSession,
|
|
user_id: uuid.UUID,
|
|
default_profiles: dict[str, str],
|
|
) -> None:
|
|
"""Validate that all profile IDs in default_profiles belong to the user."""
|
|
for tool_type_id, profile_id_str in default_profiles.items():
|
|
try:
|
|
profile_uuid = uuid.UUID(profile_id_str)
|
|
except ValueError:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=f"Invalid profile ID for tool type {tool_type_id}: {profile_id_str}",
|
|
)
|
|
profile = await session.get(ConfigProfile, profile_uuid)
|
|
if profile is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail=f"Profile not found: {profile_id_str}",
|
|
)
|
|
if profile.user_id != user_id:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail=f"Profile does not belong to user: {profile_id_str}",
|
|
)
|
|
|
|
|
|
async def create_profile(
|
|
session: AsyncSession,
|
|
user_id: uuid.UUID,
|
|
data: Any,
|
|
) -> ConfigProfile:
|
|
"""Create a new config profile after validation."""
|
|
existing = await session.execute(
|
|
select(ConfigProfile)
|
|
.where(
|
|
ConfigProfile.user_id == user_id,
|
|
ConfigProfile.name == data.name,
|
|
)
|
|
.options(selectinload(ConfigProfile.includes))
|
|
)
|
|
if existing.scalar_one_or_none() is not None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_409_CONFLICT,
|
|
detail=f"Profile with name '{data.name}' already exists",
|
|
)
|
|
|
|
project_uuid = uuid.UUID(data.project_id) if data.project_id else None
|
|
tool_uuid = uuid.UUID(data.tool_type_id) if data.tool_type_id else None
|
|
await check_access(session, user_id, project_uuid, tool_uuid)
|
|
|
|
if data.git_mounts:
|
|
git_mounts_data = [
|
|
m.model_dump() if hasattr(m, "model_dump") else m for m in data.git_mounts
|
|
]
|
|
await validate_git_mounts(session, user_id, git_mounts_data, project_uuid)
|
|
|
|
size = calculate_profile_size(data.model_dump())
|
|
if size > MAX_PROFILE_SIZE_BYTES:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
|
detail="Profile size exceeds 10MB limit",
|
|
)
|
|
|
|
profile = ConfigProfile(
|
|
user_id=user_id,
|
|
name=data.name,
|
|
description=data.description,
|
|
project_id=project_uuid,
|
|
tool_type_id=tool_uuid,
|
|
env_vars=data.env_vars,
|
|
runtime_hints=data.runtime_hints,
|
|
mounts=[m.model_dump() for m in data.mounts],
|
|
git_mounts=[m.model_dump() for m in data.git_mounts],
|
|
files=data.files,
|
|
is_default=data.is_default,
|
|
)
|
|
session.add(profile)
|
|
await session.commit()
|
|
|
|
result = await session.execute(
|
|
select(ConfigProfile)
|
|
.where(ConfigProfile.id == profile.id)
|
|
.options(selectinload(ConfigProfile.includes))
|
|
)
|
|
return result.scalar_one()
|
|
|
|
|
|
async def update_profile(
|
|
session: AsyncSession,
|
|
profile: ConfigProfile,
|
|
data: Any,
|
|
) -> ConfigProfile:
|
|
"""Update a config profile after validation."""
|
|
update_data = data.model_dump(exclude_unset=True)
|
|
|
|
if "name" in update_data:
|
|
existing = await session.execute(
|
|
select(ConfigProfile).where(
|
|
ConfigProfile.user_id == profile.user_id,
|
|
ConfigProfile.name == update_data["name"],
|
|
ConfigProfile.id != profile.id,
|
|
)
|
|
)
|
|
if existing.scalar_one_or_none() is not None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_409_CONFLICT,
|
|
detail=f"Profile with name '{update_data['name']}' already exists",
|
|
)
|
|
|
|
project_uuid = (
|
|
uuid.UUID(update_data["project_id"])
|
|
if "project_id" in update_data and update_data["project_id"]
|
|
else (profile.project_id if "project_id" not in update_data else None)
|
|
)
|
|
tool_uuid = (
|
|
uuid.UUID(update_data["tool_type_id"])
|
|
if "tool_type_id" in update_data and update_data["tool_type_id"]
|
|
else (profile.tool_type_id if "tool_type_id" not in update_data else None)
|
|
)
|
|
await check_access(session, profile.user_id, project_uuid, tool_uuid)
|
|
|
|
if "git_mounts" in update_data and update_data["git_mounts"] is not None:
|
|
git_mounts_data = [
|
|
m.model_dump() if hasattr(m, "model_dump") else m
|
|
for m in update_data["git_mounts"]
|
|
]
|
|
await validate_git_mounts(
|
|
session, profile.user_id, git_mounts_data, project_uuid
|
|
)
|
|
|
|
current_data = profile_to_response(profile)
|
|
merged = {**current_data, **update_data}
|
|
size = calculate_profile_size(merged)
|
|
if size > MAX_PROFILE_SIZE_BYTES:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
|
detail="Profile size exceeds 10MB limit",
|
|
)
|
|
|
|
for field_name, value in update_data.items():
|
|
if field_name in ("project_id", "tool_type_id"):
|
|
value = uuid.UUID(value) if value else None
|
|
elif field_name == "mounts" and value is not None:
|
|
value = [m.model_dump() if not isinstance(m, dict) else m for m in value]
|
|
elif field_name == "git_mounts" and value is not None:
|
|
value = [m.model_dump() if not isinstance(m, dict) else m for m in value]
|
|
setattr(profile, field_name, value)
|
|
|
|
await session.commit()
|
|
|
|
result = await session.execute(
|
|
select(ConfigProfile)
|
|
.where(ConfigProfile.id == profile.id)
|
|
.options(selectinload(ConfigProfile.includes))
|
|
)
|
|
return result.scalar_one()
|
|
|
|
|
|
async def update_includes(
|
|
session: AsyncSession,
|
|
profile: ConfigProfile,
|
|
included_ids: list[uuid.UUID],
|
|
user_id: uuid.UUID,
|
|
) -> ConfigProfile:
|
|
"""Replace profile includes after cycle check."""
|
|
for inc_uuid in included_ids:
|
|
inc_profile = await session.get(ConfigProfile, inc_uuid)
|
|
if inc_profile is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail=f"Included profile not found: {inc_uuid}",
|
|
)
|
|
if inc_profile.user_id != user_id:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail=f"Not authorized to include profile: {inc_uuid}",
|
|
)
|
|
if inc_uuid == profile.id:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail="Profile cannot include itself",
|
|
)
|
|
|
|
from src.services.config.config_profile_resolver import check_include_cycle
|
|
|
|
cycle = await check_include_cycle(session, profile.id, None)
|
|
if cycle is None and included_ids:
|
|
for inc_uuid in included_ids:
|
|
cycle = await check_include_cycle(session, profile.id, inc_uuid)
|
|
if cycle is not None:
|
|
break
|
|
|
|
if cycle is not None:
|
|
cycle_str = " -> ".join(str(c) for c in cycle)
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=f"Include cycle detected: {cycle_str}",
|
|
)
|
|
|
|
result = await session.execute(
|
|
select(ConfigProfileInclude).where(
|
|
ConfigProfileInclude.profile_id == profile.id
|
|
)
|
|
)
|
|
for existing in result.scalars().all():
|
|
await session.delete(existing)
|
|
await session.flush()
|
|
|
|
for order_index, inc_uuid in enumerate(included_ids):
|
|
include = ConfigProfileInclude(
|
|
profile_id=profile.id,
|
|
included_profile_id=inc_uuid,
|
|
order_index=order_index,
|
|
)
|
|
session.add(include)
|
|
await session.flush()
|
|
await session.commit()
|
|
|
|
result = await session.execute(
|
|
select(ConfigProfile).where(ConfigProfile.id == profile.id)
|
|
)
|
|
return result.scalar_one()
|