Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
20 changes: 12 additions & 8 deletions src/ai/backend/manager/api/gql_legacy/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,10 +46,13 @@
ScopeType,
)
from ai.backend.manager.models.rbac.context import ClientContext
from ai.backend.manager.models.resource_slot import AgentResourceRow
from ai.backend.manager.models.scaling_group import ScalingGroupRow
from ai.backend.manager.models.user import UserRole, users
from ai.backend.manager.repositories.agent.query import QueryConditions, QueryOrders
from ai.backend.manager.repositories.agent.query import (
QueryConditions,
QueryOrders,
fetch_actual_occupied_slots,
)
from ai.backend.manager.services.agent.actions.update_resource_group import (
UpdateAgentResourceGroupAction,
)
Expand Down Expand Up @@ -783,11 +786,6 @@ async def batch_load(
query = (
sa.select(AgentRow)
.where(AgentRow.id.in_(agent_ids))
.options(
sa.orm.selectinload(AgentRow.agent_resource_rows).joinedload(
AgentResourceRow.slot_type_row
)
)
.order_by(
AgentRow.id,
)
Expand All @@ -800,7 +798,13 @@ async def batch_load(
async with graph_ctx.db.begin_readonly_session() as session:
result = await session.scalars(query)
agent_list = result.unique().all()
return [cls.from_data(agent.to_data()) for agent in agent_list]
occupied_slots = await fetch_actual_occupied_slots(
session, [AgentId(agent.id) for agent in agent_list]
)
return [
cls.from_data(agent.to_data(occupied_slots[AgentId(agent.id)]))
for agent in agent_list
]

@classmethod
async def load_count(
Expand Down
23 changes: 23 additions & 0 deletions src/ai/backend/manager/api/gql_legacy/image.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,7 @@
ResourceLimit,
ResourceLimitInput,
batch_multiresult_in_scalar_stream,
batch_result_in_scalar_stream,
extract_object_uuid,
generate_sql_info_for_gql_connection,
)
Expand Down Expand Up @@ -500,6 +501,28 @@ async def batch_load_by_name_and_arch(
lambda row: (row.name, row.architecture),
)

@classmethod
async def batch_load_by_ids(
cls,
graph_ctx: GraphQueryContext,
image_ids: Sequence[ImageID],
) -> Sequence[ImageNode | None]:
"""Load images by row id, regardless of status."""
query = (
sa.select(ImageRow)
.where(ImageRow.id.in_(image_ids))
.options(selectinload(ImageRow.aliases))
)
async with graph_ctx.db.begin_readonly_session() as db_session:
return await batch_result_in_scalar_stream(
graph_ctx,
db_session,
query,
cls,
image_ids,
lambda row: row.id,
)

@classmethod
async def batch_load_by_image_identifier(
cls,
Expand Down
35 changes: 18 additions & 17 deletions src/ai/backend/manager/api/gql_legacy/kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,20 +15,20 @@
from dateutil.parser import parse as dtparse
from graphene.types.datetime import DateTime as GQLDateTime
from sqlalchemy.engine.row import Row
from sqlalchemy.orm import noload, selectinload
from sqlalchemy.orm import noload

from ai.backend.common.types import (
AccessKey,
AgentId,
BinarySize,
ImageID,
KernelId,
SessionId,
)
from ai.backend.manager.api.gql_legacy.stat_converter import LegacyLiveStatConverter
from ai.backend.manager.data.kernel.types import KernelStatus
from ai.backend.manager.defs import DEFAULT_ROLE
from ai.backend.manager.models.group import groups
from ai.backend.manager.models.image import ImageRow
from ai.backend.manager.models.kernel import (
AGENT_RESOURCE_OCCUPYING_KERNEL_STATUSES,
DEFAULT_KERNEL_ORDERING,
Expand Down Expand Up @@ -220,6 +220,8 @@ class ComputeContainer(graphene.ObjectType): # type: ignore[misc]
class Meta:
interfaces = (Item,)

_image_id: ImageID | None = None

# identity
idx = graphene.Int() # legacy
role = graphene.String() # legacy
Expand Down Expand Up @@ -286,7 +288,6 @@ def parse_row(cls, ctx: GraphQueryContext, row: KernelRow) -> Mapping[str, Any]:
"session_id": row.session_id,
# image
"image": row.image,
"image_object": ImageNode.from_row(ctx, row.image_row),
"architecture": row.architecture,
"registry": row.registry,
# status
Expand Down Expand Up @@ -314,7 +315,18 @@ def from_row(cls, ctx: GraphQueryContext, row: KernelRow | None) -> ComputeConta
if row is None:
return None
props = cls.parse_row(ctx, row)
return cls(**props)
obj = cls(**props)
obj._image_id = row.image_id
return obj

async def resolve_image_object(self, info: graphene.ResolveInfo) -> ImageNode | None:
if self._image_id is None:
return None
graph_ctx: GraphQueryContext = info.context
loader = graph_ctx.dataloader_manager.get_loader_by_func(
graph_ctx, ImageNode.batch_load_by_ids
)
return cast(ImageNode | None, await loader.load(self._image_id))

# last_stat also fetches data from Redis, meaning that
# both live_stat and last_stat will reference same data from same source
Expand Down Expand Up @@ -425,7 +437,6 @@ async def load_slice(
.where(KernelRow.session_id == session_id)
.limit(limit)
.offset(offset)
.options(selectinload(KernelRow.image_row).options(selectinload(ImageRow.aliases)))
)
if cluster_role is not None:
query = query.where(KernelRow.cluster_role == cluster_role)
Expand Down Expand Up @@ -456,7 +467,6 @@ async def batch_load_by_session(
sa.select(KernelRow)
# TODO: use "owner session ID" when we implement multi-container session
.where(KernelRow.session_id.in_(session_ids))
.options(selectinload(KernelRow.image_row).options(selectinload(ImageRow.aliases)))
)
async with ctx.db.begin_readonly_session() as conn:
return await batch_multiresult(
Expand All @@ -476,11 +486,7 @@ async def batch_load_by_agent_id(
*,
status: KernelStatus | None = None,
) -> Sequence[Sequence[ComputeContainer]]:
query_stmt = (
sa.select(KernelRow)
.where(KernelRow.agent.in_(agent_ids))
.options(selectinload(KernelRow.image_row).options(selectinload(ImageRow.aliases)))
)
query_stmt = sa.select(KernelRow).where(KernelRow.agent.in_(agent_ids))
kernel_status: tuple[KernelStatus, ...]
if status is not None:
kernel_status = (status,)
Expand Down Expand Up @@ -511,12 +517,7 @@ async def batch_load_detail(
.where(
(KernelRow.id.in_(container_ids)),
)
.options(
noload("*"),
selectinload(KernelRow.group_row),
selectinload(KernelRow.user_row),
selectinload(KernelRow.image_row),
)
.options(noload("*"))
)
if domain_name is not None:
query = query.where(KernelRow.domain_name == domain_name)
Expand Down
4 changes: 2 additions & 2 deletions src/ai/backend/manager/api/gql_legacy/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from graphql.type import GraphQLField, get_named_type, is_leaf_type
from opentelemetry import trace
from opentelemetry.trace import StatusCode
from sqlalchemy.orm import joinedload, selectinload
from sqlalchemy.orm import selectinload

from ai.backend.common.clients.valkey_client.valkey_image.client import ValkeyImageClient
from ai.backend.common.clients.valkey_client.valkey_live.client import ValkeyLiveClient
Expand Down Expand Up @@ -2615,7 +2615,7 @@ async def resolve_session_pending_queue(
stmt = (
sa.select(SessionRow)
.where(SessionRow.id.in_(pending_sessions))
.options(selectinload(SessionRow.kernels), joinedload(SessionRow.user))
.options(selectinload(SessionRow.kernels))
)
query_result = await db_session.scalars(stmt)
for row in query_result:
Expand Down
50 changes: 23 additions & 27 deletions src/ai/backend/manager/api/gql_legacy/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
from dateutil.parser import parse as dtparse
from graphene.types.datetime import DateTime as GQLDateTime
from sqlalchemy.engine.row import Row
from sqlalchemy.orm import joinedload, selectinload
from sqlalchemy.orm import selectinload

from ai.backend.common import validators as tx
from ai.backend.common.defs.session import SESSION_PRIORITY_MAX, SESSION_PRIORITY_MIN
Expand Down Expand Up @@ -67,7 +67,6 @@
from ai.backend.manager.models.types import (
QueryCondition,
QueryOption,
join_by_related_field,
load_related_field,
)
from ai.backend.manager.models.user import UserRole, UserRow
Expand Down Expand Up @@ -194,6 +193,13 @@ def parse_value(value: str) -> ComputeSessionPermission:
return ComputeSessionPermission(value)


def _join_owner_rows(stmt: sa.sql.Select[Any]) -> sa.sql.Select[Any]:
"""Join the owning user and project so their columns are filterable/orderable."""
return stmt.join(UserRow, SessionRow.user_uuid == UserRow.uuid).join(
GroupRow, SessionRow.group_id == GroupRow.id
)


class _HasVFID(Protocol):
vfid: VFolderID

Expand Down Expand Up @@ -310,7 +316,7 @@ async def get_node(
stmt = (
sa.select(SessionRow)
.where(SessionRow.id == uuid.UUID(raw_session_id))
.options(selectinload(SessionRow.kernels), joinedload(SessionRow.user))
.options(selectinload(SessionRow.kernels))
)
query_result = await db_session.scalar(stmt)
if query_result is None:
Expand All @@ -322,21 +328,21 @@ async def get_node(
def _add_basic_options_to_query(
cls, stmt: sa.sql.Select[Any], is_count: bool = False
) -> sa.sql.Select[Any]:
options = [
join_by_related_field(SessionRow.user),
join_by_related_field(SessionRow.group),
]
# The user and project joins back the `full_name`/`user_email`/`group_name` filters.
stmt = _join_owner_rows(stmt)
if not is_count:
options = [
*options,
load_related_field(SessionRow.kernel_load_option()),
load_related_field(SessionRow.user_load_option(already_joined=True)),
load_related_field(SessionRow.project_load_option(already_joined=True)),
]
for option in options:
stmt = option(stmt)
stmt = stmt.options(SessionRow.kernel_load_option())
return stmt

async def resolve_owner(self, info: graphene.ResolveInfo) -> UserNode | None:
if self.user_id is None:
return None
graph_ctx: GraphQueryContext = info.context
loader = graph_ctx.dataloader_manager.get_loader_by_func(
graph_ctx, UserNode.batch_load_by_uuids
)
return cast(UserNode | None, await loader.load(self.user_id))

async def resolve_queue_position(self, info: graphene.ResolveInfo) -> int | None:
if self.status != SessionStatus.PENDING:
return None
Expand Down Expand Up @@ -377,7 +383,6 @@ def from_row(
project_id=row.group_id,
user_id=row.user_uuid,
access_key=row.access_key,
owner=UserNode.from_row(ctx, row.user),
# status
status=row.status.name,
# status_changed=row.status_changed, # FIXME: generated attribute
Expand Down Expand Up @@ -592,13 +597,7 @@ async def resolve_graph(
query = sa.select(dependency_cte.c.id)
session_ids = (await db_sess.execute(query)).scalars().all()
# Get the session rows in the graph
query = (
sa.select(SessionRow)
.where(SessionRow.id.in_(session_ids))
.options(
selectinload(SessionRow.user),
)
)
query = sa.select(SessionRow).where(SessionRow.id.in_(session_ids))
session_rows = list((await db_sess.execute(query)).scalars().all())
await batch_populate_session_occupied_slots(db_sess, session_rows)

Expand Down Expand Up @@ -653,7 +652,6 @@ async def batch_load_by_dependee_id(
.where(SessionRow.id.in_(dependent_ids))
.options(
selectinload(SessionRow.kernels.and_(KernelRow.cluster_role == DEFAULT_ROLE)),
joinedload(SessionRow.user),
)
)
rows = list((await db_sess.execute(sess_query)).unique().scalars().all())
Expand Down Expand Up @@ -698,7 +696,6 @@ async def batch_load_by_dependent_id(
.where(SessionRow.id.in_(dependee_ids))
.options(
selectinload(SessionRow.kernels.and_(KernelRow.cluster_role == DEFAULT_ROLE)),
joinedload(SessionRow.user),
)
)
rows = list((await db_sess.execute(sess_query)).unique().scalars().all())
Expand Down Expand Up @@ -848,8 +845,7 @@ async def get_data(
query_conditions.append(by_resource_group_name(resource_group_name))
query_options: list[QueryOption] = [
load_related_field(SessionRow.kernel_load_option()),
join_by_related_field(SessionRow.user),
join_by_related_field(SessionRow.group),
_join_owner_rows,
]
session_rows = await SessionRow.list_session_by_condition(
query_conditions, query_options, db=ctx.db
Expand Down
18 changes: 18 additions & 0 deletions src/ai/backend/manager/api/gql_legacy/user.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,7 @@
PaginatedList,
batch_multiresult,
batch_result,
batch_result_in_scalar_stream,
generate_sql_info_for_gql_connection,
)
from .gql_relay import AsyncNode, Connection, ConnectionResolverResult
Expand Down Expand Up @@ -174,6 +175,23 @@ class Meta:
description="Added in 25.5.0.",
)

@classmethod
async def batch_load_by_uuids(
cls,
ctx: GraphQueryContext,
user_uuids: Sequence[UUID],
) -> Sequence[Self | None]:
query = sa.select(UserRow).where(UserRow.uuid.in_(user_uuids))
async with ctx.db.begin_readonly_session() as db_session:
return await batch_result_in_scalar_stream(
ctx,
db_session,
query,
cls,
user_uuids,
lambda row: row.uuid,
)

@classmethod
def from_row(cls, ctx: GraphQueryContext, row: UserRow) -> Self:
return cls(
Expand Down
14 changes: 2 additions & 12 deletions src/ai/backend/manager/models/agent/row.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
joinedload,
load_only,
mapped_column,
relationship,
selectinload,
)
from sqlalchemy.sql.expression import false, true
Expand Down Expand Up @@ -132,16 +131,7 @@ class AgentRow(Base):
default=False,
)

agent_resource_rows: Mapped[list[AgentResourceRow]] = relationship("AgentResourceRow")

def actual_occupied_slots(self) -> ResourceSlot:
occupied = ResourceSlot()
sorted_rows = sorted(self.agent_resource_rows, key=lambda r: r.slot_type_row.rank)
for resource_row in sorted_rows:
occupied[resource_row.slot_name] = resource_row.used
return occupied

def to_data(self) -> AgentData:
def to_data(self, actual_occupied_slots: ResourceSlot) -> AgentData:
return AgentData(
id=AgentId(self.id),
status=self.status,
Expand All @@ -151,7 +141,7 @@ def to_data(self) -> AgentData:
schedulable=self.schedulable,
available_slots=self.available_slots,
cached_occupied_slots=self.occupied_slots,
actual_occupied_slots=self.actual_occupied_slots(),
actual_occupied_slots=actual_occupied_slots,
addr=self.addr,
public_host=self.public_host,
first_contact=self.first_contact,
Expand Down
Loading
Loading