feat(agent): 增强多模态链路与工具调用能力
This commit is contained in:
@@ -1,3 +1,7 @@
|
||||
from core.agentscope.tools.custom.calendar import calendar_read, calendar_write
|
||||
from core.agentscope.tools.custom.calendar import (
|
||||
calendar_read,
|
||||
calendar_write,
|
||||
user_resolve,
|
||||
)
|
||||
|
||||
__all__ = ["calendar_read", "calendar_write"]
|
||||
__all__ = ["calendar_read", "calendar_write", "user_resolve"]
|
||||
|
||||
@@ -7,6 +7,7 @@ from core.auth.jwt_verifier import JwtVerifier, TokenValidationError
|
||||
from core.agentscope.tools.custom.calendar_backend_ops import (
|
||||
_execute_list_calendar_events,
|
||||
_execute_mutate_calendar_event,
|
||||
_execute_resolve_user_identity,
|
||||
)
|
||||
from core.config.settings import config
|
||||
from core.agentscope.tools.response import build_tool_response
|
||||
@@ -150,6 +151,30 @@ async def calendar_write(
|
||||
bool,
|
||||
Field(description="Whether to use the replace strategy for conflicts."),
|
||||
] = False,
|
||||
invite_user_emails: Annotated[
|
||||
list[str] | None,
|
||||
Field(description="Optional invite targets by email."),
|
||||
] = None,
|
||||
invite_user_names: Annotated[
|
||||
list[str] | None,
|
||||
Field(description="Optional invite targets by username."),
|
||||
] = None,
|
||||
invite_user_ids: Annotated[
|
||||
list[str] | None,
|
||||
Field(description="Optional invite targets by user ID (UUID string)."),
|
||||
] = None,
|
||||
invite_permission_view: Annotated[
|
||||
bool,
|
||||
Field(description="Invite permission: view."),
|
||||
] = True,
|
||||
invite_permission_edit: Annotated[
|
||||
bool,
|
||||
Field(description="Invite permission: edit."),
|
||||
] = False,
|
||||
invite_permission_invite: Annotated[
|
||||
bool,
|
||||
Field(description="Invite permission: invite others."),
|
||||
] = False,
|
||||
session: Any = None,
|
||||
owner_id: Any = None,
|
||||
user_token: str | None = None,
|
||||
@@ -240,6 +265,15 @@ async def calendar_write(
|
||||
tool_args["reminderMinutes"] = reminder_minutes
|
||||
if status is not None:
|
||||
tool_args["status"] = status
|
||||
if invite_user_emails is not None:
|
||||
tool_args["inviteUserEmails"] = invite_user_emails
|
||||
if invite_user_names is not None:
|
||||
tool_args["inviteUserNames"] = invite_user_names
|
||||
if invite_user_ids is not None:
|
||||
tool_args["inviteUserIds"] = invite_user_ids
|
||||
tool_args["invitePermissionView"] = invite_permission_view
|
||||
tool_args["invitePermissionEdit"] = invite_permission_edit
|
||||
tool_args["invitePermissionInvite"] = invite_permission_invite
|
||||
|
||||
result = await _execute_mutate_calendar_event(
|
||||
session=cast(Any, session),
|
||||
@@ -247,3 +281,34 @@ async def calendar_write(
|
||||
tool_args=tool_args,
|
||||
)
|
||||
return build_tool_response(result)
|
||||
|
||||
|
||||
async def user_resolve(
|
||||
user_email: Annotated[
|
||||
str | None,
|
||||
Field(description="User email to resolve user ID."),
|
||||
] = None,
|
||||
user_name: Annotated[
|
||||
str | None,
|
||||
Field(description="Username to resolve user ID."),
|
||||
] = None,
|
||||
session: Any = None,
|
||||
owner_id: Any = None,
|
||||
user_token: str | None = None,
|
||||
) -> Any:
|
||||
if session is None or owner_id is None:
|
||||
raise ValueError("user.resolve missing runtime preset arguments")
|
||||
if not isinstance(user_token, str) or not user_token.strip():
|
||||
return build_tool_response(_unauthorized_response())
|
||||
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
|
||||
return build_tool_response(_unauthorized_response())
|
||||
|
||||
result = await _execute_resolve_user_identity(
|
||||
session=cast(Any, session),
|
||||
owner_id=cast(UUID, owner_id),
|
||||
tool_args={
|
||||
"userEmail": user_email,
|
||||
"userName": user_name,
|
||||
},
|
||||
)
|
||||
return build_tool_response(result)
|
||||
|
||||
@@ -4,13 +4,20 @@ import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.auth.models import CurrentUser
|
||||
from services.base.supabase import supabase_service
|
||||
from models.profile import Profile
|
||||
from v1.auth.gateway import SupabaseAuthGateway
|
||||
from v1.inbox_messages.repository import SQLAlchemyInboxMessageRepository
|
||||
from v1.schedule_items.repository import SQLAlchemyScheduleItemRepository
|
||||
from v1.schedule_items.schemas import (
|
||||
ScheduleItemCreateRequest,
|
||||
ScheduleItemMetadata,
|
||||
ScheduleItemShareRequest,
|
||||
ScheduleItemStatus,
|
||||
ScheduleItemUpdateRequest,
|
||||
)
|
||||
@@ -72,9 +79,196 @@ def _service(session: AsyncSession, owner_id: UUID) -> ScheduleItemService:
|
||||
repository=SQLAlchemyScheduleItemRepository(session),
|
||||
session=session,
|
||||
current_user=CurrentUser(id=owner_id),
|
||||
inbox_repository=SQLAlchemyInboxMessageRepository(session),
|
||||
)
|
||||
|
||||
|
||||
def _parse_string_list(value: object, *, field_name: str) -> list[str]:
|
||||
if value is None:
|
||||
return []
|
||||
if not isinstance(value, list):
|
||||
raise ValueError(f"{field_name} must be a list of strings")
|
||||
parsed: list[str] = []
|
||||
for item in value:
|
||||
if not isinstance(item, str) or not item.strip():
|
||||
raise ValueError(f"{field_name} must be a list of non-empty strings")
|
||||
parsed.append(item.strip())
|
||||
return parsed
|
||||
|
||||
|
||||
def _list_auth_users() -> list[object]:
|
||||
admin_client = supabase_service.get_admin_client()
|
||||
users: list[object] = []
|
||||
page = 1
|
||||
while page <= 100:
|
||||
response = admin_client.auth.admin.list_users(page=page, per_page=100)
|
||||
batch = (
|
||||
list(response)
|
||||
if isinstance(response, list)
|
||||
else list(getattr(response, "users", []))
|
||||
)
|
||||
users.extend(batch)
|
||||
if len(batch) < 100:
|
||||
break
|
||||
page += 1
|
||||
return users
|
||||
|
||||
|
||||
async def _get_profile_username(*, session: AsyncSession, user_id: UUID) -> str | None:
|
||||
stmt = select(Profile.username).where(Profile.id == user_id)
|
||||
return (await session.execute(stmt)).scalar_one_or_none()
|
||||
|
||||
|
||||
async def _get_profile_by_username(
|
||||
*, session: AsyncSession, username: str
|
||||
) -> Profile | None:
|
||||
stmt = (
|
||||
select(Profile)
|
||||
.where(Profile.username == username)
|
||||
.where(Profile.deleted_at.is_(None))
|
||||
)
|
||||
return (await session.execute(stmt)).scalar_one_or_none()
|
||||
|
||||
|
||||
def _find_auth_email_by_user_id(*, users: list[object], user_id: UUID) -> str | None:
|
||||
target = str(user_id)
|
||||
for user in users:
|
||||
if str(getattr(user, "id", "")) == target:
|
||||
email = getattr(user, "email", None)
|
||||
if isinstance(email, str) and email.strip():
|
||||
return email.strip()
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_identity(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
user_email: str | None,
|
||||
user_name: str | None,
|
||||
) -> dict[str, object]:
|
||||
email = user_email.strip().lower() if isinstance(user_email, str) else ""
|
||||
name = user_name.strip() if isinstance(user_name, str) else ""
|
||||
if bool(email) == bool(name):
|
||||
raise ValueError("provide exactly one of user_email or user_name")
|
||||
|
||||
if email:
|
||||
auth_gateway = SupabaseAuthGateway()
|
||||
user = await auth_gateway.get_user_by_email(email)
|
||||
user_id = UUID(user.id)
|
||||
username = await _get_profile_username(session=session, user_id=user_id)
|
||||
return {
|
||||
"userId": str(user_id),
|
||||
"email": user.email,
|
||||
"username": username,
|
||||
"matchedBy": "email",
|
||||
}
|
||||
|
||||
profile = await _get_profile_by_username(session=session, username=name)
|
||||
if profile is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
users = _list_auth_users()
|
||||
email_value = _find_auth_email_by_user_id(users=users, user_id=profile.id)
|
||||
return {
|
||||
"userId": str(profile.id),
|
||||
"email": email_value,
|
||||
"username": profile.username,
|
||||
"matchedBy": "username",
|
||||
}
|
||||
|
||||
|
||||
def _invite_permission(tool_args: dict[str, object]) -> dict[str, bool]:
|
||||
return {
|
||||
"permission_view": bool(tool_args.get("invitePermissionView", True)),
|
||||
"permission_edit": bool(tool_args.get("invitePermissionEdit", False)),
|
||||
"permission_invite": bool(tool_args.get("invitePermissionInvite", False)),
|
||||
}
|
||||
|
||||
|
||||
async def _share_event_with_invitees(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
owner_id: UUID,
|
||||
event_id: UUID,
|
||||
tool_args: dict[str, object],
|
||||
) -> dict[str, object] | None:
|
||||
email_targets = _parse_string_list(
|
||||
tool_args.get("inviteUserEmails"),
|
||||
field_name="inviteUserEmails",
|
||||
)
|
||||
name_targets = _parse_string_list(
|
||||
tool_args.get("inviteUserNames"),
|
||||
field_name="inviteUserNames",
|
||||
)
|
||||
id_targets = _parse_string_list(
|
||||
tool_args.get("inviteUserIds"),
|
||||
field_name="inviteUserIds",
|
||||
)
|
||||
if not email_targets and not name_targets and not id_targets:
|
||||
return None
|
||||
|
||||
users = _list_auth_users() if id_targets else []
|
||||
emails = {item.lower() for item in email_targets}
|
||||
for user_id_raw in id_targets:
|
||||
try:
|
||||
user_id = UUID(user_id_raw)
|
||||
except ValueError as exc:
|
||||
raise ValueError("inviteUserIds must contain valid UUID strings") from exc
|
||||
resolved_email = _find_auth_email_by_user_id(users=users, user_id=user_id)
|
||||
if resolved_email is None:
|
||||
raise HTTPException(status_code=404, detail="Invite user email not found")
|
||||
emails.add(resolved_email.lower())
|
||||
for username in name_targets:
|
||||
resolved = await _resolve_identity(
|
||||
session=session,
|
||||
user_email=None,
|
||||
user_name=username,
|
||||
)
|
||||
resolved_email = resolved.get("email")
|
||||
if not isinstance(resolved_email, str) or not resolved_email:
|
||||
raise HTTPException(status_code=404, detail="Invite user email not found")
|
||||
emails.add(resolved_email.lower())
|
||||
|
||||
service = _service(session, owner_id)
|
||||
permission = _invite_permission(tool_args)
|
||||
invited: list[str] = []
|
||||
for email in sorted(emails):
|
||||
request = ScheduleItemShareRequest(email=email, **permission)
|
||||
await service.share(event_id, request)
|
||||
invited.append(email)
|
||||
return {
|
||||
"count": len(invited),
|
||||
"emails": invited,
|
||||
"permission": permission,
|
||||
}
|
||||
|
||||
|
||||
async def _execute_resolve_user_identity(
|
||||
*,
|
||||
session: AsyncSession,
|
||||
owner_id: UUID,
|
||||
tool_args: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
del owner_id
|
||||
user_email_raw = tool_args.get("userEmail")
|
||||
user_name_raw = tool_args.get("userName")
|
||||
user_email = user_email_raw if isinstance(user_email_raw, str) else None
|
||||
user_name = user_name_raw if isinstance(user_name_raw, str) else None
|
||||
resolved = await _resolve_identity(
|
||||
session=session,
|
||||
user_email=user_email,
|
||||
user_name=user_name,
|
||||
)
|
||||
return {
|
||||
"type": "user_lookup.v1",
|
||||
"version": "v1",
|
||||
"data": {
|
||||
"ok": True,
|
||||
**resolved,
|
||||
},
|
||||
"actions": [],
|
||||
}
|
||||
|
||||
|
||||
def _resolve_metadata(tool_args: dict[str, object]) -> ScheduleItemMetadata:
|
||||
location = tool_args.get("location")
|
||||
location_value = location.strip() if isinstance(location, str) else None
|
||||
@@ -185,6 +379,12 @@ async def _execute_create(
|
||||
)
|
||||
event_data = _event_payload(created)
|
||||
event_id = str(event_data["id"])
|
||||
invite_result = await _share_event_with_invitees(
|
||||
session=service._session,
|
||||
owner_id=service.require_user_id(),
|
||||
event_id=UUID(event_id),
|
||||
tool_args=tool_args,
|
||||
)
|
||||
return {
|
||||
"type": "calendar_card.v1",
|
||||
"version": "v1",
|
||||
@@ -193,12 +393,13 @@ async def _execute_create(
|
||||
"sourceType": "agent_generated",
|
||||
"ok": True,
|
||||
"message": "日程已创建",
|
||||
"inviteResult": invite_result,
|
||||
},
|
||||
"actions": [
|
||||
{
|
||||
"type": "link",
|
||||
"label": "查看详情",
|
||||
"target": f"/calendar/events/{event_id}",
|
||||
"target": f"/schedule-items/{event_id}",
|
||||
}
|
||||
],
|
||||
}
|
||||
@@ -274,6 +475,12 @@ async def _execute_update(
|
||||
ScheduleItemUpdateRequest.model_validate(update_data),
|
||||
)
|
||||
event_data = _event_payload(updated)
|
||||
invite_result = await _share_event_with_invitees(
|
||||
session=service._session,
|
||||
owner_id=service.require_user_id(),
|
||||
event_id=UUID(str(event_data["id"])),
|
||||
tool_args=tool_args,
|
||||
)
|
||||
return {
|
||||
"type": "calendar_card.v1",
|
||||
"version": "v1",
|
||||
@@ -282,12 +489,13 @@ async def _execute_update(
|
||||
"sourceType": "agent_generated",
|
||||
"ok": True,
|
||||
"message": "日程已更新",
|
||||
"inviteResult": invite_result,
|
||||
},
|
||||
"actions": [
|
||||
{
|
||||
"type": "link",
|
||||
"label": "查看详情",
|
||||
"target": f"/calendar/events/{event_data['id']}",
|
||||
"target": f"/schedule-items/{event_data['id']}",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
@@ -4,8 +4,9 @@ from dataclasses import dataclass
|
||||
|
||||
|
||||
TOOL_APPROVAL_REQUIRED: dict[str, bool] = {
|
||||
"calendar.read": False,
|
||||
"calendar.write": False,
|
||||
"calendar_read": False,
|
||||
"calendar_write": False,
|
||||
"user_resolve": False,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,11 @@ from uuid import UUID
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from core.agentscope.tools.custom.calendar import calendar_read, calendar_write
|
||||
from core.agentscope.tools.custom.calendar import (
|
||||
calendar_read,
|
||||
calendar_write,
|
||||
user_resolve,
|
||||
)
|
||||
from core.agentscope.tools.hitl_middleware import register_tool_middlewares
|
||||
from core.agentscope.tools.tool_meta import TOOL_META
|
||||
|
||||
@@ -25,10 +29,12 @@ class ToolGroup:
|
||||
|
||||
|
||||
TOOL_GROUPS: dict[str, ToolGroup] = {
|
||||
"intent": ToolGroup(stage="intent", tool_names=frozenset({"calendar.read"})),
|
||||
"intent": ToolGroup(
|
||||
stage="intent", tool_names=frozenset({"calendar_read", "user_resolve"})
|
||||
),
|
||||
"execution": ToolGroup(
|
||||
stage="execution",
|
||||
tool_names=frozenset({"calendar.read", "calendar.write"}),
|
||||
tool_names=frozenset({"calendar_read", "calendar_write", "user_resolve"}),
|
||||
),
|
||||
"report": ToolGroup(stage="report", tool_names=frozenset()),
|
||||
}
|
||||
@@ -49,7 +55,7 @@ def _load_custom_tool_bindings(
|
||||
) -> list[CustomToolBinding]:
|
||||
return [
|
||||
CustomToolBinding(
|
||||
name="calendar.read",
|
||||
name="calendar_read",
|
||||
func=calendar_read,
|
||||
preset_kwargs={
|
||||
"session": session,
|
||||
@@ -58,7 +64,7 @@ def _load_custom_tool_bindings(
|
||||
},
|
||||
),
|
||||
CustomToolBinding(
|
||||
name="calendar.write",
|
||||
name="calendar_write",
|
||||
func=calendar_write,
|
||||
preset_kwargs={
|
||||
"session": session,
|
||||
@@ -66,6 +72,15 @@ def _load_custom_tool_bindings(
|
||||
"user_token": user_token or "",
|
||||
},
|
||||
),
|
||||
CustomToolBinding(
|
||||
name="user_resolve",
|
||||
func=user_resolve,
|
||||
preset_kwargs={
|
||||
"session": session,
|
||||
"owner_id": owner_id,
|
||||
"user_token": user_token or "",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user