refactor: 重整 schemas 作用域并统一用户上下文模型

This commit is contained in:
zl-q
2026-03-13 01:01:54 +08:00
parent f201babb48
commit fb3c649db7
42 changed files with 4205 additions and 2013 deletions
@@ -3,5 +3,13 @@ from core.agentscope.tools.custom.calendar import (
calendar_read,
calendar_write,
)
from core.agentscope.tools.custom.user_lookup import (
user_lookup,
)
__all__ = ["calendar_read", "calendar_write", "calendar_share"]
__all__ = [
"calendar_read",
"calendar_write",
"calendar_share",
"user_lookup",
]
@@ -1,21 +1,38 @@
from __future__ import annotations
import re
from datetime import datetime, timedelta, timezone
from typing import Annotated, Any, Literal, cast
from uuid import UUID
from fastapi import HTTPException
from pydantic import Field
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
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_share_calendar_event,
)
from core.agentscope.tools.tool_response_builder import build_tool_response
from core.agentscope.schemas.ui_schema import (
build_calendar_list,
build_calendar_operation,
from core.agentscope.tools.tool_response_builder import (
build_success_response,
build_error_response,
)
from core.agentscope.schemas.runtime_models import ToolOutputContent
from core.config.settings import config
from core.auth.models import CurrentUser
from services.base.supabase import supabase_service
from models.profile import Profile
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,
)
from v1.schedule_items.service import ScheduleItemService
_HEX_COLOR_PATTERN = re.compile(r"^#[0-9A-Fa-f]{6}$")
def _verify_user_token(*, user_token: str, owner_id: UUID) -> bool:
@@ -35,65 +52,68 @@ def _verify_user_token(*, user_token: str, owner_id: UUID) -> bool:
return isinstance(subject, str) and subject == str(owner_id)
def _failure_response(
*,
card_type: Literal["calendar_event_list.v1", "calendar_operation.v1"],
operation: str | None,
code: str,
message: str,
) -> dict[str, object]:
if card_type == "calendar_event_list.v1":
return build_calendar_list(
items=[],
page=1,
page_size=20,
total=0,
) | {"data": {"ok": False, "code": code, "message": message}}
return build_calendar_operation(
operation=operation or "operation",
ok=False,
message=message,
code=code,
)
def _authorized_or_response(
*,
session: Any,
owner_id: Any,
user_token: str | None,
card_type: Literal["calendar_event_list.v1", "calendar_operation.v1"],
operation: str | None,
) -> tuple[Any, UUID] | dict[str, object]:
if session is None or owner_id is None:
raise ValueError("calendar tool missing runtime preset arguments")
if not isinstance(user_token, str) or not user_token.strip():
return _failure_response(
card_type=card_type,
operation=operation,
code="UNAUTHORIZED",
message="calendar tool requires validated user token",
)
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
return _failure_response(
card_type=card_type,
operation=operation,
code="UNAUTHORIZED",
message="calendar tool requires validated user token",
)
return cast(Any, session), cast(UUID, owner_id)
def _map_exception(exc: Exception) -> tuple[str, str]:
def _map_exception(exc: Exception) -> tuple[str, str, bool]:
"""Map exception to error code, message, and retryable flag."""
if isinstance(exc, HTTPException):
detail = exc.detail
if isinstance(detail, str) and detail.strip():
return "OPERATION_FAILED", detail.strip()
return "OPERATION_FAILED", "calendar operation failed"
return "OPERATION_FAILED", detail.strip(), True
return "OPERATION_FAILED", "日历操作失败", True
if isinstance(exc, ValueError):
return "INVALID_ARGUMENT", str(exc)
return "INTERNAL_ERROR", "calendar operation failed"
return "INVALID_ARGUMENT", str(exc), False
return "INTERNAL_ERROR", "日历操作失败", True
def _create_service(session: AsyncSession, owner_id: UUID) -> ScheduleItemService:
return ScheduleItemService(
repository=SQLAlchemyScheduleItemRepository(session),
session=session,
current_user=CurrentUser(id=owner_id),
inbox_repository=SQLAlchemyInboxMessageRepository(session),
)
def _event_to_dict(event: object) -> dict[str, Any]:
"""Convert ScheduleItem entity to dict."""
event_id = str(getattr(event, "id"))
metadata = getattr(event, "metadata", None)
location_value = getattr(metadata, "location", None)
color_value = getattr(metadata, "color", None) or "#4F46E5"
reminder_minutes_value = getattr(metadata, "reminder_minutes", None)
return {
"id": event_id,
"title": getattr(event, "title"),
"description": getattr(event, "description"),
"startAt": getattr(event, "start_at").isoformat(),
"endAt": getattr(event, "end_at").isoformat()
if getattr(event, "end_at") is not None
else None,
"timezone": getattr(event, "timezone"),
"location": location_value,
"color": color_value,
"reminderMinutes": reminder_minutes_value,
}
def _build_metadata(
location: str | None,
color: str | None,
reminder_minutes: int | None,
) -> ScheduleItemMetadata:
"""Build ScheduleItemMetadata from parameters."""
location_value = location.strip() if location and location.strip() else None
raw_color = color.strip() if color and color.strip() else "#4F46E5"
color_value = raw_color if _HEX_COLOR_PATTERN.match(raw_color) else "#4F46E5"
reminder_value: int | None = None
if reminder_minutes is not None:
if reminder_minutes < 0 or reminder_minutes > 10080:
raise ValueError("reminderMinutes must be 0..10080")
reminder_value = reminder_minutes
return ScheduleItemMetadata(
location=location_value,
color=color_value,
reminder_minutes=reminder_value,
)
async def calendar_read(
@@ -112,37 +132,69 @@ async def calendar_read(
session: Any = None,
owner_id: Any = None,
user_token: str | None = None,
) -> Any:
auth_result = _authorized_or_response(
session=session,
owner_id=owner_id,
user_token=user_token,
card_type="calendar_event_list.v1",
operation=None,
)
if isinstance(auth_result, dict):
return build_tool_response(auth_result)
runtime_session, runtime_owner_id = auth_result
) -> ToolOutputContent:
"""
Read calendar events with optional filtering and pagination.
"""
if session is None or owner_id is None:
return build_error_response(
code="MISSING_RUNTIME_ARGS",
message="日历工具缺少运行时参数",
retryable=False,
)
if not isinstance(user_token, str) or not user_token.strip():
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
try:
result = await _execute_list_calendar_events(
session=runtime_session,
owner_id=runtime_owner_id,
tool_args={"query": query, "page": page, "pageSize": page_size},
service = _create_service(cast(AsyncSession, session), cast(UUID, owner_id))
items, total = await service.list_paginated(page=page, page_size=page_size)
total_pages = max(1, (total + page_size - 1) // page_size) if total else 0
return build_success_response(
title="日程列表",
summary=f"{total} 个日程",
payload={
"ok": True,
"message": "已获取日程列表",
},
items=[_event_to_dict(item) for item in items],
kv_pairs=[
{"key": "total", "label": "总数", "value": total, "copyable": False},
{"key": "page", "label": "当前页", "value": page, "copyable": False},
{
"key": "page_size",
"label": "每页",
"value": page_size,
"copyable": False,
},
{
"key": "total_pages",
"label": "总页数",
"value": total_pages,
"copyable": False,
},
],
)
except Exception as exc:
code, message = _map_exception(exc)
return build_tool_response(
_failure_response(
card_type="calendar_event_list.v1",
operation=None,
code=code,
message=message,
)
code, message, retryable = _map_exception(exc)
return build_error_response(
code=code,
message=message,
retryable=retryable,
)
return build_tool_response(result)
async def calendar_write(
operation: Annotated[
@@ -169,7 +221,7 @@ async def calendar_write(
str | None,
Field(description="Event end time in ISO 8601 format."),
] = None,
timezone: Annotated[
event_timezone: Annotated[
str | None,
Field(description="IANA timezone name for the event.", max_length=50),
] = None,
@@ -197,58 +249,211 @@ async def calendar_write(
session: Any = None,
owner_id: Any = None,
user_token: str | None = None,
) -> Any:
auth_result = _authorized_or_response(
session=session,
owner_id=owner_id,
user_token=user_token,
card_type="calendar_operation.v1",
operation=operation,
)
if isinstance(auth_result, dict):
return build_tool_response(auth_result)
runtime_session, runtime_owner_id = auth_result
) -> ToolOutputContent:
"""
Write calendar event: create, update, or delete.
"""
if session is None or owner_id is None:
return build_error_response(
code="MISSING_RUNTIME_ARGS",
message="日历工具缺少运行时参数",
retryable=False,
)
tool_args: dict[str, object] = {"operation": operation, "replace": replace}
if event_id is not None:
tool_args["eventId"] = event_id
if title is not None:
tool_args["title"] = title
if description is not None:
tool_args["description"] = description
if start_at is not None:
tool_args["startAt"] = start_at
if end_at is not None:
tool_args["endAt"] = end_at
if timezone is not None:
tool_args["timezone"] = timezone
if location is not None:
tool_args["location"] = location
if color is not None:
tool_args["color"] = color
if reminder_minutes is not None:
tool_args["reminderMinutes"] = reminder_minutes
if status is not None:
tool_args["status"] = status
if not isinstance(user_token, str) or not user_token.strip():
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
try:
result = await _execute_mutate_calendar_event(
session=runtime_session,
owner_id=runtime_owner_id,
tool_args=tool_args,
)
except Exception as exc:
code, message = _map_exception(exc)
return build_tool_response(
_failure_response(
card_type="calendar_operation.v1",
operation=operation,
code=code,
message=message,
service = _create_service(cast(AsyncSession, session), cast(UUID, owner_id))
if operation == "create":
parsed_start = _parse_datetime(start_at) if start_at else None
if parsed_start is None:
parsed_start = datetime.now(timezone.utc) + timedelta(hours=1)
parsed_end = _parse_datetime(end_at) if end_at else None
tz = (
event_timezone.strip()
if event_timezone and event_timezone.strip()
else "Asia/Shanghai"
)
created = await service.create_agent_generated(
ScheduleItemCreateRequest(
title=title.strip() if title and title.strip() else "新的日程",
description=description.strip()
if description and description.strip()
else None,
start_at=parsed_start,
end_at=parsed_end,
timezone=tz,
metadata=_build_metadata(location, color, reminder_minutes),
)
)
event_dict = _event_to_dict(created)
return build_success_response(
title="日程已创建",
summary=f"日程「{created.title}」已创建",
payload={"ok": True, "operation": "create"},
items=[event_dict],
kv_pairs=[
{
"key": "title",
"label": "标题",
"value": created.title,
"copyable": True,
},
{
"key": "start_at",
"label": "开始时间",
"value": created.start_at.isoformat(),
"copyable": True,
},
],
)
if operation == "update":
if not event_id:
return build_error_response(
code="INVALID_ARGUMENT",
message="更新日程需要提供 event_id",
retryable=False,
)
parsed_event_id = UUID(event_id)
update_data: dict[str, Any] = {}
if title:
update_data["title"] = title.strip()
if description:
update_data["description"] = description.strip()
if start_at:
update_data["start_at"] = _parse_datetime(start_at)
if end_at:
update_data["end_at"] = _parse_datetime(end_at)
if event_timezone:
update_data["timezone"] = event_timezone.strip()
if status:
try:
update_data["status"] = ScheduleItemStatus(status)
except ValueError:
return build_error_response(
code="INVALID_ARGUMENT",
message="status 必须是 active, completed, canceled, archived 之一",
retryable=False,
)
if location or color or reminder_minutes is not None:
existing = await service.get_by_id(parsed_event_id)
metadata_dump = (
existing.metadata.model_dump() if existing.metadata else {}
)
if location:
metadata_dump["location"] = location.strip() or None
if color:
color_str = color.strip()
if not color_str:
metadata_dump["color"] = None
elif _HEX_COLOR_PATTERN.match(color_str):
metadata_dump["color"] = color_str
else:
return build_error_response(
code="INVALID_ARGUMENT",
message="color 必须是十六进制颜色值如 #4F46E5",
retryable=False,
)
if reminder_minutes is not None:
if reminder_minutes < 0 or reminder_minutes > 10080:
return build_error_response(
code="INVALID_ARGUMENT",
message="reminderMinutes 必须在 0-10080 之间",
retryable=False,
)
metadata_dump["reminder_minutes"] = reminder_minutes
update_data["metadata"] = ScheduleItemMetadata.model_validate(
metadata_dump
)
updated = await service.update(
parsed_event_id, ScheduleItemUpdateRequest.model_validate(update_data)
)
event_dict = _event_to_dict(updated)
return build_success_response(
title="日程已更新",
summary=f"日程「{updated.title}」已更新",
payload={"ok": True, "operation": "update"},
items=[event_dict],
kv_pairs=[
{
"key": "title",
"label": "标题",
"value": updated.title,
"copyable": True,
},
{
"key": "start_at",
"label": "开始时间",
"value": updated.start_at.isoformat(),
"copyable": True,
},
],
)
if operation == "delete":
if not event_id:
return build_error_response(
code="INVALID_ARGUMENT",
message="删除日程需要提供 event_id",
retryable=False,
)
await service.delete(UUID(event_id))
return build_success_response(
title="日程已删除",
summary=f"日程 {event_id} 已删除",
payload={"ok": True, "operation": "delete", "event_id": event_id},
items=[],
kv_pairs=[
{
"key": "event_id",
"label": "已删除日程ID",
"value": event_id,
"copyable": True,
},
],
)
return build_error_response(
code="INVALID_ARGUMENT",
message="无效的操作类型",
retryable=False,
)
return build_tool_response(result)
except Exception as exc:
code, message, retryable = _map_exception(exc)
return build_error_response(
code=code,
message=message,
retryable=retryable,
)
def _parse_datetime(value: str | None) -> datetime | None:
if not value:
return None
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
except ValueError:
return None
async def calendar_share(
@@ -283,46 +488,163 @@ async def calendar_share(
session: Any = None,
owner_id: Any = None,
user_token: str | None = None,
) -> Any:
auth_result = _authorized_or_response(
session=session,
owner_id=owner_id,
user_token=user_token,
card_type="calendar_operation.v1",
operation="share",
)
if isinstance(auth_result, dict):
return build_tool_response(auth_result)
runtime_session, runtime_owner_id = auth_result
) -> ToolOutputContent:
"""
Share a calendar event with other users.
"""
if session is None or owner_id is None:
return build_error_response(
code="MISSING_RUNTIME_ARGS",
message="日历工具缺少运行时参数",
retryable=False,
)
tool_args: dict[str, object] = {
"eventId": event_id,
"invitePermissionView": invite_permission_view,
"invitePermissionEdit": invite_permission_edit,
"invitePermissionInvite": invite_permission_invite,
}
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
if not isinstance(user_token, str) or not user_token.strip():
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
return build_error_response(
code="UNAUTHORIZED",
message="日历工具需要有效的用户令牌",
retryable=False,
)
if not invite_user_emails and not invite_user_names and not invite_user_ids:
return build_error_response(
code="INVALID_ARGUMENT",
message="请提供至少一个邀请目标(邮箱、用户名或用户ID)",
retryable=False,
)
try:
result = await _execute_share_calendar_event(
session=runtime_session,
owner_id=runtime_owner_id,
tool_args=tool_args,
service = _create_service(cast(AsyncSession, session), cast(UUID, owner_id))
target_uuid = UUID(event_id)
emails: set[str] = set()
if invite_user_emails:
emails = {e.strip().lower() for e in invite_user_emails if e and e.strip()}
if invite_user_ids:
users = _list_auth_users()
for uid in invite_user_ids:
try:
user_uuid = UUID(uid)
email = _find_auth_email(users, user_uuid)
if email:
emails.add(email.lower())
except ValueError:
pass
if invite_user_names:
for username in invite_user_names:
if not username or not username.strip():
continue
profile = await _get_profile_by_username(
cast(AsyncSession, session), username.strip()
)
if profile:
users = _list_auth_users()
email = _find_auth_email(users, profile.id)
if email:
emails.add(email.lower())
if not emails:
return build_error_response(
code="NOT_FOUND",
message="未找到任何有效的邀请目标",
retryable=False,
)
permission = {
"permission_view": invite_permission_view,
"permission_edit": invite_permission_edit,
"permission_invite": invite_permission_invite,
}
invited: list[str] = []
for email in sorted(emails):
await service.share(
target_uuid, ScheduleItemShareRequest(email=email, **permission)
)
invited.append(email)
return build_success_response(
title="日程已分享",
summary=f"已邀请 {len(invited)}",
payload={
"ok": True,
"operation": "share",
"invited": invited,
"permission": permission,
},
items=[],
kv_pairs=[
{
"key": "event_id",
"label": "日程ID",
"value": event_id,
"copyable": True,
},
{
"key": "invited_count",
"label": "已邀请人数",
"value": len(invited),
"copyable": False,
},
{
"key": "invited_emails",
"label": "被邀请人",
"value": ", ".join(invited),
"copyable": False,
},
],
)
except Exception as exc:
code, message = _map_exception(exc)
return build_tool_response(
_failure_response(
card_type="calendar_operation.v1",
operation="share",
code=code,
message=message,
)
code, message, retryable = _map_exception(exc)
return build_error_response(
code=code,
message=message,
retryable=retryable,
)
return build_tool_response(result)
def _list_auth_users() -> list[Any]:
admin_client = supabase_service.get_admin_client()
users: list[Any] = []
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
def _find_auth_email(users: list[Any], 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 _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()
@@ -1,557 +0,0 @@
from __future__ import annotations
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,
)
from v1.schedule_items.service import ScheduleItemService
_HEX_COLOR_PATTERN = re.compile(r"^#[0-9A-Fa-f]{6}$")
def _parse_datetime(value: object) -> datetime | None:
if not isinstance(value, str) or not value:
return None
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
except ValueError:
return None
def _parse_positive_int(
value: object,
*,
default: int,
minimum: int,
maximum: int,
) -> int:
if isinstance(value, bool):
return default
candidate: int | float | str
if isinstance(value, (int, float, str)):
candidate = value
else:
return default
if isinstance(candidate, str):
candidate = candidate.strip()
try:
parsed = int(candidate)
except (TypeError, ValueError):
return default
if parsed < minimum:
return minimum
if parsed > maximum:
return maximum
return parsed
def _parse_event_id(value: object) -> UUID:
if not isinstance(value, str) or not value.strip():
raise ValueError("eventId is required")
try:
return UUID(value)
except ValueError as exc:
raise ValueError("eventId must be a valid UUID") from exc
def _service(session: AsyncSession, owner_id: UUID) -> ScheduleItemService:
return 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
color = tool_args.get("color")
raw_color = color.strip() if isinstance(color, str) and color.strip() else "#4F46E5"
color_value = raw_color if _HEX_COLOR_PATTERN.match(raw_color) else "#4F46E5"
reminder_raw = tool_args.get("reminderMinutes")
reminder_value: int | None = None
if isinstance(reminder_raw, bool):
reminder_value = None
elif isinstance(reminder_raw, (int, float, str)):
try:
parsed = int(str(reminder_raw).strip())
if parsed < 0 or parsed > 10080:
raise ValueError("reminderMinutes must be 0..10080")
reminder_value = parsed
except ValueError as exc:
raise ValueError("reminderMinutes must be an integer in 0..10080") from exc
return ScheduleItemMetadata(
location=location_value,
color=color_value,
reminder_minutes=reminder_value,
)
def _event_payload(event: object) -> dict[str, object]:
event_id = str(getattr(event, "id"))
metadata = getattr(event, "metadata", None)
location_value = getattr(metadata, "location", None)
color_value = getattr(metadata, "color", None) or "#4F46E5"
reminder_minutes_value = getattr(metadata, "reminder_minutes", None)
return {
"id": event_id,
"title": getattr(event, "title"),
"description": getattr(event, "description"),
"startAt": getattr(event, "start_at").isoformat(),
"endAt": getattr(event, "end_at").isoformat()
if getattr(event, "end_at") is not None
else None,
"timezone": getattr(event, "timezone"),
"location": location_value,
"color": color_value,
"reminderMinutes": reminder_minutes_value,
}
async def _execute_list_calendar_events(
session: AsyncSession,
owner_id: UUID,
tool_args: dict[str, object],
) -> dict[str, object]:
page = _parse_positive_int(
tool_args.get("page"),
default=1,
minimum=1,
maximum=100000,
)
page_size = _parse_positive_int(
tool_args.get("pageSize"),
default=20,
minimum=1,
maximum=100,
)
service = _service(session, owner_id)
items, total = await service.list_paginated(page=page, page_size=page_size)
total_pages = max(1, (total + page_size - 1) // page_size) if total else 0
return {
"type": "calendar_event_list.v1",
"version": "v1",
"data": {
"items": [_event_payload(item) for item in items],
"pagination": {
"page": page,
"pageSize": page_size,
"total": total,
"totalPages": total_pages,
},
"ok": True,
"message": "已获取日程列表",
},
"actions": [],
}
async def _execute_create(
*,
service: ScheduleItemService,
tool_args: dict[str, object],
) -> dict[str, object]:
title = str(tool_args.get("title", "新的日程")).strip() or "新的日程"
description = str(tool_args.get("description", "")).strip() or None
start_at = _parse_datetime(tool_args.get("startAt"))
if start_at is None:
start_at = datetime.now(timezone.utc) + timedelta(hours=1)
end_at = _parse_datetime(tool_args.get("endAt"))
timezone_value = (
str(tool_args.get("timezone", "Asia/Shanghai")).strip() or "Asia/Shanghai"
)
created = await service.create_agent_generated(
ScheduleItemCreateRequest(
title=title,
description=description,
start_at=start_at,
end_at=end_at,
timezone=timezone_value,
metadata=_resolve_metadata(tool_args),
)
)
event_data = _event_payload(created)
event_id = str(event_data["id"])
return {
"type": "calendar_card.v1",
"version": "v1",
"data": {
**event_data,
"sourceType": "agent_generated",
"ok": True,
"message": "日程已创建",
},
"actions": [
{
"type": "link",
"label": "查看详情",
"target": f"/schedule-items/{event_id}",
}
],
}
async def _execute_update(
*,
service: ScheduleItemService,
tool_args: dict[str, object],
) -> dict[str, object]:
event_id = _parse_event_id(tool_args.get("eventId"))
update_data: dict[str, object] = {}
for source_key, target_key in (
("title", "title"),
("description", "description"),
("timezone", "timezone"),
):
value = tool_args.get(source_key)
if isinstance(value, str):
update_data[target_key] = value.strip()
start_at = _parse_datetime(tool_args.get("startAt"))
if start_at is not None:
update_data["start_at"] = start_at
end_at = _parse_datetime(tool_args.get("endAt"))
if end_at is not None:
update_data["end_at"] = end_at
status_value = tool_args.get("status")
if isinstance(status_value, str) and status_value.strip():
try:
update_data["status"] = ScheduleItemStatus(status_value.strip().lower())
except ValueError as exc:
raise ValueError(
"status must be one of: active, completed, canceled, archived"
) from exc
has_location = isinstance(tool_args.get("location"), str)
has_color = isinstance(tool_args.get("color"), str)
has_reminder = "reminderMinutes" in tool_args
if has_location or has_color or has_reminder:
existing = await service.get_by_id(event_id)
metadata_dump = (
existing.metadata.model_dump() if existing.metadata is not None else {}
)
if has_location:
metadata_dump["location"] = str(tool_args.get("location")).strip() or None
if has_color:
color = str(tool_args.get("color")).strip()
if not color:
metadata_dump["color"] = None
elif _HEX_COLOR_PATTERN.match(color):
metadata_dump["color"] = color
else:
raise ValueError("color must be a hex string like #RRGGBB")
if has_reminder:
reminder_raw = tool_args.get("reminderMinutes")
if reminder_raw is None:
metadata_dump["reminder_minutes"] = None
elif isinstance(reminder_raw, bool):
raise ValueError("reminderMinutes must be an integer in 0..10080")
else:
try:
reminder = int(str(reminder_raw).strip())
except ValueError as exc:
raise ValueError(
"reminderMinutes must be an integer in 0..10080"
) from exc
if reminder < 0 or reminder > 10080:
raise ValueError("reminderMinutes must be 0..10080")
metadata_dump["reminder_minutes"] = reminder
update_data["metadata"] = ScheduleItemMetadata.model_validate(metadata_dump)
updated = await service.update(
event_id,
ScheduleItemUpdateRequest.model_validate(update_data),
)
event_data = _event_payload(updated)
return {
"type": "calendar_card.v1",
"version": "v1",
"data": {
**event_data,
"sourceType": "agent_generated",
"ok": True,
"message": "日程已更新",
},
"actions": [
{
"type": "link",
"label": "查看详情",
"target": f"/schedule-items/{event_data['id']}",
}
],
}
async def _execute_delete(
*,
service: ScheduleItemService,
tool_args: dict[str, object],
) -> dict[str, object]:
event_id = _parse_event_id(tool_args.get("eventId"))
await service.delete(event_id)
return {
"type": "calendar_operation.v1",
"version": "v1",
"data": {
"operation": "delete",
"id": str(event_id),
"ok": True,
"message": "日程已删除",
},
"actions": [],
}
async def _execute_mutate_calendar_event(
session: AsyncSession,
owner_id: UUID,
tool_args: dict[str, object],
) -> dict[str, object]:
operation_raw = tool_args.get("operation")
if not isinstance(operation_raw, str) or not operation_raw.strip():
raise ValueError("operation is required")
operation = operation_raw.strip().lower()
service = _service(session, owner_id)
if operation == "create":
return await _execute_create(service=service, tool_args=tool_args)
if operation == "update":
return await _execute_update(service=service, tool_args=tool_args)
if operation == "delete":
return await _execute_delete(service=service, tool_args=tool_args)
raise ValueError("operation must be one of: create, update, delete")
async def _execute_share_calendar_event(
*,
session: AsyncSession,
owner_id: UUID,
tool_args: dict[str, object],
) -> dict[str, object]:
event_id = _parse_event_id(tool_args.get("eventId"))
invite_result = await _share_event_with_invitees(
session=session,
owner_id=owner_id,
event_id=event_id,
tool_args=tool_args,
)
if invite_result is None:
raise ValueError(
"at least one invite target is required: inviteUserEmails, inviteUserNames, or inviteUserIds"
)
return {
"type": "calendar_operation.v1",
"version": "v1",
"data": {
"operation": "share",
"id": str(event_id),
"ok": True,
"message": "日程已分享",
"shareResult": invite_result,
},
"actions": [],
}
@@ -0,0 +1,237 @@
from __future__ import annotations
from typing import Annotated, Any, cast
from uuid import UUID
from fastapi import HTTPException
from pydantic import Field
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from core.auth.jwt_verifier import JwtVerifier, TokenValidationError
from core.agentscope.tools.tool_response_builder import (
build_success_response,
build_error_response,
)
from core.agentscope.schemas.runtime_models import ToolOutputContent
from core.config.settings import config
from models.profile import Profile
from services.base.supabase import supabase_service
from v1.auth.gateway import SupabaseAuthGateway
def _verify_user_token(*, user_token: str, owner_id: UUID) -> bool:
"""Verify the user token matches the owner_id."""
jwt_secret = config.supabase.jwt_secret
if jwt_secret is None:
return False
verifier = JwtVerifier(
issuer=str(config.supabase.jwt_issuer),
jwt_secret=jwt_secret.get_secret_value(),
jwt_algorithm=config.supabase.jwt_algorithm,
)
try:
payload = verifier.verify(user_token)
except TokenValidationError:
return False
subject = payload.get("sub")
return isinstance(subject, str) and subject == str(owner_id)
def _list_auth_users() -> list[Any]:
"""List all auth users from Supabase."""
admin_client = supabase_service.get_admin_client()
users: list[Any] = []
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
def _find_auth_email_by_user_id(*, users: list[Any], user_id: UUID) -> str | None:
"""Find user email by user ID from auth users list."""
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, Any]:
"""Resolve user identity by email or username."""
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 HTTPException(
status_code=400,
detail="请提供 email 或 username 其中之一",
)
if email:
auth_gateway = SupabaseAuthGateway()
user = await auth_gateway.get_user_by_email(email)
user_id = UUID(user.id)
stmt = (
select(Profile.username)
.where(Profile.id == user_id)
.where(Profile.deleted_at.is_(None))
)
username = (await session.execute(stmt)).scalar_one_or_none()
return {
"userId": str(user_id),
"email": user.email,
"username": username,
"matchedBy": "email",
}
stmt = (
select(Profile)
.where(Profile.username == name)
.where(Profile.deleted_at.is_(None))
)
profile = await session.execute(stmt)
profile = profile.scalar_one_or_none()
if profile is None:
raise HTTPException(status_code=404, detail="用户不存在")
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",
}
async def user_lookup(
user_email: Annotated[
str | None,
Field(description="User email address to look up."),
] = None,
user_name: Annotated[
str | None,
Field(description="Username to look up."),
] = None,
session: Any = None,
owner_id: Any = None,
user_token: str | None = None,
) -> ToolOutputContent:
"""
Look up user information by email or username.
Args:
user_email: User email address to look up.
user_name: Username to look up.
session: Database session (runtime preset).
owner_id: Current user ID (runtime preset).
user_token: Validated JWT token (runtime preset).
Returns:
ToolOutputContent with user information or error.
"""
if session is None or owner_id is None:
return build_error_response(
code="MISSING_RUNTIME_ARGS",
message="用户查找工具缺少运行时参数",
retryable=False,
)
if not isinstance(user_token, str) or not user_token.strip():
return build_error_response(
code="UNAUTHORIZED",
message="用户查找工具需要有效的用户令牌",
retryable=False,
)
if not _verify_user_token(user_token=user_token, owner_id=cast(UUID, owner_id)):
return build_error_response(
code="UNAUTHORIZED",
message="用户查找工具需要有效的用户令牌",
retryable=False,
)
try:
resolved = await _resolve_identity(
session=cast(AsyncSession, session),
user_email=user_email,
user_name=user_name,
)
user_id = resolved.get("userId", "")
email = resolved.get("email", "")
username = resolved.get("username", "")
matched_by = resolved.get("matchedBy", "")
return build_success_response(
title="用户信息",
summary=f"已找到用户: {username or email}",
payload={
"ok": True,
"userId": user_id,
"email": email,
"username": username,
"matchedBy": matched_by,
},
items=[],
kv_pairs=[
{
"key": "user_id",
"label": "用户ID",
"value": user_id,
"copyable": True,
},
{"key": "email", "label": "邮箱", "value": email, "copyable": True},
{
"key": "username",
"label": "用户名",
"value": username or "-",
"copyable": True,
},
{
"key": "matched_by",
"label": "匹配方式",
"value": matched_by,
"copyable": False,
},
],
)
except HTTPException as exc:
if exc.status_code == 404:
return build_error_response(
code="NOT_FOUND",
message=exc.detail or "用户不存在",
retryable=False,
)
return build_error_response(
code="LOOKUP_FAILED",
message=exc.detail or "用户查找失败",
retryable=True,
)
except Exception as exc:
return build_error_response(
code="INTERNAL_ERROR",
message=f"用户查找失败: {str(exc)}",
retryable=True,
)
@@ -2,7 +2,10 @@ from __future__ import annotations
from typing import Any, AsyncGenerator, Callable
from core.agentscope.tools.tool_response_builder import build_tool_response
from core.agentscope.tools.tool_response_builder import (
build_tool_response,
build_error_response,
)
from core.agentscope.tools.tool_meta import ToolMeta
@@ -58,31 +61,27 @@ def create_hitl_middleware(
return
if decision == "rejected":
yield build_tool_response(
{
"type": "tool_approval.v1",
"version": "v1",
"data": {
"status": "rejected",
"tool": tool_name,
"ok": False,
"message": "tool call rejected by reviewer",
},
}
content = build_error_response(
code="TOOL_REJECTED",
message=f"工具 {tool_name} 的调用已被审核拒绝",
retryable=False,
details={
"tool": tool_name,
"status": "rejected",
},
)
yield build_tool_response(content)
return
yield build_tool_response(
{
"type": "tool_approval.v1",
"version": "v1",
"data": {
"status": "pending",
"tool": tool_name,
"ok": False,
"message": "tool call requires approval",
},
}
content = build_error_response(
code="TOOL_PENDING_APPROVAL",
message=f"工具 {tool_name} 需要审核批准",
retryable=True,
details={
"tool": tool_name,
"status": "pending",
},
)
yield build_tool_response(content)
return hitl_middleware
@@ -1,18 +1,110 @@
from __future__ import annotations
import json
from typing import Any
from typing import TYPE_CHECKING
from agentscope.message import TextBlock
from agentscope.tool import ToolResponse
if TYPE_CHECKING:
from core.agentscope.schemas.runtime_models import ToolOutputContent
def build_tool_response(payload: dict[str, Any]):
from agentscope.message import TextBlock
from agentscope.tool import ToolResponse
def build_tool_response(
content: "ToolOutputContent",
*,
tool_name: str = "unknown",
) -> ToolResponse:
"""
Build a ToolResponse from ToolOutputContent.
Args:
content: The ToolOutputContent instance to serialize.
tool_name: Name of the tool (for debugging).
Returns:
ToolResponse with serialized content.
"""
payload = content.model_dump(mode="json", exclude_none=True)
return ToolResponse(
content=[
TextBlock(
type="text",
text=json.dumps(payload, ensure_ascii=True, separators=(",", ":")),
text=json.dumps(payload, ensure_ascii=False, separators=(",", ":")),
)
]
)
def build_success_response(
title: str | None = None,
summary: str | None = None,
payload: dict | None = None,
items: list[dict] | None = None,
kv_pairs: list[dict] | None = None,
**kwargs,
) -> "ToolOutputContent":
"""
Build a success ToolOutputContent.
Args:
title: Optional title for the response.
summary: Optional summary/description.
payload: Optional structured payload data.
items: Optional list of items (for list UI).
kv_pairs: Optional key-value pairs (for kv UI).
**kwargs: Additional fields for ToolOutputContent.
Returns:
ToolOutputContent with success status.
"""
from core.agentscope.schemas.runtime_models import ToolOutputContent
return ToolOutputContent(
title=title,
summary=summary,
payload=payload or {},
items=items or [],
kv_pairs=kv_pairs or [],
**kwargs,
)
def build_error_response(
code: str,
message: str,
retryable: bool = False,
details: dict | None = None,
**kwargs,
) -> "ToolOutputContent":
"""
Build an error ToolOutputContent.
Args:
code: Error code (e.g., NOT_FOUND, UNAUTHORIZED).
message: Human-readable error message.
retryable: Whether the operation can be retried.
details: Additional error details.
**kwargs: Additional fields for ToolOutputContent.
Returns:
ToolOutputContent with error information.
"""
from core.agentscope.schemas.runtime_models import ToolOutputContent
return ToolOutputContent(
title="操作失败",
summary=message,
payload={
"code": code,
"message": message,
"retryable": retryable,
"details": details or {},
},
items=[],
kv_pairs=[
{"key": "error_code", "label": "错误代码", "value": code},
{"key": "message", "label": "错误信息", "value": message},
],
**kwargs,
)
+20 -1
View File
@@ -11,6 +11,9 @@ from core.agentscope.tools.custom.calendar import (
calendar_read,
calendar_write,
)
from core.agentscope.tools.custom.user_lookup import (
user_lookup,
)
from core.agentscope.tools.hitl_middleware import register_tool_middlewares
from core.agentscope.tools.tool_meta import TOOL_META
@@ -32,7 +35,14 @@ TOOL_GROUPS: dict[str, ToolGroup] = {
"intent": ToolGroup(stage="intent", tool_names=frozenset({"calendar_read"})),
"execution": ToolGroup(
stage="execution",
tool_names=frozenset({"calendar_read", "calendar_write", "calendar_share"}),
tool_names=frozenset(
{
"calendar_read",
"calendar_write",
"calendar_share",
"user_lookup",
}
),
),
"report": ToolGroup(stage="report", tool_names=frozenset()),
}
@@ -79,6 +89,15 @@ def _load_custom_tool_bindings(
"user_token": user_token or "",
},
),
CustomToolBinding(
name="user_lookup",
func=user_lookup,
preset_kwargs={
"session": session,
"owner_id": owner_id,
"user_token": user_token or "",
},
),
]