Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changes/12436.enhance.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Stop the `endpoint` model row from importing the manager `registry` layer (move the scaling-group check down to its repository) so editing it no longer expands the repository test/typecheck scope.
47 changes: 1 addition & 46 deletions src/ai/backend/manager/models/endpoint/row.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,6 @@

from ai.backend.common.identifier.deployment import DeploymentID
from ai.backend.common.identifier.deployment_revision import DeploymentRevisionID
from ai.backend.common.identifier.project import ProjectID
from ai.backend.common.identifier.replica_group import ReplicaGroupID
from ai.backend.common.identifier.runtime_variant import RuntimeVariantID
from ai.backend.common.types import (
Expand All @@ -40,7 +39,6 @@
AutoScalingMetricSource,
ClusterMode,
ResourceSlot,
SessionTypes,
VFolderID,
VFolderMount,
VFolderMountOptions,
Expand Down Expand Up @@ -80,7 +78,7 @@
ScalingState,
)
from ai.backend.manager.errors.api import InvalidAPIParameters
from ai.backend.manager.errors.common import ObjectNotFound, ServiceUnavailable
from ai.backend.manager.errors.common import ObjectNotFound
from ai.backend.manager.models.base import (
GUID,
Base,
Expand All @@ -89,7 +87,6 @@
StrEnumType,
)
from ai.backend.manager.models.routing import RouteStatus
from ai.backend.manager.models.scaling_group import scaling_groups
from ai.backend.manager.models.storage import StorageSessionManager
from ai.backend.manager.models.vfolder import prepare_vfolder_mounts
from ai.backend.manager.types import MountOptionModel, UserScope
Expand Down Expand Up @@ -1239,48 +1236,6 @@ def apply_model_deployment_modifier(


class ModelServiceHelper:
@staticmethod
async def check_scaling_group(
conn: AsyncConnection,
scaling_group: str,
owner_access_key: AccessKey,
target_domain: str,
target_project: str | ProjectID,
) -> str:
"""
Wrapper of `registry.check_scaling_group()` with additional guards flavored for
model service included
"""
from ai.backend.manager.registry import check_scaling_group

checked_scaling_group = await check_scaling_group(
conn,
scaling_group,
SessionTypes.INFERENCE,
owner_access_key,
target_domain,
target_project,
)

query = (
sa.select(scaling_groups.c.wsproxy_addr, scaling_groups.c.wsproxy_api_token)
.select_from(scaling_groups)
.where(scaling_groups.c.name == checked_scaling_group)
)

result = await conn.execute(query)
sgroup = result.first()
if sgroup is None:
raise ServiceUnavailable("Scaling group not found")
wsproxy_addr = sgroup.wsproxy_addr
if not wsproxy_addr:
raise ServiceUnavailable("No coordinator configured for this resource group")

if not sgroup.wsproxy_api_token:
raise ServiceUnavailable("Scaling group not ready to start model service")

return checked_scaling_group

@staticmethod
async def check_extra_mounts(
conn: AsyncConnection,
Expand Down
49 changes: 46 additions & 3 deletions src/ai/backend/manager/repositories/model_serving/repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import sqlalchemy as sa
from pydantic import HttpUrl
from sqlalchemy.exc import IntegrityError, NoResultFound, StatementError
from sqlalchemy.ext.asyncio import AsyncConnection
from sqlalchemy.ext.asyncio import AsyncSession as SASession
from sqlalchemy.orm import selectinload

Expand All @@ -24,6 +25,7 @@
from ai.backend.common.types import (
AccessKey,
ResourceSlot,
SessionTypes,
)
from ai.backend.manager.config.loader.legacy_etcd_loader import LegacyEtcdLoader
from ai.backend.manager.data.deployment.types import RouteHealthStatus
Expand All @@ -44,7 +46,7 @@
from ai.backend.manager.data.permission.types import RBACElementRef
from ai.backend.manager.data.vfolder.types import VFolderOwnershipType
from ai.backend.manager.errors.api import InvalidAPIParameters
from ai.backend.manager.errors.common import GenericForbidden, ObjectNotFound
from ai.backend.manager.errors.common import GenericForbidden, ObjectNotFound, ServiceUnavailable
from ai.backend.manager.errors.resource import DatabaseConnectionUnavailable
from ai.backend.manager.errors.service import EndpointNotFound
from ai.backend.manager.models.deployment_revision import DeploymentRevisionRow
Expand All @@ -71,6 +73,7 @@
from ai.backend.manager.models.vfolder import VFolderRow, VFolderUsageMode
from ai.backend.manager.models.vfolder.row import query_accessible_vfolders, vfolders
from ai.backend.manager.registry import AgentRegistry
from ai.backend.manager.registry import check_scaling_group as registry_check_scaling_group
from ai.backend.manager.repositories.base import (
BatchQuerier,
Creator,
Expand Down Expand Up @@ -111,6 +114,46 @@ class ModelServingRepository:
def __init__(self, db: ExtendedAsyncSAEngine) -> None:
self._db = db

async def _check_inference_scaling_group(
self,
conn: AsyncConnection,
scaling_group: str,
owner_access_key: AccessKey,
target_domain: str,
target_project: str | ProjectID,
) -> str:
"""
Wrapper of ``registry.check_scaling_group()`` with additional guards flavored for
model service included.
"""
checked_scaling_group = await registry_check_scaling_group(
conn,
scaling_group,
SessionTypes.INFERENCE,
owner_access_key,
target_domain,
target_project,
)

query = (
sa.select(scaling_groups.c.wsproxy_addr, scaling_groups.c.wsproxy_api_token)
.select_from(scaling_groups)
.where(scaling_groups.c.name == checked_scaling_group)
)

result = await conn.execute(query)
sgroup = result.first()
if sgroup is None:
raise ServiceUnavailable("Scaling group not found")
wsproxy_addr = sgroup.wsproxy_addr
if not wsproxy_addr:
raise ServiceUnavailable("No coordinator configured for this resource group")

if not sgroup.wsproxy_api_token:
raise ServiceUnavailable("Scaling group not ready to start model service")

return checked_scaling_group

@model_serving_repository_resilience.apply()
async def get_endpoint_by_id(self, endpoint_id: uuid.UUID) -> EndpointData | None:
"""
Expand Down Expand Up @@ -822,7 +865,7 @@ async def _do_mutate() -> MutationResult:
if conn is None:
raise DatabaseConnectionUnavailable("Database connection is not available")

await ModelServiceHelper.check_scaling_group(
await self._check_inference_scaling_group(
conn,
endpoint_row.resource_group,
AccessKey(session_owner.main_access_key),
Expand Down Expand Up @@ -948,7 +991,7 @@ async def resolve_model_service_validation_context(
used by ``SchedulerRepository.prepare_vfolder_mounts``.
"""
async with self._db.begin_readonly() as conn:
checked_scaling_group = await ModelServiceHelper.check_scaling_group(
checked_scaling_group = await self._check_inference_scaling_group(
conn,
scaling_group,
owner_access_key,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -382,9 +382,10 @@ async def test_modify_endpoint_fields_does_not_create_revision(

with (
with_user(user_context),
patch(
"ai.backend.manager.repositories.model_serving.repository.ModelServiceHelper",
check_scaling_group=AsyncMock(),
patch.object(
repository,
"_check_inference_scaling_group",
AsyncMock(),
),
):
result = await repository.modify_endpoint_fields(
Expand Down
Loading