Files
workphone-sdk/sdk/app/routers/fleet.py
2026-08-10 05:42:42 +08:00

410 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
多手机设备管理接口。
面向存客宝/触客宝/超管等外部系统:把单机 SDK 能力包装成可筛选、
可批量执行、可审计的 Fleet API。单机能力仍由 devices/unified 路由负责。
"""
from __future__ import annotations
import asyncio
import uuid
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
from services.device_id_util import device_id_md5
from services.ws_hub import ws_hub
router = APIRouter()
WRITE_ACTIONS = {
"send_message",
"batch_send",
"mass_send",
"add_friend",
"batch_add_friend",
"post_moments",
"like_moments",
"comment_moments",
"create_group",
"invite_to_group",
"remove_from_group",
"set_group_notice",
"set_group_name",
"set_group_welcome",
"delete_tag",
"create_tag",
"tag_add",
"tag_remove",
}
def _result_success(result: Dict[str, Any]) -> bool:
"""只接受执行端明确 success=true禁止以 HTTP/code 单独判定成功。"""
if not isinstance(result, dict):
return False
payload = result.get("data") if isinstance(result.get("data"), dict) else result
return payload.get("success") is True
def _normalize_agent_result(action: str, result: Dict[str, Any]) -> Dict[str, Any]:
"""将 Agent 原始动作回执归一化为 Fleet 的标准可验收结构。"""
normalized = dict(result or {})
data = dict(normalized.get("data") or {})
raw = normalized.get("raw_rpc_receipt") or dict(normalized)
trace_id = (normalized.get("trace_id") or normalized.get("command_id")
or raw.get("trace_id") or raw.get("command_id"))
# Android Agent 的截图字段是 image_base64Fleet 标准字段为 base64。
# 有有效图片即是该只读采集动作的业务成功,不依赖 HTTP 200 单独判定。
if action == "screenshot" and int(normalized.get("code", 500)) == 200 and data.get("image_base64"):
data.setdefault("base64", data["image_base64"])
data["success"] = True
normalized["readback"] = {
"kind": "screenshot",
"image_bytes_base64": len(str(data["image_base64"])),
"channel": data.get("channel") or "agent_internal_screenshot",
}
normalized["data"] = data
normalized["trace_id"] = trace_id
normalized["raw_rpc_receipt"] = raw
return normalized
def _timeout_result(device_id: str, timeout: int) -> Dict[str, Any]:
"""统一 Fleet 超时回执,保留可重试元数据。"""
trace_id = uuid.uuid4().hex
receipt = {
"operation": "fleet_execute",
"device_id": device_id,
"channel": "websocket/timeout",
"trace_id": trace_id,
}
return {
"device_id": device_id,
"trace_id": trace_id,
"channel_used": "websocket/timeout",
"raw_rpc_receipt": receipt,
"readback": None,
"success": False,
"result": {
"code": 504,
"success": False,
"error_code": "timeout",
"error_message": f"设备执行超过 {timeout}s 超时",
"retryable": True,
"retry_after_seconds": 1,
"channel_used": "websocket/timeout",
"trace_id": trace_id,
"raw_rpc_receipt": receipt,
"readback": None,
},
}
class FleetOperationReceipt(BaseModel):
"""Fleet 统一离线/失败回执。"""
code: int
success: bool
data: dict = Field(default_factory=dict)
error_code: Optional[str] = None
error_message: Optional[str] = None
retryable: bool = False
trace_id: Optional[str] = None
channel_used: str
raw_rpc_receipt: Optional[dict] = None
readback: Optional[dict] = None
def _fleet_offline_receipt(req: "FleetExecuteRequest") -> JSONResponse:
"""Fleet 无 WSS 目标时返回结构化 503不转 ADB 或本机通道。"""
trace_id = req.trace_id or uuid.uuid4().hex
receipt = {
"operation": "fleet_execute",
"device_ids": req.device_ids or [],
"project_id": req.project_id,
"channel": "websocket/offline",
"trace_id": trace_id,
}
return JSONResponse(
status_code=503,
content={
"code": 503,
"success": False,
"data": {"action": req.action, "target_count": 0, "results": []},
"error_code": "device_offline",
"error_message": "没有匹配的 WSS 在线设备",
"retryable": True,
"trace_id": trace_id,
"channel_used": "websocket/offline",
"raw_rpc_receipt": receipt,
"readback": None,
},
)
class FleetExecuteRequest(BaseModel):
"""批量执行请求。"""
device_ids: Optional[List[str]] = None
project_id: Optional[str] = None
all_online: bool = False
platform: str = "wechat"
action: str
params: Dict[str, Any] = Field(default_factory=dict)
hook_only: bool = False
timeout: int = 60
max_concurrency: int = 3
dry_run: bool = False
confirm: bool = False
trace_id: Optional[str] = None
def _normalize_status(device: dict) -> str:
"""在线态只信任当前进程中的 WSS 连接,历史登记不能冒充在线。"""
return "online" if ws_hub.is_online(device.get("device_id", "")) else "offline"
async def _select_devices(
*,
device_ids: Optional[List[str]] = None,
project_id: Optional[str] = None,
all_online: bool = False,
status: str = "",
capability: str = "",
) -> List[dict]:
from routers.devices import list_ws_managed_devices
devices = await list_ws_managed_devices()
selected = []
wanted = set(device_ids or [])
for device in devices:
did = device.get("device_id", "")
if wanted and did not in wanted:
continue
if project_id and str(device.get("project_id") or "") != str(project_id):
continue
if all_online and not ws_hub.is_online(did):
continue
if status and _normalize_status(device) != status:
continue
if capability and capability not in (device.get("capabilities") or []):
continue
normalized = dict(device)
normalized["status"] = _normalize_status(normalized)
normalized["online"] = ws_hub.is_online(did)
normalized["device_id_md5"] = normalized.get("device_id_md5") or device_id_md5(did)
selected.append(normalized)
return selected
@router.get("/fleet/summary", response_model=dict, responses={503: {"model": FleetOperationReceipt}})
async def fleet_summary(project_id: str = ""):
"""所有手机/项目手机总览。"""
devices = await _select_devices(project_id=project_id or None)
online = [d for d in devices if d.get("online")]
offline = [d for d in devices if not d.get("online")]
by_project: Dict[str, int] = {}
for device in devices:
pid = str(device.get("project_id") or "")
by_project[pid] = by_project.get(pid, 0) + 1
return {
"code": 200,
"data": {
"total": len(devices),
"online": len(online),
"adb": 0,
"offline": len(offline),
"by_project": by_project,
"device_ids": [d.get("device_id") for d in devices],
"online_device_ids": [d.get("device_id") for d in online],
},
}
@router.get("/fleet/devices", response_model=dict, responses={503: {"model": FleetOperationReceipt}})
async def fleet_devices(
project_id: str = "",
status: str = "",
capability: str = "",
online_only: bool = False,
):
"""按项目、状态、能力筛选设备列表。"""
devices = await _select_devices(
project_id=project_id or None,
status=status,
capability=capability,
all_online=online_only,
)
return {"code": 200, "data": {"devices": devices, "count": len(devices)}}
@router.post("/fleet/execute", response_model=dict, responses={503: {"model": FleetOperationReceipt}})
async def fleet_execute(req: FleetExecuteRequest):
"""在多台在线手机上批量执行同一个设备动作。"""
batch_trace_id = req.trace_id or uuid.uuid4().hex
devices = await _select_devices(
device_ids=req.device_ids,
project_id=req.project_id,
all_online=req.all_online or not req.device_ids,
status="online",
)
if not devices:
return _fleet_offline_receipt(req)
is_write = req.action in WRITE_ACTIONS
if is_write and not (req.dry_run or req.confirm):
return {
"code": 200,
"success": False,
"trace_id": batch_trace_id,
"channel_used": "fleet/confirm_required",
"raw_rpc_receipt": {
"operation": "fleet_execute",
"action": req.action,
"confirm_required": True,
"trace_id": batch_trace_id,
},
"readback": None,
"data": {
"success": False,
"confirm_required": True,
"reason": "批量写类动作需要 confirm=true可先 dry_run=true 查看目标设备",
"action": req.action,
"target_count": len(devices),
"targets": [d.get("device_id") for d in devices],
},
}
if req.dry_run:
return {
"code": 200,
"success": True,
"trace_id": batch_trace_id,
"channel_used": "fleet/dry_run",
"raw_rpc_receipt": {
"operation": "fleet_execute",
"action": req.action,
"dry_run": True,
"executed": False,
"trace_id": batch_trace_id,
},
"readback": None,
"data": {
"success": True,
"dry_run": True,
"executed": False,
"action": req.action,
"target_count": len(devices),
"targets": [d.get("device_id") for d in devices],
},
}
sem = asyncio.Semaphore(max(1, min(int(req.max_concurrency or 1), 10)))
timeout = max(5, min(int(req.timeout or 60), 300))
async def run_one(device: dict) -> dict:
did = device.get("device_id", "")
async with sem:
try:
# Fleet 与单设备微信路由共用 unified + device_transport
# 确保 WS/Hook 选路、hook_only、错误码和真实回执口径一致。
from routers.unified import _execute_skill
result = await _execute_skill(
did,
req.platform,
req.action,
req.params or {},
timeout=timeout,
# 微信默认固定走 WSS Agent → 手机本机 Frida RPC调用方
# 仅在非微信平台时可维持原有 generic Agent 语义。
hook_only=req.hook_only or req.platform == "wechat",
)
result = _normalize_agent_result(req.action, result)
return {
"device_id": did,
"device_id_md5": device_id_md5(did),
"success": _result_success(result),
"trace_id": result.get("trace_id"),
"channel_used": result.get("channel_used") or result.get("_channel_used") or "websocket/agent",
"raw_rpc_receipt": result.get("raw_rpc_receipt"),
"readback": result.get("readback") or result.get("db_readback"),
"result": result,
}
except asyncio.TimeoutError:
return _timeout_result(did, timeout)
except Exception as exc:
trace_id = uuid.uuid4().hex
receipt = {
"operation": "fleet_execute",
"device_id": did,
"channel": "websocket/error",
"exception": type(exc).__name__,
"trace_id": trace_id,
}
return {
"device_id": did,
"device_id_md5": device_id_md5(did),
"success": False,
"trace_id": trace_id,
"channel_used": "websocket/error",
"raw_rpc_receipt": receipt,
"readback": None,
"result": {
"code": 503,
"success": False,
"error_code": "execution_failed",
"error_message": str(exc),
"retryable": True,
"retry_after_seconds": 1,
"channel_used": "websocket/error",
"trace_id": trace_id,
"raw_rpc_receipt": receipt,
"readback": None,
},
}
results = await asyncio.gather(*(run_one(device) for device in devices))
ok = sum(1 for item in results if item.get("success"))
body = {
"code": 200,
"success": ok == len(results),
"trace_id": batch_trace_id,
"channel_used": "frida_rpc" if req.platform == "wechat" else "websocket/agent",
"raw_rpc_receipt": {
"operation": "fleet_execute",
"action": req.action,
"trace_id": batch_trace_id,
"results": [item.get("raw_rpc_receipt") for item in results],
},
"readback": [item.get("readback") for item in results],
"data": {
"success": ok == len(results),
"action": req.action,
"target_count": len(results),
"success_count": ok,
"failed_count": len(results) - ok,
"results": results,
},
}
if ok == 0 and results and all(int(item.get("result", {}).get("code", 200)) >= 500 for item in results):
trace_id = batch_trace_id
body.update({
"code": 503,
"error_code": "fleet_execution_failed",
"error_message": "所有 WSS Agent 执行目标失败",
"retryable": True,
"trace_id": trace_id,
"channel_used": "websocket/agent",
"raw_rpc_receipt": {"operation": "fleet_execute", "trace_id": trace_id, "results": results},
"readback": None,
})
return JSONResponse(status_code=503, content=body)
return body