Compare commits

..
8 changed files with 729 additions and 28 deletions
+21 -2
View File
@@ -1,16 +1,30 @@
import os import os
from sqlalchemy import create_engine from sqlalchemy import create_engine, event
from sqlalchemy.orm import sessionmaker, declarative_base, Session from sqlalchemy.orm import sessionmaker, declarative_base, Session
BASE_DIR = os.path.dirname(os.path.abspath(__file__)) BASE_DIR = os.path.dirname(os.path.abspath(__file__))
DB_FILE = os.path.join(BASE_DIR, "avatar.db") DB_FILE = os.path.join(BASE_DIR, "avatar.db")
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}") DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
engine = create_engine( engine = create_engine(
DATABASE_URL, DATABASE_URL,
connect_args={"check_same_thread": False} if DATABASE_URL.startswith("sqlite:") else {}, connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {},
) )
if IS_SQLITE:
@event.listens_for(engine, "connect")
def _configure_sqlite_connection(dbapi_connection, _connection_record):
cursor = dbapi_connection.cursor()
try:
cursor.execute("PRAGMA synchronous=NORMAL")
cursor.execute("PRAGMA busy_timeout=30000")
finally:
cursor.close()
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
Base = declarative_base() Base = declarative_base()
@@ -26,6 +40,10 @@ def get_db():
def init_db(): def init_db():
import models import models
if IS_SQLITE:
with engine.connect() as conn:
conn.exec_driver_sql("PRAGMA journal_mode=WAL")
conn.commit()
Base.metadata.create_all(bind=engine) Base.metadata.create_all(bind=engine)
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试) # 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
@@ -45,6 +63,7 @@ def init_db():
("token_account", "total_consumed", "BIGINT DEFAULT 0"), ("token_account", "total_consumed", "BIGINT DEFAULT 0"),
("token_account", "created_at", "TIMESTAMP"), ("token_account", "created_at", "TIMESTAMP"),
("token_account", "updated_at", "TIMESTAMP"), ("token_account", "updated_at", "TIMESTAMP"),
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
) )
_normalize_optional_unique_values() _normalize_optional_unique_values()
_normalize_takeover_delays() _normalize_takeover_delays()
+2 -1
View File
@@ -120,6 +120,7 @@ class TakeoverMessage(Base):
direction = Column(String, nullable=False) # incoming | outgoing direction = Column(String, nullable=False) # incoming | outgoing
message_type = Column(Integer, default=0) message_type = Column(Integer, default=0)
content = Column(Text, default="") content = Column(Text, default="")
attachment_id = Column(String, nullable=True)
is_avatar = Column(Boolean, default=False) is_avatar = Column(Boolean, default=False)
send_time = Column(DateTime, nullable=False) send_time = Column(DateTime, nullable=False)
created_at = Column(DateTime, server_default=func.now()) created_at = Column(DateTime, server_default=func.now())
@@ -263,7 +264,7 @@ class ChatAttachment(Base):
__tablename__ = "chat_attachments" __tablename__ = "chat_attachments"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex) id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, default="", index=True) avatar_id = Column(String, nullable=False, default="", index=True)
uploader_kind = Column(String, default="owner") # owner | public uploader_kind = Column(String, default="owner") # owner | public | boxim
filename = Column(String, default="") filename = Column(String, default="")
mime_type = Column(String, default="") mime_type = Column(String, default="")
file_size = Column(Integer, default=0) file_size = Column(Integer, default=0)
+32 -5
View File
@@ -236,11 +236,40 @@ async def _analyze_uploaded_image(
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024)))) max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
content = await file.read(max_bytes + 1) content = await file.read(max_bytes + 1)
filename = os.path.basename(file.filename or "图片")[:255] filename = os.path.basename(file.filename or "图片")[:255]
try:
return _analyze_image_bytes(
db,
avatar,
content,
filename=filename,
mime_type=file.content_type or "",
uploader_kind=uploader_kind,
)
except ImageValidationError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
except InsufficientTokensError:
raise
except RuntimeError as exc:
raise HTTPException(status_code=502, detail=str(exc)) from exc
finally:
content = b""
def _analyze_image_bytes(
db: Session,
avatar: Avatar,
content: bytes,
*,
filename: str,
mime_type: str,
uploader_kind: str,
) -> ChatAttachment:
"""Analyze image bytes from either HTTP upload or BOXIM without persisting raw data."""
attachment = ChatAttachment( attachment = ChatAttachment(
avatar_id=avatar.id, avatar_id=avatar.id,
uploader_kind=uploader_kind, uploader_kind=uploader_kind,
filename=filename, filename=filename,
mime_type=(file.content_type or "")[:100], mime_type=(mime_type or "")[:100],
file_size=len(content), file_size=len(content),
status="processing", status="processing",
expires_at=_attachment_expiry(), expires_at=_attachment_expiry(),
@@ -312,7 +341,7 @@ async def _analyze_uploaded_image(
attachment.status = "failed" attachment.status = "failed"
attachment.warning = str(exc) attachment.warning = str(exc)
db.commit() db.commit()
raise HTTPException(status_code=400, detail=str(exc)) from exc raise
except InsufficientTokensError: except InsufficientTokensError:
attachment.status = "failed" attachment.status = "failed"
attachment.warning = "积分余额不足" attachment.warning = "积分余额不足"
@@ -328,9 +357,7 @@ async def _analyze_uploaded_image(
avatar.id, avatar.id,
type(exc).__name__, type(exc).__name__,
) )
raise HTTPException(status_code=502, detail=str(exc)) from exc raise
finally:
content = b""
def _normalize_question(value: str) -> str: def _normalize_question(value: str) -> str:
@@ -0,0 +1,151 @@
"""Parse and safely download image payloads from BOXIM private messages."""
import ipaddress
import json
import os
import socket
from dataclasses import dataclass
from pathlib import PurePosixPath
from urllib.parse import unquote, urljoin, urlsplit
import httpx
MAX_REDIRECTS = 3
class BoxIMImageError(RuntimeError):
pass
@dataclass(frozen=True)
class DownloadedBoxIMImage:
content: bytes
filename: str
mime_type: str
source_url: str
def parse_boxim_image_url(content: str, *, base_url: str = "") -> str:
try:
payload = json.loads(content or "")
except (TypeError, ValueError) as exc:
raise BoxIMImageError("BOXIM 图片消息格式无效") from exc
if not isinstance(payload, dict):
raise BoxIMImageError("BOXIM 图片消息格式无效")
value = payload.get("originUrl") or payload.get("thumbUrl") or payload.get("url")
if not isinstance(value, str) or not value.strip():
raise BoxIMImageError("BOXIM 图片消息缺少图片地址")
value = value.strip()
if value.startswith("/"):
if not base_url:
raise BoxIMImageError("BOXIM 图片地址不完整")
value = urljoin(f"{base_url.rstrip('/')}/", value)
return value
def _configured_hosts(name: str) -> set[str]:
return {
value.strip().lower().rstrip(".")
for value in os.getenv(name, "").split(",")
if value.strip()
}
def _host_matches(host: str, configured: set[str]) -> bool:
return any(host == value or host.endswith(f".{value}") for value in configured)
def _resolved_addresses(host: str, port: int) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
try:
return {
ipaddress.ip_address(item[4][0])
for item in socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
}
except (OSError, ValueError) as exc:
raise BoxIMImageError("BOXIM 图片地址无法解析") from exc
def _is_safe_remote_url(url: str) -> None:
parsed = urlsplit(url)
scheme = parsed.scheme.lower()
allow_http = os.getenv("BOXIM_IMAGE_ALLOW_HTTP", "").lower() in {"1", "true", "yes"}
if scheme not in ({"https", "http"} if allow_http else {"https"}):
raise BoxIMImageError("BOXIM 图片地址必须使用 HTTPS")
if parsed.username or parsed.password or not parsed.hostname:
raise BoxIMImageError("BOXIM 图片地址无效")
host = parsed.hostname.lower().rstrip(".")
allowed_hosts = _configured_hosts("BOXIM_IMAGE_ALLOWED_HOSTS")
if allowed_hosts and not _host_matches(host, allowed_hosts):
raise BoxIMImageError("BOXIM 图片地址不在允许的域名范围内")
private_hosts = _configured_hosts("BOXIM_IMAGE_PRIVATE_HOSTS")
try:
addresses = {ipaddress.ip_address(host)}
except ValueError:
addresses = _resolved_addresses(host, parsed.port or (443 if scheme == "https" else 80))
if not addresses:
raise BoxIMImageError("BOXIM 图片地址无法解析")
if _host_matches(host, private_hosts):
return
if any(not address.is_global for address in addresses):
raise BoxIMImageError("BOXIM 图片地址指向受限网络")
def _filename_from_url(url: str) -> str:
value = unquote(PurePosixPath(urlsplit(url).path).name).strip()
value = value.replace("\x00", "")
return (value or "boxim-image")[:255]
def download_boxim_image(
content: str,
*,
base_url: str = "",
transport: httpx.BaseTransport | None = None,
) -> DownloadedBoxIMImage:
"""Download one BOXIM image without redirects or oversized responses escaping checks."""
url = parse_boxim_image_url(content, base_url=base_url)
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
timeout = max(1.0, min(float(os.getenv("BOXIM_IMAGE_TIMEOUT_SECONDS", "15")), 60.0))
with httpx.Client(
timeout=timeout,
follow_redirects=False,
trust_env=False,
transport=transport,
) as client:
for _ in range(MAX_REDIRECTS + 1):
_is_safe_remote_url(url)
try:
with client.stream("GET", url, headers={"Accept": "image/*"}) as response:
if response.status_code in {301, 302, 303, 307, 308}:
location = response.headers.get("location", "").strip()
if not location:
raise BoxIMImageError("BOXIM 图片跳转地址无效")
url = urljoin(url, location)
continue
response.raise_for_status()
raw_length = response.headers.get("content-length", "")
if raw_length.isdigit() and int(raw_length) > max_bytes:
raise BoxIMImageError("BOXIM 图片超过大小限制")
chunks = bytearray()
for chunk in response.iter_bytes():
chunks.extend(chunk)
if len(chunks) > max_bytes:
raise BoxIMImageError("BOXIM 图片超过大小限制")
if not chunks:
raise BoxIMImageError("BOXIM 图片内容为空")
return DownloadedBoxIMImage(
content=bytes(chunks),
filename=_filename_from_url(url),
mime_type=response.headers.get("content-type", "").split(";", 1)[0][:100],
source_url=url,
)
except BoxIMImageError:
raise
except (httpx.HTTPError, OSError) as exc:
raise BoxIMImageError("BOXIM 图片下载失败") from exc
raise BoxIMImageError("BOXIM 图片跳转次数过多")
@@ -3,6 +3,7 @@
import asyncio import asyncio
import hashlib import hashlib
import logging import logging
import os
import re import re
import secrets import secrets
import time import time
@@ -13,12 +14,19 @@ from sqlalchemy.orm import Session
from models import ( from models import (
Avatar, Avatar,
ChatAttachment,
TakeoverCursor, TakeoverCursor,
TakeoverMessage, TakeoverMessage,
TakeoverReplyTask, TakeoverReplyTask,
User, User,
) )
from services.boxim_client import BoxIMClient, BoxIMError from services.boxim_client import BoxIMClient, BoxIMError
from services.boxim_image_service import (
BoxIMImageError,
download_boxim_image,
parse_boxim_image_url,
)
from services.vision_service import ImageValidationError
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -37,6 +45,12 @@ HUMAN_PAUSE_SECONDS = 600
RATE_LIMIT_WINDOW_SECONDS = 300 RATE_LIMIT_WINDOW_SECONDS = 300
RATE_LIMIT_MAX_REPLIES = 5 RATE_LIMIT_MAX_REPLIES = 5
AVATAR_LOCAL_ID_PREFIX = "880" AVATAR_LOCAL_ID_PREFIX = "880"
BOXIM_TEXT_MESSAGE_TYPE = 0
BOXIM_IMAGE_MESSAGE_TYPE = 1
BOXIM_IMAGE_PROMPT = "请看看这张图片。"
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
MAX_RECENT_IMAGE_CONTEXTS = 3
def _utcnow() -> datetime: def _utcnow() -> datetime:
@@ -107,6 +121,14 @@ def _configured_reply_delay(avatar: Avatar, fallback: int | None = None) -> int:
return delay return delay
def _event_prompt(event: TakeoverMessage) -> str:
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
return event.content.strip()
if event.message_type == BOXIM_IMAGE_MESSAGE_TYPE:
return BOXIM_IMAGE_PROMPT
return ""
class TakeoverService: class TakeoverService:
"""Poll BOXIM, honor the owner grace period, then generate and send one reply.""" """Poll BOXIM, honor the owner grace period, then generate and send one reply."""
@@ -129,6 +151,7 @@ class TakeoverService:
self._sessions: dict[str, dict] = {} self._sessions: dict[str, dict] = {}
self._poll_lock = asyncio.Lock() self._poll_lock = asyncio.Lock()
self._process_lock = asyncio.Lock() self._process_lock = asyncio.Lock()
self._persist_lock = asyncio.Lock()
async def poll_and_process_messages(self): async def poll_and_process_messages(self):
"""Run one complete cycle for callers that do not use the split scheduler.""" """Run one complete cycle for callers that do not use the split scheduler."""
@@ -389,13 +412,6 @@ class TakeoverService:
max_message_id = _numeric_id(cursor.last_message_id) max_message_id = _numeric_id(cursor.last_message_id)
read_receipts: dict[str, int] = {} read_receipts: dict[str, int] = {}
for message in messages: for message in messages:
self._record_message(
db,
avatar,
cursor.boxim_owner_id,
message,
schedule_reply=not priming,
)
message_id = _numeric_id(message.get("id")) message_id = _numeric_id(message.get("id"))
max_message_id = max(max_message_id, message_id) max_message_id = max(max_message_id, message_id)
send_id = str(message.get("sendId") or "") send_id = str(message.get("sendId") or "")
@@ -410,11 +426,22 @@ class TakeoverService:
session["access_token"], peer_id, message_id session["access_token"], peer_id, message_id
) )
cursor.last_message_id = str(max_message_id) # Keep SQLite write transactions short. The read-receipt request above
cursor.initialized = True # can block on the network and must not hold the database write lock.
cursor.last_polled_at = self.now() async with self._persist_lock:
cursor.last_error = "" for message in messages:
db.commit() self._record_message(
db,
avatar,
cursor.boxim_owner_id,
message,
schedule_reply=not priming,
)
cursor.last_message_id = str(max_message_id)
cursor.initialized = True
cursor.last_polled_at = self.now()
cursor.last_error = ""
db.commit()
return True return True
except Exception: except Exception:
db.rollback() db.rollback()
@@ -498,8 +525,27 @@ class TakeoverService:
if not is_avatar: if not is_avatar:
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied") self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_replied")
return return
if not schedule_reply or event.message_type != 0 or not event.content.strip(): if not schedule_reply or event.message_type not in {
BOXIM_TEXT_MESSAGE_TYPE,
BOXIM_IMAGE_MESSAGE_TYPE,
}:
return return
if event.message_type == BOXIM_TEXT_MESSAGE_TYPE and not event.content.strip():
return
if event.message_type == BOXIM_IMAGE_MESSAGE_TYPE:
try:
parse_boxim_image_url(
event.content,
base_url=getattr(self.boxim, "im_base_url", ""),
)
except BoxIMImageError as exc:
logger.warning(
"Ignored invalid BOXIM image message %s for avatar %s: %s",
message_id,
avatar.id,
exc,
)
return
if (now - send_time).total_seconds() > self.max_message_age_seconds: if (now - send_time).total_seconds() > self.max_message_age_seconds:
logger.info( logger.info(
"Ignored stale BOXIM message %s for avatar %s (age=%ss)", "Ignored stale BOXIM message %s for avatar %s (age=%ss)",
@@ -606,7 +652,16 @@ class TakeoverService:
task.status = "cancelled" task.status = "cancelled"
task.cancel_reason = "newer_incoming_message" task.cancel_reason = "newer_incoming_message"
task.locked_at = None task.locked_at = None
prompt_parts.append(event.content.strip()) if event.message_type == BOXIM_TEXT_MESSAGE_TYPE:
for image_event in self._recent_unhandled_images(
db,
avatar,
event,
source_ids,
):
prompt_parts.append(_event_prompt(image_event))
source_ids.append(image_event.boxim_message_id)
prompt_parts.append(_event_prompt(event))
source_ids.append(event.boxim_message_id) source_ids.append(event.boxim_message_id)
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:] prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
due_at = max( due_at = max(
@@ -631,6 +686,54 @@ class TakeoverService:
) )
) )
@staticmethod
def _recent_unhandled_images(
db: Session,
avatar: Avatar,
event: TakeoverMessage,
current_source_ids: list[str],
) -> list[TakeoverMessage]:
"""Recover a recent image that an older deployment recorded without a task."""
threshold = event.send_time - timedelta(seconds=IMAGE_CONTEXT_LOOKBACK_SECONDS)
candidates = (
db.query(TakeoverMessage)
.filter(
TakeoverMessage.avatar_id == avatar.id,
TakeoverMessage.owner_id == avatar.owner_id,
TakeoverMessage.peer_id == event.peer_id,
TakeoverMessage.direction == "incoming",
TakeoverMessage.message_type == BOXIM_IMAGE_MESSAGE_TYPE,
TakeoverMessage.is_avatar.is_(False),
TakeoverMessage.send_time >= threshold,
TakeoverMessage.send_time <= event.send_time,
)
.order_by(TakeoverMessage.send_time.desc())
.limit(MAX_RECENT_IMAGE_CONTEXTS)
.all()
)
if not candidates:
return []
handled_ids = set(current_source_ids)
task_sources = (
db.query(TakeoverReplyTask.source_message_ids)
.filter(
TakeoverReplyTask.avatar_id == avatar.id,
TakeoverReplyTask.owner_id == avatar.owner_id,
TakeoverReplyTask.peer_id == event.peer_id,
TakeoverReplyTask.created_at >= threshold,
)
.all()
)
for (source_message_ids,) in task_sources:
handled_ids.update(source_message_ids or [])
return [
image
for image in reversed(candidates)
if image.boxim_message_id not in handled_ids
]
async def _prepare_replies(self) -> int: async def _prepare_replies(self) -> int:
db = self.session_factory() db = self.session_factory()
try: try:
@@ -665,6 +768,50 @@ class TakeoverService:
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids)) results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
return sum(bool(result) for result in results) return sum(bool(result) for result in results)
def _takeover_image_attachment(
self,
db: Session,
avatar: Avatar,
event: TakeoverMessage,
) -> ChatAttachment:
now = self.now()
if event.attachment_id:
cached = db.get(ChatAttachment, event.attachment_id)
if cached and cached.status == "ready" and cached.expires_at > now:
cached.used_at = now
db.commit()
return cached
downloaded = download_boxim_image(
event.content,
base_url=getattr(
self.boxim,
"im_base_url",
os.getenv("BOXIM_API_BASE_URL", "https://im.99hui.com/api"),
),
)
from routers.chat import _analyze_image_bytes
attachment = _analyze_image_bytes(
db,
avatar,
downloaded.content,
filename=downloaded.filename,
mime_type=downloaded.mime_type,
uploader_kind="boxim",
)
event.attachment_id = attachment.id
attachment.used_at = now
db.commit()
logger.info(
"BOXIM image analyzed message=%s attachment=%s avatar=%s category=%s",
event.boxim_message_id,
attachment.id,
avatar.id,
attachment.category,
)
return attachment
def _generate_reply(self, task_id: str) -> bool: def _generate_reply(self, task_id: str) -> bool:
db = self.session_factory() db = self.session_factory()
try: try:
@@ -683,6 +830,21 @@ class TakeoverService:
db.commit() db.commit()
excluded_ids = set(task.source_message_ids or []) excluded_ids = set(task.source_message_ids or [])
source_events = {
event.boxim_message_id: event
for event in (
db.query(TakeoverMessage)
.filter(
TakeoverMessage.owner_id == task.owner_id,
TakeoverMessage.peer_id == task.peer_id,
TakeoverMessage.avatar_id == task.avatar_id,
TakeoverMessage.boxim_message_id.in_(excluded_ids),
)
.all()
if excluded_ids
else []
)
}
events = ( events = (
db.query(TakeoverMessage) db.query(TakeoverMessage)
.filter( .filter(
@@ -694,9 +856,31 @@ class TakeoverService:
.limit(30) .limit(30)
.all() .all()
) )
image_attachments = []
image_failed = False
for message_id in (task.source_message_ids or [])[-3:]:
event = source_events.get(message_id)
if not event or event.message_type != BOXIM_IMAGE_MESSAGE_TYPE:
continue
try:
image_attachments.append(
self._takeover_image_attachment(db, avatar, event)
)
except (BoxIMImageError, ImageValidationError) as exc:
image_failed = True
logger.warning(
"BOXIM image unavailable message=%s avatar=%s: %s",
event.boxim_message_id,
avatar.id,
exc,
)
history = [] history = []
for event in reversed(events): for event in reversed(events):
if event.boxim_message_id in excluded_ids or not event.content.strip(): if (
event.boxim_message_id in excluded_ids
or event.message_type != BOXIM_TEXT_MESSAGE_TYPE
or not event.content.strip()
):
continue continue
if event.direction == "incoming" and event.is_avatar: if event.direction == "incoming" and event.is_avatar:
continue continue
@@ -708,10 +892,21 @@ class TakeoverService:
) )
history = history[-10:] history = history[-10:]
from routers.chat import _resolve_reply from routers.chat import _attachment_contexts, _resolve_reply
result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover") image_contexts = _attachment_contexts(image_attachments)
answer = _plain_text_reply(result.get("answer", "")) if image_failed and not image_contexts:
answer = BOXIM_IMAGE_UNAVAILABLE_REPLY
else:
result = _resolve_reply(
db,
avatar,
task.prompt,
history,
usage_source="takeover",
image_contexts=image_contexts,
)
answer = _plain_text_reply(result.get("answer", ""))
db.refresh(task) db.refresh(task)
if task.status != "generating": if task.status != "generating":
return False return False
@@ -0,0 +1,71 @@
import ipaddress
import json
import httpx
import pytest
from services.boxim_image_service import (
BoxIMImageError,
download_boxim_image,
parse_boxim_image_url,
)
def test_parse_boxim_image_prefers_origin_and_supports_relative_url():
content = json.dumps({"originUrl": "/files/original.png", "thumbUrl": "/thumb.png"})
assert parse_boxim_image_url(content, base_url="https://im.example/api") == (
"https://im.example/files/original.png"
)
def test_download_boxim_image_streams_public_https(monkeypatch):
monkeypatch.setattr(
"services.boxim_image_service._resolved_addresses",
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
)
transport = httpx.MockTransport(
lambda request: httpx.Response(
200,
headers={"content-type": "image/png"},
content=b"png-bytes",
request=request,
)
)
image = download_boxim_image(
json.dumps({"originUrl": "https://cdn.example/case%20photo.png"}),
transport=transport,
)
assert image.content == b"png-bytes"
assert image.filename == "case photo.png"
assert image.mime_type == "image/png"
def test_download_boxim_image_rejects_private_network_url():
with pytest.raises(BoxIMImageError, match="受限网络"):
download_boxim_image(
json.dumps({"originUrl": "https://127.0.0.1/private.png"}),
transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
)
def test_download_boxim_image_stops_oversized_stream(monkeypatch):
monkeypatch.setenv("CHAT_IMAGE_MAX_BYTES", "1024")
monkeypatch.setattr(
"services.boxim_image_service._resolved_addresses",
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
)
transport = httpx.MockTransport(
lambda request: httpx.Response(
200,
headers={"content-length": "2048"},
request=request,
)
)
with pytest.raises(BoxIMImageError, match="超过大小限制"):
download_boxim_image(
json.dumps({"originUrl": "https://cdn.example/large.png"}),
transport=transport,
)
@@ -17,6 +17,7 @@ from routers.chat import (
_resolve_reply, _resolve_reply,
) )
from services.chat_attachment_service import purge_expired_chat_attachments from services.chat_attachment_service import purge_expired_chat_attachments
from services.token_billing import InsufficientTokensError
from services.vision_service import PreparedImage from services.vision_service import PreparedImage
@@ -75,6 +76,22 @@ def test_non_owner_cannot_upload_chat_image(authorization_context):
assert response.status_code == 403 assert response.status_code == 403
def test_image_upload_preserves_insufficient_points_response(authorization_context):
context = authorization_context
with patch(
"routers.chat._analyze_image_bytes",
side_effect=InsufficientTokensError("积分余额不足"),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("private.png", b"image-bytes", "image/png")},
)
assert response.status_code == 402
assert response.json()["detail"] == "积分余额不足"
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context): def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
context = authorization_context context = authorization_context
db = SessionLocal() db = SessionLocal()
@@ -1,6 +1,7 @@
"""End-to-end service tests for BOXIM takeover timing and human priority.""" """End-to-end service tests for BOXIM takeover timing and human priority."""
import asyncio import asyncio
import json
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from threading import Barrier from threading import Barrier
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
@@ -10,8 +11,9 @@ from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
from database import Base from database import Base
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User from models import Avatar, ChatAttachment, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
from services.boxim_client import BoxIMError from services.boxim_client import BoxIMError
from services.boxim_image_service import DownloadedBoxIMImage
from services.takeover_service import ( from services.takeover_service import (
AVATAR_LOCAL_ID_PREFIX, AVATAR_LOCAL_ID_PREFIX,
TakeoverService, TakeoverService,
@@ -83,6 +85,29 @@ class ConcurrentPollingBoxIM(FakeBoxIM):
return [] return []
class ConcurrentMessagePollingBoxIM(ConcurrentPollingBoxIM):
async def fetch_private_messages(self, access_token, min_id="0"):
await super().fetch_private_messages(access_token, min_id)
owner_id = 100 if access_token == "prod-huihui-token" else 101
return [
{
"id": owner_id,
"localId": owner_id,
"sendId": owner_id + 100,
"recvId": owner_id,
"sendTime": 1_700_000_000_000,
"type": 0,
"content": "并发写入测试",
}
]
async def mark_private_messages_read(self, access_token, friend_id, message_id):
await asyncio.sleep(0.05)
self.read_receipts.append(
{"friendId": str(friend_id), "messageId": str(message_id)}
)
@pytest.fixture @pytest.fixture
def service_context(tmp_path): def service_context(tmp_path):
engine = create_engine( engine = create_engine(
@@ -171,6 +196,199 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
db.close() db.close()
@pytest.mark.asyncio
async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
boxim.messages.append(
{
"id": 111,
"localId": 111,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps(
{
"originUrl": "https://cdn.example/case.png",
"thumbUrl": "https://cdn.example/case-thumb.png",
}
),
}
)
await service.poll_and_process_messages()
db = session_factory()
try:
scheduled = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
assert scheduled.status == "pending"
assert scheduled.prompt == "请看看这张图片。"
finally:
db.close()
clock.advance(3)
def analyze(db, avatar, content, **kwargs):
assert content == b"image-content"
attachment = ChatAttachment(
avatar_id=avatar.id,
uploader_kind=kwargs["uploader_kind"],
filename=kwargs["filename"],
mime_type="image/jpeg",
file_size=len(content),
status="ready",
category="medical_document",
summary="一张门诊病例",
extracted_text="主诉:咳嗽三天",
structured_data={"medical": {"chief_complaint": "咳嗽三天"}},
warning="请核对原始资料",
expires_at=clock.now() + timedelta(hours=24),
)
db.add(attachment)
db.commit()
db.refresh(attachment)
return attachment
downloaded = DownloadedBoxIMImage(
content=b"image-content",
filename="case.png",
mime_type="image/png",
source_url="https://cdn.example/case.png",
)
with (
patch("services.takeover_service.download_boxim_image", return_value=downloaded),
patch("routers.chat._analyze_image_bytes", side_effect=analyze) as analyzer,
patch("routers.chat._resolve_reply", return_value={"answer": "这份资料里写的是咳嗽三天。"}) as resolver,
):
await service.poll_and_process_messages()
analyzer.assert_called_once()
assert resolver.call_args.args[2] == "请看看这张图片。"
image_contexts = resolver.call_args.kwargs["image_contexts"]
assert image_contexts[0]["summary"] == "一张门诊病例"
assert image_contexts[0]["extractedText"] == "主诉:咳嗽三天"
assert [item["content"] for item in boxim.sent] == ["这份资料里写的是咳嗽三天。"]
db = session_factory()
try:
event = db.query(TakeoverMessage).filter_by(boxim_message_id="111").one()
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="111").one()
assert event.attachment_id
assert db.get(ChatAttachment, event.attachment_id).uploader_kind == "boxim"
assert task.status == "sent"
with patch(
"services.takeover_service.download_boxim_image",
side_effect=AssertionError("cached image must not be downloaded again"),
):
cached = service._takeover_image_attachment(db, db.get(Avatar, "avatar-1"), event)
assert cached.id == event.attachment_id
finally:
db.close()
@pytest.mark.asyncio
async def test_followup_text_recovers_recent_image_recorded_without_task(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
image_message = {
"id": 113,
"localId": 113,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps(
{
"originUrl": "https://cdn.example/case.png",
"thumbUrl": "https://cdn.example/case-thumb.png",
}
),
}
db = session_factory()
try:
avatar = db.get(Avatar, "avatar-1")
service._record_message(db, avatar, "100", image_message, schedule_reply=False)
cursor = db.query(TakeoverCursor).one()
cursor.last_message_id = "113"
db.commit()
finally:
db.close()
clock.advance(60)
boxim.messages.extend(
[
image_message,
{
"id": 114,
"localId": 114,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 0,
"content": "请帮我看看这张图",
},
]
)
await service.poll_messages()
db = session_factory()
try:
task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="114").one()
assert task.source_message_ids == ["113", "114"]
assert task.prompt == "请看看这张图片。\n请帮我看看这张图"
finally:
db.close()
clock.advance(1)
boxim.messages.append(
{
"id": 115,
"localId": 115,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 0,
"content": "图里写了什么",
}
)
await service.poll_messages()
db = session_factory()
try:
latest = db.query(TakeoverReplyTask).filter_by(trigger_message_id="115").one()
assert latest.source_message_ids == ["113", "114", "115"]
assert latest.source_message_ids.count("113") == 1
finally:
db.close()
@pytest.mark.asyncio
async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
boxim.messages.append(
{
"id": 112,
"localId": 112,
"sendId": 200,
"recvId": 100,
"sendTime": clock.millis(),
"type": 1,
"content": json.dumps({"width": 100, "height": 100}),
}
)
await service.poll_and_process_messages()
db = session_factory()
try:
assert db.query(TakeoverMessage).filter_by(boxim_message_id="112").one()
assert db.query(TakeoverReplyTask).count() == 0
finally:
db.close()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_default_reply_delay_is_three_minutes(service_context): async def test_default_reply_delay_is_three_minutes(service_context):
session_factory, service, boxim, clock = service_context session_factory, service, boxim, clock = service_context
@@ -230,7 +448,7 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
finally: finally:
db.close() db.close()
boxim = ConcurrentPollingBoxIM() boxim = ConcurrentMessagePollingBoxIM()
service = TakeoverService( service = TakeoverService(
session_factory, session_factory,
boxim, boxim,
@@ -244,6 +462,8 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
db = session_factory() db = session_factory()
try: try:
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2 assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
assert db.query(TakeoverMessage).count() == 2
assert len(boxim.read_receipts) == 2
finally: finally:
db.close() db.close()