Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 31 additions & 6 deletions src/goofish_cli/commands/message/send.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,28 @@
"""message send — 向指定会话发送一条消息(文本/图片)。

写操作,走限流 + 熔断。未知 cid 时可传 --item-id 自动创建单聊。

发送前必须等 `/s/vulcan`(见 `wait_ready`),并校验服务端 ack;
未拿到 `code=200` 一律报错,不返回假成功。
"""

import asyncio
from typing import Any, Literal

from goofish_cli.core import Session, Strategy, command
from goofish_cli.core.errors import GoofishError
from goofish_cli.core.guard import watch
from goofish_cli.core.limiter import acquire
from goofish_cli.core.token import get_access_token
from goofish_cli.core.ws import (
connect,
create_chat,
heartbeat_loop,
recv_ack,
register,
send_image,
send_text,
wait_ready,
)


Expand All @@ -25,7 +31,7 @@
name="send",
description="向会话发送消息(text/image)。text 必填,image 走 url+wh",
strategy=Strategy.COOKIE,
columns=["cid", "toid", "kind", "ok", "mid"],
columns=["cid", "toid", "kind", "ok", "mid", "message_id"],
write=True,
)
def send(
Expand Down Expand Up @@ -71,9 +77,15 @@ async def _send(
await register(ws, session, token)
hb = asyncio.create_task(heartbeat_loop(ws))
try:
# `/r/` 请求必须等 `/s/vulcan` 之后才发,否则服务端回 code 400。
if not await wait_ready(ws):
raise GoofishError("IM 连接未就绪(等待 /s/vulcan 超时),未发送")

if item_id:
await create_chat(ws, myid=session.unb, toid=toid, item_id=item_id)
await asyncio.sleep(0.5)
create_mid = await create_chat(
ws, myid=session.unb, toid=toid, item_id=item_id
)
await recv_ack(ws, create_mid)

if kind == "text":
if not text:
Expand All @@ -95,8 +107,21 @@ async def _send(
)
else:
raise ValueError(f"不支持的 kind: {kind}")
# 等一轮 ack 回包,避免 WS 提前关
await asyncio.sleep(1.0)
ack = await recv_ack(ws, mid)
finally:
hb.cancel()
return {"cid": cid, "toid": toid, "kind": kind, "ok": True, "mid": mid}

ack_code = (ack or {}).get("code")
if ack_code != 200:
detail = "无回包(超时)" if ack is None else f"ack_code={ack_code}"
raise GoofishError(f"发送未被服务端确认:{detail}")

body = (ack or {}).get("body") or {}
return {
"cid": cid,
"toid": toid,
"kind": kind,
"ok": True,
"mid": mid,
"message_id": str(body.get("messageId", "")),
}
55 changes: 55 additions & 0 deletions src/goofish_cli/core/ws.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,61 @@ def build_ack(msg: dict[str, Any]) -> dict[str, Any]:
return ack


async def _recv_json(ws: ClientConnection, *, timeout: float) -> dict[str, Any] | None:
"""收一帧并解析成 dict。超时或非 JSON 返回 None。"""
try:
raw = await asyncio.wait_for(ws.recv(), timeout=timeout)
except TimeoutError:
return None
try:
parsed = json.loads(raw)
except (json.JSONDecodeError, TypeError):
return None
return parsed if isinstance(parsed, dict) else None


async def wait_ready(ws: ClientConnection, *, timeout: float = 15.0) -> bool:
"""等服务端推 `/s/vulcan`,期间对下行帧回 ack。就绪返回 True。

`/r/` 请求在 `/s/vulcan` 到达之前发出会被服务端以 `code 400` 拒绝。
`/reg` 自身返回 200,所以「注册成功」并不代表可以发请求。
`list_user_messages()` 一直是等到 `/s/vulcan` 才发 `/r/` 请求;
`message send` 缺这一步,导致发送恒被拒。
"""
loop = asyncio.get_running_loop()
deadline = loop.time() + timeout
while True:
remaining = deadline - loop.time()
if remaining <= 0:
return False
frame = await _recv_json(ws, timeout=min(3.0, remaining))
if frame is None:
continue
with suppress(Exception):
await ws.send(json.dumps(build_ack(frame)))
if frame.get("lwp") == "/s/vulcan":
return True


async def recv_ack(
ws: ClientConnection, mid: str, *, timeout: float = 10.0
) -> dict[str, Any] | None:
"""读取 `mid` 对应的响应帧,期间继续对下行推送回 ack。超时返回 None。"""
loop = asyncio.get_running_loop()
deadline = loop.time() + timeout
while True:
remaining = deadline - loop.time()
if remaining <= 0:
return None
frame = await _recv_json(ws, timeout=min(3.0, remaining))
if frame is None:
continue
if (frame.get("headers") or {}).get("mid") == mid:
return frame
with suppress(Exception):
await ws.send(json.dumps(build_ack(frame)))


async def send_text(
ws: ClientConnection, *, myid: str, cid: str, toid: str, text: str
) -> str:
Expand Down
88 changes: 88 additions & 0 deletions tests/test_ws_ready.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
"""wait_ready / recv_ack 的时序逻辑(不连 WebSocket)。

背景:`/r/` 请求在服务端推 `/s/vulcan` 之前发出会被以 `code 400` 拒绝,
而 `/reg` 自身返回 200,容易误判连接已就绪。
"""

import asyncio
import json

import pytest

from goofish_cli.core.ws import recv_ack, wait_ready


class FakeWS:
"""按脚本吐帧的假 ws;记录本端发出的内容。"""

def __init__(self, frames: list[dict]):
self._frames = list(frames)
self.sent: list[dict] = []

async def recv(self) -> str:
if not self._frames:
await asyncio.sleep(3600) # 模拟没有更多下行
return json.dumps(self._frames.pop(0))

async def send(self, raw: str) -> None:
self.sent.append(json.loads(raw))


def test_wait_ready_returns_on_vulcan():
ws = FakeWS([
{"headers": {"mid": "reg-1"}, "code": 200},
{"headers": {"mid": "p-1", "sid": "s"}, "lwp": "/s/sync"},
{"headers": {"mid": "p-2", "sid": "s"}, "lwp": "/s/vulcan"},
])
assert asyncio.run(wait_ready(ws, timeout=5.0)) is True
# 每个下行帧都回了 ack
assert len(ws.sent) == 3
assert all(a["code"] == 200 for a in ws.sent)


def test_wait_ready_times_out_without_vulcan():
ws = FakeWS([{"headers": {"mid": "p-1"}, "lwp": "/s/sync"}])
assert asyncio.run(wait_ready(ws, timeout=0.3)) is False


def test_recv_ack_matches_by_mid():
ws = FakeWS([
{"headers": {"mid": "other"}, "lwp": "/s/sync"},
{"headers": {"mid": "mine"}, "code": 200, "body": {"messageId": "srv-1"}},
])
ack = asyncio.run(recv_ack(ws, "mine", timeout=5.0))
assert ack is not None
assert ack["code"] == 200
assert ack["body"]["messageId"] == "srv-1"
# 不匹配的下行帧被 ack,匹配的那帧不回 ack
assert [a["headers"]["mid"] for a in ws.sent] == ["other"]


def test_recv_ack_surfaces_rejection():
"""服务端拒绝时必须把 400 透出来,不能当成成功。"""
ws = FakeWS([{"headers": {"mid": "mine"}, "code": 400}])
ack = asyncio.run(recv_ack(ws, "mine", timeout=5.0))
assert ack is not None
assert ack["code"] == 400


def test_recv_ack_times_out():
ws = FakeWS([])
assert asyncio.run(recv_ack(ws, "mine", timeout=0.3)) is None


@pytest.mark.parametrize("junk", ["not json", "[1,2,3]"])
def test_recv_json_skips_non_dict_frames(junk):
"""非 JSON / 非 dict 帧不能让循环崩掉。"""

class JunkWS(FakeWS):
async def recv(self) -> str:
if self._frames:
return json.dumps(self._frames.pop(0))
if not getattr(self, "_junked", False):
self._junked = True
return junk
await asyncio.sleep(3600)

ws = JunkWS([{"headers": {"mid": "p"}, "lwp": "/s/vulcan"}])
assert asyncio.run(wait_ready(ws, timeout=5.0)) is True