refactor: use content block helpers and shared url dedupe; run format and lint

This commit is contained in:
aran-yogesh 2026-02-24 12:11:24 -08:00
parent fb23ef6c15
commit f82408cd61
9 changed files with 41 additions and 93 deletions

View file

@ -9,7 +9,7 @@ import contextlib
import os import os
import time import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any from typing import Any
from deepagents.backends.protocol import ( from deepagents.backends.protocol import (
ExecuteResponse, ExecuteResponse,
@ -19,7 +19,6 @@ from deepagents.backends.protocol import (
WriteResult, WriteResult,
) )
from deepagents.backends.sandbox import BaseSandbox from deepagents.backends.sandbox import BaseSandbox
from langsmith.sandbox import Sandbox, SandboxClient, SandboxTemplate from langsmith.sandbox import Sandbox, SandboxClient, SandboxTemplate
@ -148,9 +147,7 @@ class LangSmithBackend(BaseSandbox):
responses: list[FileDownloadResponse] = [] responses: list[FileDownloadResponse] = []
for path in paths: for path in paths:
content = self._sandbox.read(path) content = self._sandbox.read(path)
responses.append( responses.append(FileDownloadResponse(path=path, content=content, error=None))
FileDownloadResponse(path=path, content=content, error=None)
)
return responses return responses
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]: def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
@ -209,10 +206,7 @@ class LangSmithProvider(SandboxProvider):
template_name=resolved_template_name, timeout=timeout template_name=resolved_template_name, timeout=timeout
) )
except Exception as e: except Exception as e:
msg = ( msg = f"Failed to create sandbox from template '{resolved_template_name}': {e}"
f"Failed to create sandbox from template "
f"'{resolved_template_name}': {e}"
)
raise RuntimeError(msg) from e raise RuntimeError(msg) from e
# Verify sandbox is ready by polling # Verify sandbox is ready by polling

View file

