mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-03 02:13:28 +00:00
Merge pull request #963 from langchain-ai/yogesh/multimodal-inputs
feat: add multimodal image support for Linear comments and queued run
This commit is contained in:
commit
a5bf5fe970
4 changed files with 300 additions and 17 deletions
|
|
@ -10,10 +10,13 @@ from __future__ import annotations
|
||||||
import logging
|
import logging
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
from langchain.agents.middleware import AgentState, before_model
|
from langchain.agents.middleware import AgentState, before_model
|
||||||
from langgraph.config import get_config, get_store
|
from langgraph.config import get_config, get_store
|
||||||
from langgraph.runtime import Runtime
|
from langgraph.runtime import Runtime
|
||||||
|
|
||||||
|
from ..utils.multimodal import fetch_image_block
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -23,6 +26,25 @@ class LinearNotifyState(AgentState):
|
||||||
linear_messages_sent_count: int
|
linear_messages_sent_count: int
|
||||||
|
|
||||||
|
|
||||||
|
async def _build_blocks_from_payload(
|
||||||
|
payload: dict[str, Any],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
text = payload.get("text", "")
|
||||||
|
image_urls = payload.get("image_urls", []) or []
|
||||||
|
blocks: list[dict[str, Any]] = []
|
||||||
|
if text:
|
||||||
|
blocks.append({"type": "text", "text": text})
|
||||||
|
|
||||||
|
if not image_urls:
|
||||||
|
return blocks
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
for image_url in image_urls:
|
||||||
|
image_block = await fetch_image_block(image_url, client)
|
||||||
|
if image_block:
|
||||||
|
blocks.append(image_block)
|
||||||
|
return blocks
|
||||||
|
|
||||||
|
|
||||||
@before_model(state_schema=LinearNotifyState)
|
@before_model(state_schema=LinearNotifyState)
|
||||||
async def check_message_queue_before_model( # noqa: PLR0911
|
async def check_message_queue_before_model( # noqa: PLR0911
|
||||||
state: LinearNotifyState, # noqa: ARG001
|
state: LinearNotifyState, # noqa: ARG001
|
||||||
|
|
@ -80,11 +102,21 @@ async def check_message_queue_before_model( # noqa: PLR0911
|
||||||
thread_id,
|
thread_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
content_blocks = [
|
content_blocks: list[dict[str, Any]] = []
|
||||||
{"type": "text", "text": msg.get("content", "")}
|
for msg in queued_messages:
|
||||||
for msg in queued_messages
|
content = msg.get("content")
|
||||||
if msg.get("content")
|
if isinstance(content, dict) and ("text" in content or "image_urls" in content):
|
||||||
]
|
logger.debug("Queued message contains text + image URLs")
|
||||||
|
blocks = await _build_blocks_from_payload(content)
|
||||||
|
content_blocks.extend(blocks)
|
||||||
|
continue
|
||||||
|
if isinstance(content, list):
|
||||||
|
logger.debug("Queued message contains %d content block(s)", len(content))
|
||||||
|
content_blocks.extend(content)
|
||||||
|
continue
|
||||||
|
if isinstance(content, str) and content:
|
||||||
|
logger.debug("Queued message contains text content")
|
||||||
|
content_blocks.append({"type": "text", "text": content})
|
||||||
|
|
||||||
if not content_blocks:
|
if not content_blocks:
|
||||||
return None
|
return None
|
||||||
|
|
|
||||||
83
apps/agent/agent/utils/multimodal.py
Normal file
83
apps/agent/agent/utils/multimodal.py
Normal file
|
|
@ -0,0 +1,83 @@
|
||||||
|
"""Utilities for building multimodal content blocks."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import base64
|
||||||
|
import logging
|
||||||
|
import mimetypes
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from langchain_core.messages.content import create_image_block
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
IMAGE_MARKDOWN_RE = re.compile(r"!\[[^\]]*\]\((https?://[^\s)]+)\)")
|
||||||
|
IMAGE_URL_RE = re.compile(
|
||||||
|
r"(https?://[^\s)]+\.(?:png|jpe?g|gif|webp|bmp|tiff)(?:\?[^\s)]+)?)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_image_urls(text: str) -> list[str]:
|
||||||
|
"""Extract image URLs from markdown image syntax and direct image links."""
|
||||||
|
if not text:
|
||||||
|
return []
|
||||||
|
|
||||||
|
urls: list[str] = []
|
||||||
|
urls.extend(IMAGE_MARKDOWN_RE.findall(text))
|
||||||
|
urls.extend(IMAGE_URL_RE.findall(text))
|
||||||
|
|
||||||
|
deduped = dedupe_urls(urls)
|
||||||
|
if deduped:
|
||||||
|
logger.debug("Extracted %d image URL(s)", len(deduped))
|
||||||
|
return deduped
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_image_block(
|
||||||
|
image_url: str,
|
||||||
|
client: httpx.AsyncClient,
|
||||||
|
) -> dict[str, Any] | None:
|
||||||
|
"""Fetch image bytes and build an image content block."""
|
||||||
|
try:
|
||||||
|
logger.debug("Fetching image from %s", image_url)
|
||||||
|
headers = None
|
||||||
|
if "uploads.linear.app" in image_url:
|
||||||
|
linear_api_key = os.environ.get("LINEAR_API_KEY", "")
|
||||||
|
if linear_api_key:
|
||||||
|
headers = {"Authorization": linear_api_key}
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"LINEAR_API_KEY not set; cannot authenticate image fetch for %s",
|
||||||
|
image_url,
|
||||||
|
)
|
||||||
|
response = await client.get(image_url, headers=headers)
|
||||||
|
response.raise_for_status()
|
||||||
|
content_type = response.headers.get("Content-Type", "").split(";")[0].strip()
|
||||||
|
if not content_type:
|
||||||
|
guessed, _ = mimetypes.guess_type(image_url)
|
||||||
|
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")
|
||||||
|
logger.info(
|
||||||
|
"Fetched image %s (%s, %d bytes)",
|
||||||
|
image_url,
|
||||||
|
content_type,
|
||||||
|
len(response.content),
|
||||||
|
)
|
||||||
|
return create_image_block(base64=encoded, mime_type=content_type)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to fetch image from %s", image_url)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def dedupe_urls(urls: list[str]) -> list[str]:
|
||||||
|
return list(dict.fromkeys(urls))
|
||||||
|
|
@ -11,10 +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 dedupe_urls, extract_image_urls, fetch_image_block
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -426,7 +428,9 @@ async def is_thread_active(thread_id: str) -> bool:
|
||||||
return status == "busy"
|
return status == "busy"
|
||||||
|
|
||||||
|
|
||||||
async def queue_message_for_thread(thread_id: str, message_content: str) -> bool:
|
async def queue_message_for_thread(
|
||||||
|
thread_id: str, message_content: str | list[dict[str, Any]] | dict[str, Any]
|
||||||
|
) -> bool:
|
||||||
"""Queue a message for a thread that is currently active.
|
"""Queue a message for a thread that is currently active.
|
||||||
|
|
||||||
Stores the message in the langgraph store, namespaced to the thread.
|
Stores the message in the langgraph store, namespaced to the thread.
|
||||||
|
|
@ -435,7 +439,7 @@ async def queue_message_for_thread(thread_id: str, message_content: str) -> bool
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
thread_id: The LangGraph thread ID
|
thread_id: The LangGraph thread ID
|
||||||
message_content: The message content to queue
|
message_content: The message content to queue (text or content blocks)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if successfully queued, False otherwise
|
True if successfully queued, False otherwise
|
||||||
|
|
@ -571,9 +575,19 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915
|
||||||
|
|
||||||
title = full_issue.get("title", "No title")
|
title = full_issue.get("title", "No title")
|
||||||
description = full_issue.get("description") or "No description"
|
description = full_issue.get("description") or "No description"
|
||||||
|
image_urls: list[str] = []
|
||||||
|
description_image_urls = extract_image_urls(description)
|
||||||
|
if description_image_urls:
|
||||||
|
image_urls.extend(description_image_urls)
|
||||||
|
logger.debug(
|
||||||
|
"Found %d image URL(s) in issue description",
|
||||||
|
len(description_image_urls),
|
||||||
|
)
|
||||||
|
|
||||||
comments = full_issue.get("comments", {}).get("nodes", [])
|
comments = full_issue.get("comments", {}).get("nodes", [])
|
||||||
comments_text = ""
|
comments_text = ""
|
||||||
|
triggering_comment = issue_data.get("triggering_comment", "")
|
||||||
|
triggering_comment_id = issue_data.get("triggering_comment_id", "")
|
||||||
|
|
||||||
bot_message_prefixes = (
|
bot_message_prefixes = (
|
||||||
"🔐 **GitHub Authentication Required**",
|
"🔐 **GitHub Authentication Required**",
|
||||||
|
|
@ -585,32 +599,75 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915
|
||||||
"❌ **Agent Error**",
|
"❌ **Agent Error**",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
comment_ids: set[str] = set()
|
||||||
|
comment_id_to_index: dict[str, int] = {}
|
||||||
if comments:
|
if comments:
|
||||||
last_bot_comment_idx = -1
|
last_bot_comment_idx = -1
|
||||||
for i, comment in enumerate(comments):
|
for i, comment in enumerate(comments):
|
||||||
|
comment_id = comment.get("id", "")
|
||||||
|
if comment_id:
|
||||||
|
comment_ids.add(comment_id)
|
||||||
|
comment_id_to_index[comment_id] = i
|
||||||
body = comment.get("body", "")
|
body = comment.get("body", "")
|
||||||
if any(body.startswith(prefix) for prefix in bot_message_prefixes):
|
if any(body.startswith(prefix) for prefix in bot_message_prefixes):
|
||||||
last_bot_comment_idx = i
|
last_bot_comment_idx = i
|
||||||
|
|
||||||
relevant_comments = []
|
relevant_comments = []
|
||||||
for i, comment in enumerate(comments):
|
trigger_index = None
|
||||||
if i <= last_bot_comment_idx:
|
if triggering_comment_id:
|
||||||
continue
|
trigger_index = comment_id_to_index.get(triggering_comment_id)
|
||||||
body = comment.get("body", "")
|
if trigger_index is not None:
|
||||||
if "@openswe" in body.lower():
|
relevant_comments = comments[trigger_index:]
|
||||||
relevant_comments.append(comment)
|
logger.debug(
|
||||||
relevant_comments.extend(comments[i + 1 :])
|
"Using triggering comment index %d to build relevant comments",
|
||||||
break
|
trigger_index,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
for i, comment in enumerate(comments):
|
||||||
|
if i <= last_bot_comment_idx:
|
||||||
|
continue
|
||||||
|
body = comment.get("body", "")
|
||||||
|
if "@openswe" in body.lower():
|
||||||
|
relevant_comments.append(comment)
|
||||||
|
relevant_comments.extend(comments[i + 1 :])
|
||||||
|
break
|
||||||
|
|
||||||
if relevant_comments:
|
if relevant_comments:
|
||||||
comments_text = "\n\n## Comments:\n"
|
comments_text = "\n\n## Comments:\n"
|
||||||
for comment in relevant_comments:
|
for comment in relevant_comments:
|
||||||
author = comment.get("user", {}).get("name", "Unknown")
|
author = comment.get("user", {}).get("name", "Unknown")
|
||||||
body = comment.get("body", "")
|
body = comment.get("body", "")
|
||||||
|
body_image_urls = extract_image_urls(body)
|
||||||
|
if body_image_urls:
|
||||||
|
image_urls.extend(body_image_urls)
|
||||||
|
logger.debug(
|
||||||
|
"Found %d image URL(s) in comment by %s",
|
||||||
|
len(body_image_urls),
|
||||||
|
author,
|
||||||
|
)
|
||||||
if any(body.startswith(prefix) for prefix in bot_message_prefixes):
|
if any(body.startswith(prefix) for prefix in bot_message_prefixes):
|
||||||
continue
|
continue
|
||||||
comments_text += f"\n**{author}:** {body}\n"
|
comments_text += f"\n**{author}:** {body}\n"
|
||||||
|
|
||||||
|
if triggering_comment and triggering_comment_id not in comment_ids:
|
||||||
|
if not comments_text:
|
||||||
|
comments_text = "\n\n## Comments:\n"
|
||||||
|
trigger_author = comment_author.get("name", "Unknown")
|
||||||
|
trigger_body = triggering_comment
|
||||||
|
trigger_image_urls = extract_image_urls(trigger_body)
|
||||||
|
if trigger_image_urls:
|
||||||
|
image_urls.extend(trigger_image_urls)
|
||||||
|
logger.debug(
|
||||||
|
"Found %d image URL(s) in triggering comment by %s",
|
||||||
|
len(trigger_image_urls),
|
||||||
|
trigger_author,
|
||||||
|
)
|
||||||
|
comments_text += f"\n**{trigger_author}:** {trigger_body}\n"
|
||||||
|
logger.debug(
|
||||||
|
"Appended triggering comment %s not present in issue comments list",
|
||||||
|
triggering_comment_id or "<missing-id>",
|
||||||
|
)
|
||||||
|
|
||||||
prompt = (
|
prompt = (
|
||||||
f"Please work on the following issue:\n\n"
|
f"Please work on the following issue:\n\n"
|
||||||
f"## Title: {title}\n\n"
|
f"## Title: {title}\n\n"
|
||||||
|
|
@ -619,6 +676,18 @@ 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]] = [create_text_block(prompt)]
|
||||||
|
if image_urls:
|
||||||
|
image_urls = dedupe_urls(image_urls)
|
||||||
|
logger.info("Preparing %d image(s) for multimodal content", len(image_urls))
|
||||||
|
logger.debug("Image URLs: %s", image_urls)
|
||||||
|
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
for image_url in image_urls:
|
||||||
|
image_block = await fetch_image_block(image_url, client)
|
||||||
|
if image_block:
|
||||||
|
content_blocks.append(image_block)
|
||||||
|
logger.info("Built %d content block(s) for prompt", len(content_blocks))
|
||||||
|
|
||||||
identifier = full_issue.get("identifier", "") or issue_data.get("identifier", "")
|
identifier = full_issue.get("identifier", "") or issue_data.get("identifier", "")
|
||||||
linear_project_id = ""
|
linear_project_id = ""
|
||||||
|
|
@ -652,9 +721,10 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915
|
||||||
thread_id,
|
thread_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
queued_payload = {"text": prompt, "image_urls": image_urls}
|
||||||
queued = await queue_message_for_thread(
|
queued = await queue_message_for_thread(
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
message_content=prompt,
|
message_content=queued_payload,
|
||||||
)
|
)
|
||||||
|
|
||||||
if queued:
|
if queued:
|
||||||
|
|
@ -669,7 +739,7 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915
|
||||||
await langgraph_client.runs.create(
|
await langgraph_client.runs.create(
|
||||||
thread_id,
|
thread_id,
|
||||||
"agent",
|
"agent",
|
||||||
input={"messages": [{"role": "user", "content": prompt}]},
|
input={"messages": [{"role": "user", "content": content_blocks}]},
|
||||||
config={"configurable": configurable},
|
config={"configurable": configurable},
|
||||||
if_not_exists="create",
|
if_not_exists="create",
|
||||||
)
|
)
|
||||||
|
|
|
||||||
98
apps/agent/tests/test_multimodal.py
Normal file
98
apps/agent/tests/test_multimodal.py
Normal file
|
|
@ -0,0 +1,98 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from agent.utils.multimodal import extract_image_urls
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_empty() -> None:
|
||||||
|
assert extract_image_urls("") == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_markdown_and_direct_dedupes() -> None:
|
||||||
|
text = (
|
||||||
|
"Here is an image  and another "
|
||||||
|
"![https://example.com/b.JPG?size=large plus a repeat https://example.com/a.png"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == [
|
||||||
|
"https://example.com/a.png",
|
||||||
|
"https://example.com/b.JPG?size=large",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_ignores_non_images() -> None:
|
||||||
|
text = "Not images: https://example.com/file.pdf and https://example.com/noext"
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_markdown_syntax() -> None:
|
||||||
|
text = "Check out this screenshot: "
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == ["https://example.com/screenshot.png"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_direct_links() -> None:
|
||||||
|
text = "Direct link: https://example.com/photo.jpg and another https://example.com/image.gif"
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == [
|
||||||
|
"https://example.com/photo.jpg",
|
||||||
|
"https://example.com/image.gif",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_various_formats() -> None:
|
||||||
|
text = (
|
||||||
|
"Multiple formats: "
|
||||||
|
"https://example.com/image.png "
|
||||||
|
"https://example.com/photo.jpeg "
|
||||||
|
"https://example.com/pic.gif "
|
||||||
|
"https://example.com/img.webp "
|
||||||
|
"https://example.com/bitmap.bmp "
|
||||||
|
"https://example.com/scan.tiff"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == [
|
||||||
|
"https://example.com/image.png",
|
||||||
|
"https://example.com/photo.jpeg",
|
||||||
|
"https://example.com/pic.gif",
|
||||||
|
"https://example.com/img.webp",
|
||||||
|
"https://example.com/bitmap.bmp",
|
||||||
|
"https://example.com/scan.tiff",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_with_query_params() -> None:
|
||||||
|
text = "Image with params: https://cdn.example.com/image.png?width=800&height=600"
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == ["https://cdn.example.com/image.png?width=800&height=600"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_case_insensitive() -> None:
|
||||||
|
text = "Mixed case: https://example.com/Image.PNG and https://example.com/photo.JpEg"
|
||||||
|
|
||||||
|
assert extract_image_urls(text) == [
|
||||||
|
"https://example.com/Image.PNG",
|
||||||
|
"https://example.com/photo.JpEg",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_deduplication() -> None:
|
||||||
|
text = "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"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_image_urls_mixed_markdown_and_direct() -> None:
|
||||||
|
text = (
|
||||||
|
"Markdown:  "
|
||||||
|
"and direct: https://example.com/direct.jpg "
|
||||||
|
"and another markdown "
|
||||||
|
)
|
||||||
|
|
||||||
|
result = extract_image_urls(text)
|
||||||
|
assert set(result) == {
|
||||||
|
"https://example.com/markdown.png",
|
||||||
|
"https://example.com/direct.jpg",
|
||||||
|
"https://example.com/another.gif",
|
||||||
|
}
|
||||||
|
assert len(result) == 3
|
||||||
Loading…
Add table
Reference in a new issue