Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
672019830d | ||
|
|
e720baa21e |
@@ -51,22 +51,32 @@ def embed(texts):
|
|||||||
if api_url:
|
if api_url:
|
||||||
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
||||||
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
||||||
payload = json.dumps({"input": texts, "model": model}).encode("utf-8")
|
try:
|
||||||
req = urllib.request.Request(
|
batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10")))
|
||||||
api_url,
|
except ValueError:
|
||||||
data=payload,
|
batch_size = 10
|
||||||
headers={
|
embeddings = []
|
||||||
"Content-Type": "application/json",
|
for start in range(0, len(texts), batch_size):
|
||||||
"Authorization": f"Bearer {api_key}" if api_key else "",
|
batch = texts[start:start + batch_size]
|
||||||
},
|
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
||||||
method="POST",
|
req = urllib.request.Request(
|
||||||
)
|
api_url,
|
||||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
data=payload,
|
||||||
data = json.loads(resp.read().decode("utf-8"))
|
headers={
|
||||||
items = data["data"]
|
"Content-Type": "application/json",
|
||||||
if items and "index" in items[0]:
|
"Authorization": f"Bearer {api_key}" if api_key else "",
|
||||||
items = sorted(items, key=lambda x: x["index"])
|
},
|
||||||
return [item["embedding"] for item in items]
|
method="POST",
|
||||||
|
)
|
||||||
|
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||||
|
data = json.loads(resp.read().decode("utf-8"))
|
||||||
|
items = data["data"]
|
||||||
|
if items and "index" in items[0]:
|
||||||
|
items = sorted(items, key=lambda x: x["index"])
|
||||||
|
if len(items) != len(batch):
|
||||||
|
raise ValueError("embedding response count does not match request")
|
||||||
|
embeddings.extend(item["embedding"] for item in items)
|
||||||
|
return embeddings
|
||||||
return _hash_embedding(texts)
|
return _hash_embedding(texts)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -143,14 +143,28 @@ def on_startup():
|
|||||||
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
|
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
|
||||||
takeover_scheduler = AsyncIOScheduler()
|
takeover_scheduler = AsyncIOScheduler()
|
||||||
takeover_scheduler.add_job(
|
takeover_scheduler.add_job(
|
||||||
takeover_service.poll_and_process_messages,
|
takeover_service.poll_messages,
|
||||||
trigger=IntervalTrigger(seconds=poll_interval),
|
trigger=IntervalTrigger(seconds=poll_interval),
|
||||||
id="takeover_message_poll",
|
id="takeover_message_poll",
|
||||||
max_instances=1,
|
max_instances=1,
|
||||||
coalesce=True,
|
coalesce=True,
|
||||||
)
|
)
|
||||||
|
process_interval = max(
|
||||||
|
0.25, float(os.getenv("TAKEOVER_PROCESS_INTERVAL_SECONDS", "0.5"))
|
||||||
|
)
|
||||||
|
takeover_scheduler.add_job(
|
||||||
|
takeover_service.process_reply_tasks,
|
||||||
|
trigger=IntervalTrigger(seconds=process_interval),
|
||||||
|
id="takeover_reply_process",
|
||||||
|
max_instances=1,
|
||||||
|
coalesce=True,
|
||||||
|
)
|
||||||
takeover_scheduler.start()
|
takeover_scheduler.start()
|
||||||
logger.info("BOXIM takeover scheduler started (interval=%ss)", poll_interval)
|
logger.info(
|
||||||
|
"BOXIM takeover scheduler started (poll=%ss, process=%ss)",
|
||||||
|
poll_interval,
|
||||||
|
process_interval,
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
stop_takeover_scheduler()
|
stop_takeover_scheduler()
|
||||||
logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}")
|
logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}")
|
||||||
|
|||||||
@@ -86,24 +86,34 @@ class TakeoverService:
|
|||||||
self.reply_delay_seconds = reply_delay_seconds
|
self.reply_delay_seconds = reply_delay_seconds
|
||||||
self.now = now
|
self.now = now
|
||||||
self._sessions: dict[str, dict] = {}
|
self._sessions: dict[str, dict] = {}
|
||||||
self._run_lock = asyncio.Lock()
|
self._poll_lock = asyncio.Lock()
|
||||||
|
self._process_lock = asyncio.Lock()
|
||||||
|
|
||||||
async def poll_and_process_messages(self):
|
async def poll_and_process_messages(self):
|
||||||
"""Run one complete cycle; polling always happens before reply dispatch."""
|
"""Run one complete cycle for callers that do not use the split scheduler."""
|
||||||
if self._run_lock.locked():
|
await self.poll_messages()
|
||||||
|
await self.process_reply_tasks()
|
||||||
|
|
||||||
|
async def poll_messages(self):
|
||||||
|
"""Fetch BOXIM events without blocking reply generation and dispatch."""
|
||||||
|
if self._poll_lock.locked():
|
||||||
return
|
return
|
||||||
async with self._run_lock:
|
async with self._poll_lock:
|
||||||
self._recover_stuck_tasks()
|
self._recover_stuck_tasks()
|
||||||
avatar_ids = self._enabled_avatar_ids()
|
avatar_ids = self._enabled_avatar_ids()
|
||||||
self._cancel_disabled_tasks(set(avatar_ids))
|
self._cancel_disabled_tasks(set(avatar_ids))
|
||||||
for avatar_id in avatar_ids:
|
for avatar_id in avatar_ids:
|
||||||
await self._sync_avatar(avatar_id)
|
await self._sync_avatar(avatar_id)
|
||||||
|
|
||||||
generated = await self._prepare_replies()
|
async def process_reply_tasks(self):
|
||||||
if generated:
|
"""Generate and send replies independently from BOXIM's long poll."""
|
||||||
# Catch a human reply sent while the model was preparing its answer.
|
if self._process_lock.locked():
|
||||||
for avatar_id in avatar_ids:
|
return
|
||||||
await self._sync_avatar(avatar_id)
|
async with self._process_lock:
|
||||||
|
self._recover_stuck_tasks()
|
||||||
|
avatar_ids = set(self._enabled_avatar_ids())
|
||||||
|
self._cancel_disabled_tasks(avatar_ids)
|
||||||
|
await self._prepare_replies()
|
||||||
await self._dispatch_ready_replies()
|
await self._dispatch_ready_replies()
|
||||||
|
|
||||||
def _enabled_avatar_ids(self) -> list[str]:
|
def _enabled_avatar_ids(self) -> list[str]:
|
||||||
@@ -454,11 +464,19 @@ class TakeoverService:
|
|||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
generated = 0
|
if not task_ids:
|
||||||
for task_id in task_ids:
|
return 0
|
||||||
if await asyncio.to_thread(self._generate_reply, task_id):
|
|
||||||
generated += 1
|
# Each conversation owns its task, so unrelated contacts can generate in
|
||||||
return generated
|
# parallel instead of one slow model response delaying every other peer.
|
||||||
|
semaphore = asyncio.Semaphore(4)
|
||||||
|
|
||||||
|
async def generate(task_id: str) -> bool:
|
||||||
|
async with semaphore:
|
||||||
|
return await asyncio.to_thread(self._generate_reply, task_id)
|
||||||
|
|
||||||
|
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
|
||||||
|
return sum(bool(result) for result in results)
|
||||||
|
|
||||||
def _generate_reply(self, task_id: str) -> bool:
|
def _generate_reply(self, task_id: str) -> bool:
|
||||||
db = self.session_factory()
|
db = self.session_factory()
|
||||||
@@ -548,8 +566,8 @@ class TakeoverService:
|
|||||||
finally:
|
finally:
|
||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
for task_id in task_ids:
|
if task_ids:
|
||||||
await self._send_task(task_id)
|
await asyncio.gather(*(self._send_task(task_id) for task_id in task_ids))
|
||||||
|
|
||||||
async def _send_task(self, task_id: str) -> bool:
|
async def _send_task(self, task_id: str) -> bool:
|
||||||
db = self.session_factory()
|
db = self.session_factory()
|
||||||
|
|||||||
@@ -1,10 +1,26 @@
|
|||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import embeddings
|
import embeddings
|
||||||
|
|
||||||
|
|
||||||
|
class FakeResponse:
|
||||||
|
def __init__(self, payload):
|
||||||
|
self.payload = payload
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *_):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def read(self):
|
||||||
|
return json.dumps(self.payload).encode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
class TextExtractionTests(unittest.TestCase):
|
class TextExtractionTests(unittest.TestCase):
|
||||||
def write_text(self, suffix, content):
|
def write_text(self, suffix, content):
|
||||||
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
||||||
@@ -28,5 +44,33 @@ class TextExtractionTests(unittest.TestCase):
|
|||||||
embeddings.extract_text(path, ".csv")
|
embeddings.extract_text(path, ".csv")
|
||||||
|
|
||||||
|
|
||||||
|
class RemoteEmbeddingTests(unittest.TestCase):
|
||||||
|
def test_large_input_is_split_into_provider_safe_batches(self):
|
||||||
|
texts = [f"chunk-{index}" for index in range(14)]
|
||||||
|
batch_sizes = []
|
||||||
|
|
||||||
|
def fake_urlopen(request, timeout):
|
||||||
|
self.assertEqual(timeout, 30)
|
||||||
|
payload = json.loads(request.data.decode("utf-8"))
|
||||||
|
batch_sizes.append(len(payload["input"]))
|
||||||
|
return FakeResponse({
|
||||||
|
"data": [
|
||||||
|
{"index": index, "embedding": [float(text.split("-")[1])]}
|
||||||
|
for index, text in enumerate(payload["input"])
|
||||||
|
]
|
||||||
|
})
|
||||||
|
|
||||||
|
with patch.dict(os.environ, {
|
||||||
|
"EMBEDDING_API_URL": "https://embedding.example/v1/embeddings",
|
||||||
|
"EMBEDDING_API_KEY": "test-key",
|
||||||
|
"EMBEDDING_MODEL": "text-embedding-v4",
|
||||||
|
"EMBEDDING_BATCH_SIZE": "10",
|
||||||
|
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
||||||
|
result = embeddings.embed(texts)
|
||||||
|
|
||||||
|
self.assertEqual(batch_sizes, [10, 4])
|
||||||
|
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -25,7 +25,8 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
|||||||
boxim = MagicMock()
|
boxim = MagicMock()
|
||||||
mock_boxim_class.return_value = boxim
|
mock_boxim_class.return_value = boxim
|
||||||
takeover = MagicMock()
|
takeover = MagicMock()
|
||||||
takeover.poll_and_process_messages = AsyncMock()
|
takeover.poll_messages = AsyncMock()
|
||||||
|
takeover.process_reply_tasks = AsyncMock()
|
||||||
mock_takeover_class.return_value = takeover
|
mock_takeover_class.return_value = takeover
|
||||||
|
|
||||||
environment = {
|
environment = {
|
||||||
@@ -46,14 +47,18 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
|||||||
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
|
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
|
||||||
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
|
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
|
||||||
|
|
||||||
scheduler.add_job.assert_called_once()
|
assert scheduler.add_job.call_count == 2
|
||||||
scheduled_callable = scheduler.add_job.call_args.args[0]
|
poll_call, process_call = scheduler.add_job.call_args_list
|
||||||
job_options = scheduler.add_job.call_args.kwargs
|
assert poll_call.args[0] is takeover.poll_messages
|
||||||
assert scheduled_callable is takeover.poll_and_process_messages
|
assert poll_call.kwargs["id"] == "takeover_message_poll"
|
||||||
assert job_options["id"] == "takeover_message_poll"
|
assert poll_call.kwargs["trigger"].interval.total_seconds() == 1
|
||||||
assert job_options["trigger"].interval.total_seconds() == 1
|
assert poll_call.kwargs["max_instances"] == 1
|
||||||
assert job_options["max_instances"] == 1
|
assert poll_call.kwargs["coalesce"] is True
|
||||||
assert job_options["coalesce"] is True
|
assert process_call.args[0] is takeover.process_reply_tasks
|
||||||
|
assert process_call.kwargs["id"] == "takeover_reply_process"
|
||||||
|
assert process_call.kwargs["trigger"].interval.total_seconds() == 0.5
|
||||||
|
assert process_call.kwargs["max_instances"] == 1
|
||||||
|
assert process_call.kwargs["coalesce"] is True
|
||||||
scheduler.start.assert_called_once_with()
|
scheduler.start.assert_called_once_with()
|
||||||
|
|
||||||
main.takeover_scheduler = None
|
main.takeover_scheduler = None
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from threading import Barrier
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -143,6 +144,39 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
|
|||||||
db.close()
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_different_contacts_generate_without_blocking_each_other(service_context):
|
||||||
|
session_factory, service, boxim, clock = service_context
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
boxim.messages.extend(
|
||||||
|
[
|
||||||
|
{"id": 13, "localId": 31, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人甲"},
|
||||||
|
{"id": 14, "localId": 32, "sendId": 300, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人乙"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
both_generating = Barrier(2, timeout=2)
|
||||||
|
|
||||||
|
def resolve(_db, _avatar, prompt, _history):
|
||||||
|
both_generating.wait()
|
||||||
|
return {"answer": f"回复{prompt[-1]}"}
|
||||||
|
|
||||||
|
with patch("routers.chat._resolve_reply", side_effect=resolve):
|
||||||
|
await service.poll_and_process_messages()
|
||||||
|
|
||||||
|
clock.advance(3)
|
||||||
|
await service.process_reply_tasks()
|
||||||
|
assert {(item["peerId"], item["content"]) for item in boxim.sent} == {
|
||||||
|
("200", "回复甲"),
|
||||||
|
("300", "回复乙"),
|
||||||
|
}
|
||||||
|
|
||||||
|
db = session_factory()
|
||||||
|
try:
|
||||||
|
assert {task.status for task in db.query(TakeoverReplyTask).all()} == {"sent"}
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_read_receipt_failure_does_not_advance_cursor(service_context):
|
async def test_read_receipt_failure_does_not_advance_cursor(service_context):
|
||||||
session_factory, service, boxim, clock = service_context
|
session_factory, service, boxim, clock = service_context
|
||||||
|
|||||||
Reference in New Issue
Block a user