mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 08:03:15 +00:00
fix: serialize Slack run dispatch (#1591)
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
93e6bb8680
commit
eac293d8a7
3 changed files with 138 additions and 28 deletions
|
|
@ -2,8 +2,11 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk import get_client
|
||||
|
|
@ -12,6 +15,25 @@ logger = logging.getLogger(__name__)
|
|||
|
||||
MAX_QUEUED_MESSAGES = 100
|
||||
|
||||
_THREAD_RUN_LOCKS: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
def get_thread_run_lock(thread_id: str) -> asyncio.Lock:
|
||||
"""Return a per-thread-id asyncio.Lock, creating one lazily if needed."""
|
||||
lock = _THREAD_RUN_LOCKS.get(thread_id)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
_THREAD_RUN_LOCKS[thread_id] = lock
|
||||
return lock
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def thread_run_lock(thread_id: str) -> AsyncIterator[None]:
|
||||
"""Serialize run dispatch for a thread."""
|
||||
lock = get_thread_run_lock(thread_id)
|
||||
async with lock:
|
||||
yield
|
||||
|
||||
|
||||
def langgraph_url() -> str:
|
||||
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
||||
|
|
|
|||
|
|
@ -124,7 +124,7 @@ from .utils.slack_feedback import (
|
|||
process_slack_reaction_added,
|
||||
process_slack_reaction_removed,
|
||||
)
|
||||
from .utils.thread_ops import is_thread_active, queue_message_for_thread
|
||||
from .utils.thread_ops import is_thread_active, queue_message_for_thread, thread_run_lock
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -1189,36 +1189,37 @@ async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[st
|
|||
source_context={"slack_thread": configurable["slack_thread"]},
|
||||
)
|
||||
|
||||
thread_active = await is_thread_active(thread_id)
|
||||
if thread_active:
|
||||
async with thread_run_lock(thread_id):
|
||||
thread_active = await is_thread_active(thread_id)
|
||||
if thread_active:
|
||||
logger.info(
|
||||
"Thread %s is active, queuing Slack message for middleware pickup",
|
||||
thread_id,
|
||||
)
|
||||
queued_payload = {"text": prompt, "image_urls": image_urls}
|
||||
queued = await queue_message_for_thread(
|
||||
thread_id=thread_id,
|
||||
message_content=queued_payload,
|
||||
)
|
||||
if queued:
|
||||
logger.info("Slack message queued for thread %s", thread_id)
|
||||
else:
|
||||
logger.error("Failed to queue Slack message for thread %s", thread_id)
|
||||
return
|
||||
|
||||
logger.info("Creating Slack LangGraph run for thread %s", thread_id)
|
||||
run = await langgraph_client.runs.create(
|
||||
thread_id,
|
||||
"agent",
|
||||
input={"messages": [{"role": "user", "content": content_blocks}]},
|
||||
config={"configurable": configurable, "metadata": _AGENT_VERSION_METADATA},
|
||||
if_not_exists="create",
|
||||
)
|
||||
logger.info(
|
||||
"Thread %s is active, queuing Slack message for middleware pickup",
|
||||
"Slack LangGraph run %s created for thread %s",
|
||||
_run_id_for_logging(run),
|
||||
thread_id,
|
||||
)
|
||||
queued_payload = {"text": prompt, "image_urls": image_urls}
|
||||
queued = await queue_message_for_thread(
|
||||
thread_id=thread_id,
|
||||
message_content=queued_payload,
|
||||
)
|
||||
if queued:
|
||||
logger.info("Slack message queued for thread %s", thread_id)
|
||||
else:
|
||||
logger.error("Failed to queue Slack message for thread %s", thread_id)
|
||||
return
|
||||
|
||||
logger.info("Creating Slack LangGraph run for thread %s", thread_id)
|
||||
run = await langgraph_client.runs.create(
|
||||
thread_id,
|
||||
"agent",
|
||||
input={"messages": [{"role": "user", "content": content_blocks}]},
|
||||
config={"configurable": configurable, "metadata": _AGENT_VERSION_METADATA},
|
||||
if_not_exists="create",
|
||||
)
|
||||
logger.info(
|
||||
"Slack LangGraph run %s created for thread %s",
|
||||
_run_id_for_logging(run),
|
||||
thread_id,
|
||||
)
|
||||
run_id = run.get("run_id")
|
||||
if is_first_mention:
|
||||
trace_message_ts = await post_slack_trace_reply(channel_id, thread_ts, thread_id)
|
||||
|
|
|
|||
|
|
@ -739,6 +739,93 @@ def test_process_slack_mention_queues_active_thread_message(
|
|||
assert "## Latest Mention Request\ninclude this screenshot" in queued_payload["text"]
|
||||
|
||||
|
||||
def test_process_slack_mention_serializes_concurrent_run_dispatch(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
_setup_slack_mention_fakes(monkeypatch, captured)
|
||||
|
||||
thread_ts = "1700000001.000100"
|
||||
expected_thread_id = generate_thread_id_from_slack_thread("C123", thread_ts)
|
||||
first_active_started = asyncio.Event()
|
||||
finish_first_active = asyncio.Event()
|
||||
active_calls: list[str] = []
|
||||
run_creates: list[dict[str, object]] = []
|
||||
queued_messages: list[dict[str, object]] = []
|
||||
|
||||
async def fake_thread_exists(thread_id: str) -> bool:
|
||||
return False
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
active_calls.append(thread_id)
|
||||
if len(active_calls) == 1:
|
||||
first_active_started.set()
|
||||
await finish_first_active.wait()
|
||||
return bool(run_creates)
|
||||
|
||||
async def fake_queue_message_for_thread(thread_id: str, message_content: object) -> bool:
|
||||
queued_messages.append({"thread_id": thread_id, "message_content": message_content})
|
||||
return True
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, thread_id: str, graph: str, **kwargs) -> dict[str, str]:
|
||||
run_creates.append({"thread_id": thread_id, "graph": graph, "kwargs": kwargs})
|
||||
return {"run_id": f"run-{len(run_creates)}"}
|
||||
|
||||
class _FakeThreadsClientForProcess:
|
||||
async def update(self, *, thread_id: str, metadata: dict) -> None:
|
||||
captured["metadata_update"] = {"thread_id": thread_id, "metadata": metadata}
|
||||
|
||||
class _FakeLangGraphClientForProcess:
|
||||
runs = _FakeRunsClient()
|
||||
threads = _FakeThreadsClientForProcess()
|
||||
|
||||
monkeypatch.setattr(webapp, "_thread_exists", fake_thread_exists)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "queue_message_for_thread", fake_queue_message_for_thread)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClientForProcess())
|
||||
|
||||
async def run_concurrent_mentions() -> None:
|
||||
first = asyncio.create_task(
|
||||
webapp.process_slack_mention(
|
||||
{
|
||||
"channel_id": "C123",
|
||||
"thread_ts": thread_ts,
|
||||
"event_ts": "1700000000.000200",
|
||||
"user_id": "U123",
|
||||
"text": "<@UBOT> first request",
|
||||
"bot_user_id": "UBOT",
|
||||
},
|
||||
{"owner": "langchain-ai", "name": "open-swe"},
|
||||
)
|
||||
)
|
||||
await first_active_started.wait()
|
||||
second = asyncio.create_task(
|
||||
webapp.process_slack_mention(
|
||||
{
|
||||
"channel_id": "C123",
|
||||
"thread_ts": thread_ts,
|
||||
"event_ts": "1700000000.000300",
|
||||
"user_id": "U123",
|
||||
"text": "<@UBOT> second request",
|
||||
"bot_user_id": "UBOT",
|
||||
},
|
||||
{"owner": "langchain-ai", "name": "open-swe"},
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
assert active_calls == [expected_thread_id]
|
||||
finish_first_active.set()
|
||||
await asyncio.gather(first, second)
|
||||
|
||||
asyncio.run(run_concurrent_mentions())
|
||||
|
||||
assert active_calls == [expected_thread_id, expected_thread_id]
|
||||
assert len(run_creates) == 1
|
||||
assert run_creates[0]["thread_id"] == expected_thread_id
|
||||
assert queued_messages[0]["thread_id"] == expected_thread_id
|
||||
|
||||
|
||||
def test_process_slack_mention_unmapped_user_blocked_and_prompted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue