diff --git a/apps/agent/agent/middleware/check_message_queue.py b/apps/agent/agent/middleware/check_message_queue.py index b85248fd..a8825761 100644 --- a/apps/agent/agent/middleware/check_message_queue.py +++ b/apps/agent/agent/middleware/check_message_queue.py @@ -10,10 +10,13 @@ from __future__ import annotations import logging from typing import Any +import httpx from langchain.agents.middleware import AgentState, before_model from langgraph.config import get_config, get_store from langgraph.runtime import Runtime +from ..utils.multimodal import fetch_image_block + logger = logging.getLogger(__name__) @@ -23,6 +26,25 @@ class LinearNotifyState(AgentState): 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) async def check_message_queue_before_model( # noqa: PLR0911 state: LinearNotifyState, # noqa: ARG001 @@ -80,11 +102,21 @@ async def check_message_queue_before_model( # noqa: PLR0911 thread_id, ) - content_blocks = [ - {"type": "text", "text": msg.get("content", "")} - for msg in queued_messages - if msg.get("content") - ] + content_blocks: list[dict[str, Any]] = [] + for msg in queued_messages: + content = 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: return None diff --git a/apps/agent/agent/utils/multimodal.py b/apps/agent/agent/utils/multimodal.py new file mode 100644 index 00000000..bc2b7ff6 --- /dev/null +++ b/apps/agent/agent/utils/multimodal.py @@ -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)) diff --git a/apps/agent/agent/webapp.py b/apps/agent/agent/webapp.py index 93296acb..1949b0e6 100644 --- a/apps/agent/agent/webapp.py +++ b/apps/agent/agent/webapp.py @@ -11,10 +11,12 @@ from typing import Any import httpx import jwt from fastapi import BackgroundTasks, FastAPI, HTTPException, Request +from langchain_core.messages.content import create_text_block from langgraph_sdk import get_client # Local import for encryption from .encryption import encrypt_token +from .utils.multimodal import dedupe_urls, extract_image_urls, fetch_image_block logger = logging.getLogger(__name__) @@ -426,7 +428,9 @@ async def is_thread_active(thread_id: str) -> bool: 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. 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: 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: 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") 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_text = "" + triggering_comment = issue_data.get("triggering_comment", "") + triggering_comment_id = issue_data.get("triggering_comment_id", "") bot_message_prefixes = ( "🔐 **GitHub Authentication Required**", @@ -585,32 +599,75 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915 "❌ **Agent Error**", ) + comment_ids: set[str] = set() + comment_id_to_index: dict[str, int] = {} if comments: last_bot_comment_idx = -1 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", "") if any(body.startswith(prefix) for prefix in bot_message_prefixes): last_bot_comment_idx = i relevant_comments = [] - 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 + trigger_index = None + if triggering_comment_id: + trigger_index = comment_id_to_index.get(triggering_comment_id) + if trigger_index is not None: + relevant_comments = comments[trigger_index:] + logger.debug( + "Using triggering comment index %d to build relevant comments", + 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: comments_text = "\n\n## Comments:\n" for comment in relevant_comments: author = comment.get("user", {}).get("name", "Unknown") 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): continue 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 "", + ) + prompt = ( f"Please work on the following issue:\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. " "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", "") linear_project_id = "" @@ -652,9 +721,10 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915 thread_id, ) + queued_payload = {"text": prompt, "image_urls": image_urls} queued = await queue_message_for_thread( thread_id=thread_id, - message_content=prompt, + message_content=queued_payload, ) if queued: @@ -669,7 +739,7 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915 await langgraph_client.runs.create( thread_id, "agent", - input={"messages": [{"role": "user", "content": prompt}]}, + input={"messages": [{"role": "user", "content": content_blocks}]}, config={"configurable": configurable}, if_not_exists="create", ) diff --git a/apps/agent/tests/test_multimodal.py b/apps/agent/tests/test_multimodal.py new file mode 100644 index 00000000..5dca4d39 --- /dev/null +++ b/apps/agent/tests/test_multimodal.py @@ -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 ![alt](https://example.com/a.png) 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: ![Screenshot](https://example.com/screenshot.png)" + + 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: ![alt text](https://example.com/markdown.png) " + "and direct: https://example.com/direct.jpg " + "and another markdown ![](https://example.com/another.gif)" + ) + + 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