@ -8,7 +8,6 @@ human messages before the next model call.
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
from typing import Any from typing import Any
import httpx import httpx
@ -38,12 +37,9 @@ async def _build_blocks_from_payload(
if not image_urls: if not image_urls:
return blocks return blocks
linear_api_key = os.environ.get("LINEAR_API_KEY", "")
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
for image_url in image_urls: for image_url in image_urls:
image_block = await fetch_image_block( image_block = await fetch_image_block(image_url, client)
image_url, client, linear_api_key=linear_api_key
)
if image_block: if image_block:
blocks.append(image_block) blocks.append(image_block)
return blocks return blocks
@ -109,17 +105,13 @@ async def check_message_queue_before_model( # noqa: PLR0911
content_blocks: list[dict[str, Any]] = [] content_blocks: list[dict[str, Any]] = []
for msg in queued_messages: for msg in queued_messages:
content = msg.get("content") content = msg.get("content")
if isinstance(content, dict) and ( if isinstance(content, dict) and ("text" in content or "image_urls" in content):
"text" in content or "image_urls" in content
):
logger.debug("Queued message contains text + image URLs") logger.debug("Queued message contains text + image URLs")
blocks = await _build_blocks_from_payload(content) blocks = await _build_blocks_from_payload(content)
content_blocks.extend(blocks) content_blocks.extend(blocks)
continue continue
if isinstance(content, list): if isinstance(content, list):
logger.debug( logger.debug("Queued message contains %d content block(s)", len(content))
"Queued message contains %d content block(s)", len(content)
)
content_blocks.extend(content) content_blocks.extend(content)
continue continue
if isinstance(content, str) and content: if isinstance(content, str) and content:

View file

@ -181,16 +181,12 @@ I've {action} pull request to address this issue:
logger.info("Changes detected, preparing PR for thread %s", thread_id) logger.info("Changes detected, preparing PR for thread %s", thread_id)
current_branch = await asyncio.to_thread( current_branch = await asyncio.to_thread(git_current_branch, sandbox_backend, repo_dir)
git_current_branch, sandbox_backend, repo_dir
)
target_branch = f"open-swe/{thread_id}" target_branch = f"open-swe/{thread_id}"
if current_branch != target_branch: if current_branch != target_branch:
await asyncio.to_thread( await asyncio.to_thread(git_checkout_branch, sandbox_backend, repo_dir, target_branch)
git_checkout_branch, sandbox_backend, repo_dir, target_branch
)
await asyncio.to_thread( await asyncio.to_thread(
git_config_user, git_config_user,

View file

@ -4,7 +4,6 @@
# Suppress deprecation warnings from langchain_core (e.g., Pydantic V1 on Python 3.14+) # Suppress deprecation warnings from langchain_core (e.g., Pydantic V1 on Python 3.14+)
# ruff: noqa: E402 # ruff: noqa: E402
import logging import logging
import os
import warnings import warnings
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -29,6 +28,7 @@ from deepagents.backends.protocol import SandboxBackendProtocol
from langchain_anthropic import ChatAnthropic from langchain_anthropic import ChatAnthropic
from .encryption import decrypt_token from .encryption import decrypt_token
from .integrations.langsmith import _create_langsmith_sandbox
from .middleware import ( from .middleware import (
ToolErrorMiddleware, ToolErrorMiddleware,
check_message_queue_before_model, check_message_queue_before_model,
@ -37,8 +37,6 @@ from .middleware import (
) )
from .prompt import construct_system_prompt from .prompt import construct_system_prompt
from .tools import commit_and_open_pr, fetch_url, http_request from .tools import commit_and_open_pr, fetch_url, http_request
from .integrations.langsmith import _create_langsmith_sandbox
client = get_client() client = get_client()
@ -46,12 +44,12 @@ SANDBOX_CREATING = "__creating__"
SANDBOX_CREATION_TIMEOUT = 180 SANDBOX_CREATION_TIMEOUT = 180
SANDBOX_POLL_INTERVAL = 1.0 SANDBOX_POLL_INTERVAL = 1.0
from .utils.sandbox_state import SANDBOX_BACKENDS, get_sandbox_id_from_metadata
from .utils.github import ( from .utils.github import (
git_has_uncommitted_changes, git_has_uncommitted_changes,
is_valid_git_repo, is_valid_git_repo,
remove_directory, remove_directory,
) )
from .utils.sandbox_state import SANDBOX_BACKENDS, get_sandbox_id_from_metadata
async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915 async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
@ -87,9 +85,7 @@ async def _clone_or_pull_repo_in_sandbox( # noqa: PLR0915
is_git_repo = await loop.run_in_executor(None, is_valid_git_repo, sandbox_backend, repo_dir) is_git_repo = await loop.run_in_executor(None, is_valid_git_repo, sandbox_backend, repo_dir)
if not is_git_repo: if not is_git_repo:
logger.warning( logger.warning("Repo directory missing or not a valid git repo at %s, removing", repo_dir)
"Repo directory missing or not a valid git repo at %s, removing", repo_dir
)
try: try:
removed = await loop.run_in_executor(None, remove_directory, sandbox_backend, repo_dir) removed = await loop.run_in_executor(None, remove_directory, sandbox_backend, repo_dir)
if not removed: if not removed:

View file

@ -181,9 +181,7 @@ def commit_and_open_pr(
"pr_url": None, "pr_url": None,
} }
base_branch = asyncio.run( base_branch = asyncio.run(get_github_default_branch(repo_owner, repo_name, github_token))
get_github_default_branch(repo_owner, repo_name, github_token)
)
pr_url, _pr_number, pr_existing = asyncio.run( pr_url, _pr_number, pr_existing = asyncio.run(
create_github_pr( create_github_pr(
repo_owner=repo_owner, repo_owner=repo_owner,

View file

@ -37,24 +37,18 @@ def remove_directory(sandbox_backend: SandboxBackendProtocol, repo_dir: str) ->
return result.exit_code == 0 return result.exit_code == 0
def git_has_uncommitted_changes( def git_has_uncommitted_changes(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> bool:
sandbox_backend: SandboxBackendProtocol, repo_dir: str
) -> bool:
"""Check whether the repo has uncommitted changes.""" """Check whether the repo has uncommitted changes."""
result = _run_git(sandbox_backend, repo_dir, "git status --porcelain") result = _run_git(sandbox_backend, repo_dir, "git status --porcelain")
return result.exit_code == 0 and bool(result.output.strip()) return result.exit_code == 0 and bool(result.output.strip())
def git_fetch_origin( def git_fetch_origin(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> ExecuteResponse:
sandbox_backend: SandboxBackendProtocol, repo_dir: str
) -> ExecuteResponse:
"""Fetch latest from origin (best-effort).""" """Fetch latest from origin (best-effort)."""
return _run_git(sandbox_backend, repo_dir, "git fetch origin 2>/dev/null || true") return _run_git(sandbox_backend, repo_dir, "git fetch origin 2>/dev/null || true")
def git_has_unpushed_commits( def git_has_unpushed_commits(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> bool:
sandbox_backend: SandboxBackendProtocol, repo_dir: str
) -> bool:
"""Check whether there are commits not pushed to upstream.""" """Check whether there are commits not pushed to upstream."""
git_log_cmd = ( git_log_cmd = (
"git log --oneline @{upstream}..HEAD 2>/dev/null " "git log --oneline @{upstream}..HEAD 2>/dev/null "
@ -75,9 +69,7 @@ def git_checkout_branch(
) -> bool: ) -> bool:
"""Checkout branch, creating it if needed.""" """Checkout branch, creating it if needed."""
safe_branch = shlex.quote(branch) safe_branch = shlex.quote(branch)
checkout_result = _run_git( checkout_result = _run_git(sandbox_backend, repo_dir, f"git checkout -b {safe_branch}")
sandbox_backend, repo_dir, f"git checkout -b {safe_branch}"
)
if checkout_result.exit_code == 0: if checkout_result.exit_code == 0:
return True return True
fallback = _run_git(sandbox_backend, repo_dir, f"git checkout {safe_branch}") fallback = _run_git(sandbox_backend, repo_dir, f"git checkout {safe_branch}")
@ -97,9 +89,7 @@ def git_config_user(
_run_git(sandbox_backend, repo_dir, f"git config user.email {safe_email}") _run_git(sandbox_backend, repo_dir, f"git config user.email {safe_email}")
def git_add_all( def git_add_all(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> ExecuteResponse:
sandbox_backend: SandboxBackendProtocol, repo_dir: str
) -> ExecuteResponse:
"""Stage all changes.""" """Stage all changes."""
return _run_git(sandbox_backend, repo_dir, "git add -A") return _run_git(sandbox_backend, repo_dir, "git add -A")
@ -112,9 +102,7 @@ def git_commit(
return _run_git(sandbox_backend, repo_dir, f"git commit -m {safe_message}") return _run_git(sandbox_backend, repo_dir, f"git commit -m {safe_message}")
def git_get_remote_url( def git_get_remote_url(sandbox_backend: SandboxBackendProtocol, repo_dir: str) -> str | None:
sandbox_backend: SandboxBackendProtocol, repo_dir: str
) -> str | None:
"""Get the origin remote URL.""" """Get the origin remote URL."""
result = _run_git(sandbox_backend, repo_dir, "git remote get-url origin") result = _run_git(sandbox_backend, repo_dir, "git remote get-url origin")
if result.exit_code != 0: if result.exit_code != 0:

View file

@ -30,26 +30,22 @@ def extract_image_urls(text: str) -> list[str]:
urls.extend(IMAGE_MARKDOWN_RE.findall(text)) urls.extend(IMAGE_MARKDOWN_RE.findall(text))
urls.extend(IMAGE_URL_RE.findall(text)) urls.extend(IMAGE_URL_RE.findall(text))
deduped = _dedupe_urls(urls) deduped = dedupe_urls(urls)
if deduped: if deduped:
logger.debug("Extracted %d image URL(s)", len(deduped)) logger.debug("Extracted %d image URL(s)", len(deduped))
return deduped return deduped
async def fetch_image_block( async def fetch_image_block(
image_url: str, image_url: str,
client: httpx.AsyncClient, client: httpx.AsyncClient,
*,
linear_api_key: str | None = None,
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""Fetch image bytes and build an image content block.""" """Fetch image bytes and build an image content block."""
try: try:
logger.debug("Fetching image from %s", image_url) logger.debug("Fetching image from %s", image_url)
headers = None headers = None
if "uploads.linear.app" in image_url: if "uploads.linear.app" in image_url:
if linear_api_key is None: linear_api_key = os.environ.get("LINEAR_API_KEY", "")
linear_api_key = os.environ.get("LINEAR_API_KEY", "")
if linear_api_key: if linear_api_key:
headers = {"Authorization": linear_api_key} headers = {"Authorization": linear_api_key}
else: else:
@ -62,7 +58,13 @@ async def fetch_image_block(
content_type = response.headers.get("Content-Type", "").split(";")[0].strip() content_type = response.headers.get("Content-Type", "").split(";")[0].strip()
if not content_type: if not content_type:
guessed, _ = mimetypes.guess_type(image_url) guessed, _ = mimetypes.guess_type(image_url)
content_type = guessed or "application/octet-stream" if not guessed:
logger.warning(
"Could not determine content type for %s; skipping image",
image_url,
)
return None
content_type = guessed
encoded = base64.b64encode(response.content).decode("ascii") encoded = base64.b64encode(response.content).decode("ascii")
logger.info( logger.info(
@ -77,12 +79,6 @@ async def fetch_image_block(
return None return None
def _dedupe_urls(urls: list[str]) -> list[str]: def dedupe_urls(urls: list[str]) -> list[str]:
seen: set[str] = set() deduped: set[str] = set(urls)
deduped: list[str] = [] return list(deduped)
for url in urls:
if url in seen:
continue
seen.add(url)
deduped.append(url)
return deduped

View file

@ -11,11 +11,12 @@ from typing import Any
import httpx import httpx
import jwt import jwt
from fastapi import BackgroundTasks, FastAPI, HTTPException, Request 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 import get_client
# Local import for encryption # Local import for encryption
from .encryption import encrypt_token from .encryption import encrypt_token
from .utils.multimodal import extract_image_urls, fetch_image_block from .utils.multimodal import dedupe_urls, extract_image_urls, fetch_image_block
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -37,7 +38,6 @@ LINEAR_API_KEY = os.environ.get("LINEAR_API_KEY", "")
X_SERVICE_AUTH_JWT_SECRET = os.environ.get("X_SERVICE_AUTH_JWT_SECRET", "") X_SERVICE_AUTH_JWT_SECRET = os.environ.get("X_SERVICE_AUTH_JWT_SECRET", "")
def get_service_jwt_token_for_user( def get_service_jwt_token_for_user(
user_id: str, tenant_id: str, expiration_seconds: int = 300 user_id: str, tenant_id: str, expiration_seconds: int = 300
) -> str: ) -> str:
@ -78,9 +78,11 @@ LINEAR_TEAM_TO_REPO: dict[str, dict[str, Any] | dict[str, str]] = {
"open-swe-v3-test": {"owner": "aran-yogesh", "name": "nimedge"}, "open-swe-v3-test": {"owner": "aran-yogesh", "name": "nimedge"},
"open-swe-dev-test": {"owner": "aran-yogesh", "name": "TalkBack"}, "open-swe-dev-test": {"owner": "aran-yogesh", "name": "TalkBack"},
}, },
"default": {"owner": "aran-yogesh", "name": "TalkBack"} # Fallback for issues without project "default": {
"owner": "aran-yogesh",
"name": "TalkBack",
}, # Fallback for issues without project
}, },
"LangChain OSS": { "LangChain OSS": {
"projects": { "projects": {
"deepagents": {"owner": "langchain-ai", "name": "deepagents"}, "deepagents": {"owner": "langchain-ai", "name": "deepagents"},
@ -93,9 +95,7 @@ LINEAR_TEAM_TO_REPO: dict[str, dict[str, Any] | dict[str, str]] = {
}, },
"default": {"owner": "langchain-ai", "name": "ai-sdr"}, "default": {"owner": "langchain-ai", "name": "ai-sdr"},
}, },
"Docs": { "Docs": {"default": {"owner": "langchain-ai", "name": "docs"}},
"default": {"owner": "langchain-ai", "name": "docs"}
},
} }
@ -676,24 +676,15 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915
"Please analyze this issue and implement the necessary changes. " "Please analyze this issue and implement the necessary changes. "
"When you're done, commit and push your changes." "When you're done, commit and push your changes."
) )
content_blocks: list[dict[str, Any]] = [{"type": "text", "text": prompt}] content_blocks: list[dict[str, Any]] = [create_text_block(prompt)]
if image_urls: if image_urls:
seen_urls: set[str] = set() image_urls = dedupe_urls(image_urls)
deduped_urls: list[str] = []
for url in image_urls:
if url in seen_urls:
continue
seen_urls.add(url)
deduped_urls.append(url)
image_urls = deduped_urls
logger.info("Preparing %d image(s) for multimodal content", len(image_urls)) logger.info("Preparing %d image(s) for multimodal content", len(image_urls))
logger.debug("Image URLs: %s", image_urls) logger.debug("Image URLs: %s", image_urls)
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
for image_url in image_urls: for image_url in image_urls:
image_block = await fetch_image_block( image_block = await fetch_image_block(image_url, client)
image_url, client, linear_api_key=LINEAR_API_KEY
)
if image_block: if image_block:
content_blocks.append(image_block) content_blocks.append(image_block)
logger.info("Built %d content block(s) for prompt", len(content_blocks)) logger.info("Built %d content block(s) for prompt", len(content_blocks))

View file

@ -77,10 +77,7 @@ def test_extract_image_urls_case_insensitive() -> None:
def test_extract_image_urls_deduplication() -> None: def test_extract_image_urls_deduplication() -> None:
text = ( text = "Same URL twice: https://example.com/image.png and again https://example.com/image.png"
"Same URL twice: https://example.com/image.png "
"and again https://example.com/image.png"
)
assert extract_image_urls(text) == ["https://example.com/image.png"] assert extract_image_urls(text) == ["https://example.com/image.png"]