refactor: 重整 schemas 作用域并统一用户上下文模型
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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 "",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user