mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
fix: persist selected repo (#1008)
* chore: Better err handling for extracting repos, always post reply * cr * fix: persist selected repo
This commit is contained in:
parent
163649197b
commit
d18318404b
2 changed files with 176 additions and 3 deletions
|
|
@ -5,6 +5,7 @@ import hmac
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
|
|
@ -12,6 +13,7 @@ import httpx
|
|||
from fastapi import BackgroundTasks, FastAPI, HTTPException, Request
|
||||
from langchain_core.messages.content import create_text_block
|
||||
from langgraph_sdk import get_client
|
||||
from langgraph_sdk.client import LangGraphClient
|
||||
|
||||
from .utils.comments import get_recent_comments
|
||||
from .utils.multimodal import dedupe_urls, extract_image_urls, fetch_image_block
|
||||
|
|
@ -242,17 +244,92 @@ def generate_thread_id_from_slack_thread(channel_id: str, thread_id: str) -> str
|
|||
return str(uuid.UUID(hex=md5_hex))
|
||||
|
||||
|
||||
def _extract_repo_config_from_thread(thread: dict[str, Any]) -> dict[str, str] | None:
|
||||
"""Extract repo config from persisted thread data."""
|
||||
metadata = thread.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
return None
|
||||
|
||||
repo = metadata.get("repo")
|
||||
if isinstance(repo, dict):
|
||||
owner = repo.get("owner")
|
||||
name = repo.get("name")
|
||||
if isinstance(owner, str) and owner and isinstance(name, str) and name:
|
||||
return {"owner": owner, "name": name}
|
||||
|
||||
owner = metadata.get("repo_owner")
|
||||
name = metadata.get("repo_name")
|
||||
if isinstance(owner, str) and owner and isinstance(name, str) and name:
|
||||
return {"owner": owner, "name": name}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _is_not_found_error(exc: Exception) -> bool:
|
||||
"""Best-effort check for LangGraph 404 errors."""
|
||||
return getattr(exc, "status_code", None) == 404
|
||||
|
||||
|
||||
async def _upsert_slack_thread_repo_metadata(
|
||||
thread_id: str, repo_config: dict[str, str], langgraph_client: LangGraphClient
|
||||
) -> None:
|
||||
"""Persist the selected repo config on the thread metadata."""
|
||||
try:
|
||||
await langgraph_client.threads.update(thread_id=thread_id, metadata={"repo": repo_config})
|
||||
except Exception as exc: # noqa: BLE001
|
||||
if _is_not_found_error(exc):
|
||||
try:
|
||||
await langgraph_client.threads.create(
|
||||
thread_id=thread_id,
|
||||
if_exists="do_nothing",
|
||||
metadata={"repo": repo_config},
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"Failed to create Slack thread %s while persisting repo metadata",
|
||||
thread_id,
|
||||
)
|
||||
return
|
||||
logger.exception(
|
||||
"Failed to persist Slack thread repo metadata for thread %s",
|
||||
thread_id,
|
||||
)
|
||||
|
||||
|
||||
async def get_slack_repo_config(message: str, channel_id: str, thread_ts: str) -> dict[str, str]:
|
||||
"""Resolve repository configuration for Slack-triggered runs."""
|
||||
owner = SLACK_REPO_OWNER.strip() or "langchain-ai"
|
||||
name = SLACK_REPO_NAME.strip() or "langchainplus"
|
||||
thread_id = generate_thread_id_from_slack_thread(channel_id, thread_ts)
|
||||
langgraph_client = get_client(url=LANGGRAPH_URL)
|
||||
|
||||
if "repo:" in message:
|
||||
import re
|
||||
should_parse_repo_from_message = False
|
||||
try:
|
||||
thread = await langgraph_client.threads.get(thread_id)
|
||||
thread_repo_config = _extract_repo_config_from_thread(thread)
|
||||
if thread_repo_config:
|
||||
owner = thread_repo_config["owner"]
|
||||
name = thread_repo_config["name"]
|
||||
else:
|
||||
logger.warning(
|
||||
"Slack thread %s exists but no repo metadata found; using default repo %s/%s",
|
||||
thread_id,
|
||||
owner,
|
||||
name,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
if _is_not_found_error(exc):
|
||||
should_parse_repo_from_message = True
|
||||
else:
|
||||
logger.exception(
|
||||
"Failed to fetch Slack thread %s for repo resolution; using default repo",
|
||||
thread_id,
|
||||
)
|
||||
|
||||
if should_parse_repo_from_message and "repo:" in message:
|
||||
match = re.search(r"repo:([^ ]+)", message)
|
||||
if match:
|
||||
repo = match.group(1)
|
||||
repo = match.group(1).strip()
|
||||
if "/" in repo:
|
||||
owner, name = repo.split("/", 1)
|
||||
|
||||
|
|
@ -659,6 +736,7 @@ async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[st
|
|||
}
|
||||
|
||||
langgraph_client = get_client(url=LANGGRAPH_URL)
|
||||
await _upsert_slack_thread_repo_metadata(thread_id, repo_config, langgraph_client)
|
||||
await langgraph_client.runs.create(
|
||||
thread_id,
|
||||
"agent",
|
||||
|
|
|
|||
|
|
@ -1,3 +1,8 @@
|
|||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import webapp
|
||||
from agent.utils.slack import (
|
||||
format_slack_messages_for_prompt,
|
||||
replace_bot_mention_with_username,
|
||||
|
|
@ -7,6 +12,30 @@ from agent.utils.slack import (
|
|||
from agent.webapp import generate_thread_id_from_slack_thread
|
||||
|
||||
|
||||
class _FakeNotFoundError(Exception):
|
||||
status_code = 404
|
||||
|
||||
|
||||
class _FakeThreadsClient:
|
||||
def __init__(self, thread: dict | None = None, raise_not_found: bool = False) -> None:
|
||||
self.thread = thread
|
||||
self.raise_not_found = raise_not_found
|
||||
self.requested_thread_id: str | None = None
|
||||
|
||||
async def get(self, thread_id: str) -> dict:
|
||||
self.requested_thread_id = thread_id
|
||||
if self.raise_not_found:
|
||||
raise _FakeNotFoundError("not found")
|
||||
if self.thread is None:
|
||||
raise AssertionError("thread must be provided when raise_not_found is False")
|
||||
return self.thread
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, threads_client: _FakeThreadsClient) -> None:
|
||||
self.threads = threads_client
|
||||
|
||||
|
||||
def test_generate_thread_id_from_slack_thread_is_deterministic() -> None:
|
||||
channel_id = "C12345"
|
||||
thread_ts = "1730900000.123456"
|
||||
|
|
@ -112,3 +141,69 @@ def test_select_slack_context_messages_detects_username_mention() -> None:
|
|||
|
||||
assert mode == "last_mention"
|
||||
assert [item["ts"] for item in selected] == ["1.0", "2.0", "3.0"]
|
||||
|
||||
|
||||
def test_get_slack_repo_config_uses_existing_thread_repo(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, str] = {}
|
||||
threads_client = _FakeThreadsClient(
|
||||
thread={"metadata": {"repo": {"owner": "saved-owner", "name": "saved-repo"}}}
|
||||
)
|
||||
|
||||
async def fake_post_slack_thread_reply(channel_id: str, thread_ts: str, text: str) -> bool:
|
||||
captured["channel_id"] = channel_id
|
||||
captured["thread_ts"] = thread_ts
|
||||
captured["text"] = text
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeClient(threads_client))
|
||||
monkeypatch.setattr(webapp, "post_slack_thread_reply", fake_post_slack_thread_reply)
|
||||
|
||||
repo = asyncio.run(
|
||||
webapp.get_slack_repo_config("please use repo:new-owner/new-repo", "C123", "1.234")
|
||||
)
|
||||
|
||||
assert repo == {"owner": "saved-owner", "name": "saved-repo"}
|
||||
assert threads_client.requested_thread_id == generate_thread_id_from_slack_thread(
|
||||
"C123", "1.234"
|
||||
)
|
||||
assert captured["text"] == "Using repository: `saved-owner/saved-repo`"
|
||||
|
||||
|
||||
def test_get_slack_repo_config_parses_message_for_new_thread(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
threads_client = _FakeThreadsClient(raise_not_found=True)
|
||||
|
||||
async def fake_post_slack_thread_reply(channel_id: str, thread_ts: str, text: str) -> bool:
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeClient(threads_client))
|
||||
monkeypatch.setattr(webapp, "post_slack_thread_reply", fake_post_slack_thread_reply)
|
||||
|
||||
repo = asyncio.run(
|
||||
webapp.get_slack_repo_config("please use repo:new-owner/new-repo", "C123", "1.234")
|
||||
)
|
||||
|
||||
assert repo == {"owner": "new-owner", "name": "new-repo"}
|
||||
|
||||
|
||||
def test_get_slack_repo_config_existing_thread_without_repo_uses_default(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
threads_client = _FakeThreadsClient(thread={"metadata": {}})
|
||||
monkeypatch.setattr(webapp, "SLACK_REPO_OWNER", "default-owner")
|
||||
monkeypatch.setattr(webapp, "SLACK_REPO_NAME", "default-repo")
|
||||
|
||||
async def fake_post_slack_thread_reply(channel_id: str, thread_ts: str, text: str) -> bool:
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeClient(threads_client))
|
||||
monkeypatch.setattr(webapp, "post_slack_thread_reply", fake_post_slack_thread_reply)
|
||||
|
||||
repo = asyncio.run(
|
||||
webapp.get_slack_repo_config("please use repo:new-owner/new-repo", "C123", "1.234")
|
||||
)
|
||||
|
||||
assert repo == {"owner": "default-owner", "name": "default-repo"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue