fix: serialize Slack run dispatch (#1591)

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Brace Sproul 2026-06-23 12:04:08 -07:00 • committed by GitHub
parent 93e6bb8680
commit eac293d8a7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 138 additions and 28 deletions

View file

@ -2,8 +2,11 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import logging import logging
import os import os
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any from typing import Any
from langgraph_sdk import get_client from langgraph_sdk import get_client
@ -12,6 +15,25 @@ logger = logging.getLogger(__name__)
MAX_QUEUED_MESSAGES = 100 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: def langgraph_url() -> str:
return os.environ.get("LANGGRAPH_URL") or os.environ.get( return os.environ.get("LANGGRAPH_URL") or os.environ.get(

View file

@ -124,7 +124,7 @@ from .utils.slack_feedback import (
process_slack_reaction_added, process_slack_reaction_added,
process_slack_reaction_removed, 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__) 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"]}, source_context={"slack_thread": configurable["slack_thread"]},
) )
thread_active = await is_thread_active(thread_id) async with thread_run_lock(thread_id):
if thread_active: 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( 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, 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") run_id = run.get("run_id")
if is_first_mention: if is_first_mention:
trace_message_ts = await post_slack_trace_reply(channel_id, thread_ts, thread_id) trace_message_ts = await post_slack_trace_reply(channel_id, thread_ts, thread_id)

View file

@ -739,6 +739,93 @@ def test_process_slack_mention_queues_active_thread_message(
assert "## Latest Mention Request\ninclude this screenshot" in queued_payload["text"] 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( def test_process_slack_mention_unmapped_user_blocked_and_prompted(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None: