diff --git a/.capy-cache/dependencies.sha256 b/.capy-cache/dependencies.sha256 new file mode 100644 index 000000000000..47f290b5a253 --- /dev/null +++ b/.capy-cache/dependencies.sha256 @@ -0,0 +1 @@ +8620910f03b78c62ebc5462fc6ef8351843f0a2525ee8213a01b0cfc28ba250b diff --git a/.capy-cache/docker-compose.capy.yml b/.capy-cache/docker-compose.capy.yml new file mode 100644 index 000000000000..268109a66294 --- /dev/null +++ b/.capy-cache/docker-compose.capy.yml @@ -0,0 +1,19 @@ +services: + clickhouse: + mem_limit: !reset null + cpus: !reset null + proxy: + extra_hosts: !override + - 'plugins:host-gateway' + - 'capture:host-gateway' + - 'capture-ai:host-gateway' + - 'capture-logs:host-gateway' + - 'replay-capture:host-gateway' + - 'feature-flags:host-gateway' + - 'hypercache-server:host-gateway' + web: + image: caddy:latest + entrypoint: socat + command: ['TCP-LISTEN:8000,fork,reuseaddr', 'TCP:host.docker.internal:8011'] + extra_hosts: + - 'host.docker.internal:host-gateway' diff --git a/.semgrep/rules/idor-team-scoped-models.yaml b/.semgrep/rules/idor-team-scoped-models.yaml index 3d4ae559159f..433aa975f682 100644 --- a/.semgrep/rules/idor-team-scoped-models.yaml +++ b/.semgrep/rules/idor-team-scoped-models.yaml @@ -84,6 +84,7 @@ rules: |Cohort |Comment |Conversation + |ConversationAttachment |ConversationRestoreToken |CoreMemory |CustomerJourney @@ -339,6 +340,7 @@ rules: |Cohort |Comment |Conversation + |ConversationAttachment |ConversationRestoreToken |CoreMemory |CustomerJourney diff --git a/ee/api/conversation.py b/ee/api/conversation.py index c76038e993f3..c932a400216e 100644 --- a/ee/api/conversation.py +++ b/ee/api/conversation.py @@ -6,7 +6,7 @@ from django.conf import settings from django.core.exceptions import ValidationError -from django.http import StreamingHttpResponse +from django.http import HttpResponse, StreamingHttpResponse from django.utils import timezone import pydantic @@ -20,6 +20,7 @@ from rest_framework.decorators import action from rest_framework.exceptions import Throttled from rest_framework.mixins import DestroyModelMixin, ListModelMixin, RetrieveModelMixin +from rest_framework.parsers import FormParser, MultiPartParser from rest_framework.request import Request from rest_framework.response import Response from rest_framework.viewsets import GenericViewSet @@ -37,6 +38,8 @@ AISustainedRateThrottle, is_team_exempt_from_ai_rate_limit, ) +from posthog.storage import object_storage +from posthog.storage.object_storage import ObjectStorageError from posthog.temporal.ai.chat_agent import ( CHAT_AGENT_STREAM_MAX_LENGTH, CHAT_AGENT_WORKFLOW_TIMEOUT, @@ -50,7 +53,14 @@ ResearchAgentWorkflowInputs, ) -from products.posthog_ai.backend.models.assistant import Conversation +from products.posthog_ai.backend.attachments import ( + MAX_ATTACHMENTS_PER_MESSAGE, + InvalidConversationAttachment, + attachment_to_ref, + process_conversation_attachment, + save_conversation_attachment, +) +from products.posthog_ai.backend.models.assistant import Conversation, ConversationAttachment from ee.billing.quota_limiting import QuotaLimitingCaches, QuotaResource, is_team_limited from ee.hogai.api.serializers import ConversationMinimalSerializer, ConversationSerializer @@ -95,6 +105,7 @@ class MessageSerializer(MessageMinimalSerializer): content = serializers.CharField( required=True, allow_null=True, # Null content means we're resuming streaming or continuing previous generation + allow_blank=True, max_length=40000, # Roughly 10k tokens ) conversation = serializers.UUIDField( @@ -108,14 +119,51 @@ class MessageSerializer(MessageMinimalSerializer): agent_mode = serializers.ChoiceField(required=False, choices=[mode.value for mode in AgentMode]) is_sandbox = serializers.BooleanField(required=False, default=False) resume_payload = serializers.JSONField(required=False, allow_null=True) + attachments = serializers.ListField( + child=serializers.UUIDField(), + required=False, + max_length=MAX_ATTACHMENTS_PER_MESSAGE, + help_text="IDs of private PNG or JPEG attachments uploaded for this exact conversation.", + ) def validate(self, attrs): data = attrs + attachment_ids = data.get("attachments") or [] + if len(set(attachment_ids)) != len(attachment_ids): + raise serializers.ValidationError({"attachments": "Attachment IDs must be unique."}) + if attachment_ids and data["content"] is None: + raise serializers.ValidationError({"attachments": "Attachments cannot be used while resuming a stream."}) + if attachment_ids and ( + data.get("agent_mode") == AgentMode.RESEARCH.value or data.get("agent_mode") == AgentMode.RESEARCH + ): + raise serializers.ValidationError({"attachments": "Attachments are not supported in Research mode."}) + if attachment_ids and (data.get("is_sandbox") or data.get("agent_mode") == AgentMode.SANDBOX.value): + raise serializers.ValidationError({"attachments": "Attachments are not supported in Sandbox mode."}) + if data["content"] == "" and not attachment_ids: + raise serializers.ValidationError({"content": "A message or attachment is required."}) + + attachment_refs: list[dict[str, str | int]] = [] + if attachment_ids: + team = self.context["team"] + user = self.context["user"] + conversation_id = data["conversation"] + attachments = ConversationAttachment.objects.for_team(team.id).filter( + id__in=attachment_ids, + team=team, + creator=user, + conversation_id=conversation_id, + ) + attachments_by_id = {attachment.id: attachment for attachment in attachments} + if len(attachments_by_id) != len(attachment_ids): + raise serializers.ValidationError({"attachments": "One or more attachments were not found."}) + attachment_refs = [attachment_to_ref(attachments_by_id[attachment_id]) for attachment_id in attachment_ids] + if data["content"] is not None: try: message = HumanMessage.model_validate( { "content": data["content"], + "attachments": attachment_refs or None, "ui_context": data.get("ui_context"), "trace_id": str(data["trace_id"]) if data.get("trace_id") else None, } @@ -152,9 +200,17 @@ class QueueMessageSerializer(serializers.Serializer): ui_context = serializers.JSONField(required=False) billing_context = serializers.JSONField(required=False) agent_mode = serializers.ChoiceField(required=False, choices=[mode.value for mode in AgentMode]) + attachments = serializers.ListField( + child=serializers.UUIDField(), + required=False, + max_length=MAX_ATTACHMENTS_PER_MESSAGE, + help_text="Attachment IDs. Attachments are not supported for queued messages.", + ) def validate(self, attrs): data = attrs + if data.get("attachments"): + raise serializers.ValidationError({"attachments": "Attachments are not supported for queued messages."}) try: HumanMessage.model_validate( { @@ -187,6 +243,26 @@ class QueueMessageUpdateSerializer(serializers.Serializer): content = serializers.CharField(required=True, allow_blank=False, max_length=40000) +class ConversationAttachmentUploadSerializer(serializers.Serializer): + image = serializers.FileField( + required=True, + help_text="A PNG or JPEG image no larger than 4 MiB.", + ) + + +class ConversationAttachmentSerializer(serializers.Serializer): + id = serializers.UUIDField(read_only=True, help_text="Attachment identifier.") + file_name = serializers.CharField(read_only=True, help_text="Sanitized display filename.") + content_type = serializers.ChoiceField( + read_only=True, + choices=["image/png", "image/jpeg"], + help_text="Verified and re-encoded image content type.", + ) + size = serializers.IntegerField(read_only=True, help_text="Sanitized image size in bytes.") + width = serializers.IntegerField(read_only=True, help_text="Decoded image width in pixels.") + height = serializers.IntegerField(read_only=True, help_text="Decoded image height in pixels.") + + @extend_schema(tags=["max"]) class ConversationViewSet( TeamAndOrgViewSetMixin, ListModelMixin, RetrieveModelMixin, DestroyModelMixin, GenericViewSet @@ -314,6 +390,34 @@ def get_serializer_class(self): return ConversationMinimalSerializer return super().get_serializer_class() + def _owned_conversation_exists(self, request: Request, conversation_id: str) -> bool: + return Conversation.objects.filter( + id=conversation_id, + team=self.team, + user=request.user, + deleted=False, + ).exists() + + def _conversation_id_is_available_to_user(self, request: Request, conversation_id: str) -> bool: + conversation = Conversation.objects.filter(id=conversation_id).only("team_id", "user_id", "deleted").first() + return conversation is None or ( + not conversation.deleted + and conversation.team_id == self.team.id + and conversation.user_id == request.user.id + ) + + def _get_scoped_attachment(self, request: Request, attachment_id: str) -> ConversationAttachment: + conversation_id = self._queue_conversation_id() + try: + return ConversationAttachment.objects.for_team(self.team.id).get( + id=attachment_id, + team=self.team, + creator=request.user, + conversation_id=conversation_id, + ) + except (ConversationAttachment.DoesNotExist, ValidationError, ValueError): + raise exceptions.NotFound("Attachment not found.") + def get_serializer_context(self): context = super().get_serializer_context() context["team"] = self.team @@ -498,6 +602,90 @@ async def async_stream( content_type="text/event-stream", ) + @extend_schema( + request=ConversationAttachmentUploadSerializer, + responses={201: ConversationAttachmentSerializer}, + description="Upload and sanitize a private PNG or JPEG attachment for this conversation.", + ) + @action( + detail=True, + methods=["POST"], + url_path="attachments", + parser_classes=[MultiPartParser, FormParser], + ) + def upload_attachment(self, request: Request, *args, **kwargs) -> Response: + conversation_id = self._queue_conversation_id() + if not self._conversation_id_is_available_to_user(request, conversation_id): + raise exceptions.NotFound("Conversation not found.") + + serializer = ConversationAttachmentUploadSerializer(data=request.data) + serializer.is_valid(raise_exception=True) + try: + processed = process_conversation_attachment(serializer.validated_data["image"]) + attachment = save_conversation_attachment( + processed=processed, + team_id=self.team.id, + creator=cast(User, request.user), + conversation_id=uuid.UUID(conversation_id), + ) + except InvalidConversationAttachment as error: + raise exceptions.ValidationError({"image": str(error)}) + except (ObjectStorageError, ValueError): + logger.exception("Failed to save conversation attachment", conversation_id=conversation_id) + raise exceptions.APIException("Could not save attachment.") + + return Response(attachment_to_ref(attachment), status=status.HTTP_201_CREATED) + + @extend_schema( + parameters=[OpenApiParameter("attachment_id", OpenApiTypes.UUID, OpenApiParameter.PATH)], + responses={(200, "image/png"): OpenApiTypes.BINARY, (200, "image/jpeg"): OpenApiTypes.BINARY}, + description="Read private attachment content scoped to the authenticated user, team, and conversation.", + ) + @action( + detail=True, + methods=["GET"], + url_path=r"attachments/(?P[^/.]+)/content", + ) + def attachment_content(self, request: Request, attachment_id: str, *args, **kwargs) -> HttpResponse: + attachment = self._get_scoped_attachment(request, attachment_id) + try: + content = object_storage.read_bytes(attachment.object_path) + except ObjectStorageError: + logger.exception("Failed to read conversation attachment", attachment_id=attachment_id) + raise exceptions.NotFound("Attachment content not found.") + if content is None: + raise exceptions.NotFound("Attachment content not found.") + + return HttpResponse( + content, + content_type=attachment.content_type, + headers={ + "Cache-Control": "private, max-age=3600", + "Content-Disposition": f'inline; filename="{attachment.file_name}"', + "X-Content-Type-Options": "nosniff", + }, + ) + + @extend_schema( + parameters=[OpenApiParameter("attachment_id", OpenApiTypes.UUID, OpenApiParameter.PATH)], + responses={204: None}, + description="Delete a private conversation attachment.", + ) + @action( + detail=True, + methods=["DELETE"], + url_path=r"attachments/(?P[^/.]+)", + ) + def delete_attachment(self, request: Request, attachment_id: str, *args, **kwargs) -> Response: + attachment = self._get_scoped_attachment(request, attachment_id) + try: + object_storage.delete(attachment.object_path) + except ObjectStorageError: + logger.exception("Failed to delete conversation attachment", attachment_id=attachment_id) + raise exceptions.APIException("Could not delete attachment.") + attachment.delete() + return Response(status=status.HTTP_204_NO_CONTENT) + @action(detail=True, methods=["GET", "POST"], url_path="queue") def queue(self, request: Request, *args, **kwargs): conversation_id = self._queue_conversation_id() diff --git a/ee/api/test/test_conversation.py b/ee/api/test/test_conversation.py index fb82bebbb34e..1facc2641db9 100644 --- a/ee/api/test/test_conversation.py +++ b/ee/api/test/test_conversation.py @@ -1,14 +1,17 @@ import uuid import datetime +from io import BytesIO from typing import Any, cast from posthog.test.base import APIBaseTest from unittest.mock import AsyncMock, patch +from django.core.files.uploadedfile import SimpleUploadedFile from django.http import StreamingHttpResponse from django.test import override_settings from django.utils import timezone +from PIL import Image from rest_framework import status from posthog.schema import ( @@ -30,7 +33,7 @@ from posthog.temporal.ai.chat_agent import ChatAgentWorkflow, ChatAgentWorkflowInputs from posthog.temporal.ai.research_agent import ResearchAgentWorkflow, ResearchAgentWorkflowInputs -from products.posthog_ai.backend.models.assistant import Conversation +from products.posthog_ai.backend.models.assistant import Conversation, ConversationAttachment from ee.api.conversation import ConversationViewSet @@ -105,6 +108,181 @@ def _create_mock_streaming_response( cast(Any, mock_response).streaming_content = [content] return mock_response + def _image_upload( + self, image_format: str = "PNG", *, size: tuple[int, int] = (4, 3), name: str | None = None + ) -> SimpleUploadedFile: + output = BytesIO() + Image.new("RGB", size, color=(200, 10, 20)).save(output, format=image_format) + content_type = "image/png" if image_format == "PNG" else "image/jpeg" + return SimpleUploadedFile(name or f"image.{image_format.lower()}", output.getvalue(), content_type=content_type) + + @override_settings(OBJECT_STORAGE_ENABLED=True) + @patch("products.posthog_ai.backend.attachments.object_storage.write") + def test_upload_png_and_send_reference_only_workflow_payload(self, mock_write): + conversation_id = str(uuid.uuid4()) + upload_response = self.client.post( + f"/api/environments/{self.team.id}/conversations/{conversation_id}/attachments/", + {"image": self._image_upload("PNG", name="../../unsafe\nname.png")}, + format="multipart", + ) + self.assertEqual(upload_response.status_code, status.HTTP_201_CREATED) + self.assertEqual(upload_response.json()["content_type"], "image/png") + self.assertEqual(upload_response.json()["file_name"], "unsafename.png") + self.assertEqual(upload_response.json()["width"], 4) + stored_bytes = mock_write.call_args.args[1] + with Image.open(BytesIO(stored_bytes)) as stored_image: + self.assertEqual(stored_image.format, "PNG") + self.assertNotIn("exif", stored_image.info) + + with ( + patch("ee.hogai.core.executor.AgentExecutor.astream", return_value=_async_generator()) as mock_stream, + patch("ee.api.conversation.StreamingHttpResponse", side_effect=self._create_mock_streaming_response), + ): + response = self.client.post( + f"/api/environments/{self.team.id}/conversations/", + { + "content": "", + "trace_id": str(uuid.uuid4()), + "conversation": conversation_id, + "attachments": [upload_response.json()["id"]], + }, + format="json", + ) + self.assertEqual(response.status_code, status.HTTP_200_OK) + workflow_inputs = mock_stream.call_args.args[1] + attachment_ref = workflow_inputs.message["attachments"][0] + self.assertEqual(attachment_ref["id"], upload_response.json()["id"]) + self.assertNotIn("data", attachment_ref) + self.assertNotIn("object_path", attachment_ref) + + @override_settings(OBJECT_STORAGE_ENABLED=True) + @patch("products.posthog_ai.backend.attachments.object_storage.write") + def test_upload_jpeg_reencodes_without_metadata(self, mock_write): + conversation_id = str(uuid.uuid4()) + response = self.client.post( + f"/api/environments/{self.team.id}/conversations/{conversation_id}/attachments/", + {"image": self._image_upload("JPEG")}, + format="multipart", + ) + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + self.assertEqual(response.json()["content_type"], "image/jpeg") + with Image.open(BytesIO(mock_write.call_args.args[1])) as stored_image: + self.assertEqual(stored_image.format, "JPEG") + self.assertNotIn("exif", stored_image.info) + + @override_settings(OBJECT_STORAGE_ENABLED=True) + @patch("products.posthog_ai.backend.attachments.object_storage.write") + def test_upload_rejects_spoofed_mime_invalid_svg_size_and_pixel_cap(self, mock_write): + conversation_id = str(uuid.uuid4()) + cases = [ + SimpleUploadedFile("spoof.jpg", self._image_upload("PNG").read(), content_type="image/jpeg"), + SimpleUploadedFile("invalid.png", b"not an image", content_type="image/png"), + SimpleUploadedFile("image.svg", b"", content_type="image/svg+xml"), + SimpleUploadedFile("large.png", b"x" * (4 * 1024 * 1024 + 1), content_type="image/png"), + self._image_upload("PNG", size=(5001, 5000), name="too-many-pixels.png"), + ] + for image in cases: + with self.subTest(image=image.name): + response = self.client.post( + f"/api/environments/{self.team.id}/conversations/{conversation_id}/attachments/", + {"image": image}, + format="multipart", + ) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + mock_write.assert_not_called() + + @override_settings(OBJECT_STORAGE_ENABLED=True) + def test_attachment_content_is_authenticated_and_exactly_scoped(self): + conversation_id = uuid.uuid4() + attachment = ConversationAttachment.objects.for_team(self.team.id).create( + team=self.team, + creator=self.user, + conversation_id=conversation_id, + file_name="private.png", + content_type="image/png", + size=3, + width=1, + height=1, + object_path="private/random", + ) + content_url = ( + f"/api/environments/{self.team.id}/conversations/{conversation_id}/attachments/{attachment.id}/content/" + ) + with patch("ee.api.conversation.object_storage.read_bytes", return_value=b"png"): + response = self.client.get(content_url) + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(response.content, b"png") + self.assertEqual(response["Cache-Control"], "private, max-age=3600") + self.assertEqual(response["X-Content-Type-Options"], "nosniff") + + self.client.logout() + self.assertEqual(self.client.get(content_url).status_code, status.HTTP_401_UNAUTHORIZED) + + self.client.force_login(self.other_user) + self.assertEqual(self.client.get(content_url).status_code, status.HTTP_404_NOT_FOUND) + + self.client.force_login(self.user) + other_team_url = content_url.replace(str(self.team.id), str(self.other_team.id), 1) + self.assertEqual(self.client.get(other_team_url).status_code, status.HTTP_404_NOT_FOUND) + wrong_conversation_url = content_url.replace(str(conversation_id), str(uuid.uuid4()), 1) + self.assertEqual(self.client.get(wrong_conversation_url).status_code, status.HTTP_404_NOT_FOUND) + + @override_settings(OBJECT_STORAGE_ENABLED=True) + @patch("products.posthog_ai.backend.attachments.object_storage.write") + def test_upload_rejects_conversation_owned_by_another_user(self, _mock_write): + conversation = Conversation.objects.create(user=self.other_user, team=self.team) + response = self.client.post( + f"/api/environments/{self.team.id}/conversations/{conversation.id}/attachments/", + {"image": self._image_upload()}, + format="multipart", + ) + self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND) + + @override_settings(OBJECT_STORAGE_ENABLED=True) + def test_attachments_are_rejected_for_research_sandbox_and_queue(self): + conversation = Conversation.objects.create(user=self.user, team=self.team) + attachment = ConversationAttachment.objects.for_team(self.team.id).create( + team=self.team, + creator=self.user, + conversation_id=conversation.id, + file_name="private.png", + content_type="image/png", + size=3, + width=1, + height=1, + object_path="private/random", + ) + for payload in [ + { + "content": "research", + "trace_id": str(uuid.uuid4()), + "conversation": str(conversation.id), + "agent_mode": AgentMode.RESEARCH.value, + "attachments": [str(attachment.id)], + }, + { + "content": "sandbox", + "trace_id": str(uuid.uuid4()), + "conversation": str(conversation.id), + "is_sandbox": True, + "attachments": [str(attachment.id)], + }, + ]: + with self.subTest(payload=payload): + response = self.client.post( + f"/api/environments/{self.team.id}/conversations/", + payload, + format="json", + ) + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + + queue_response = self.client.post( + f"/api/environments/{self.team.id}/conversations/{conversation.id}/queue/", + {"content": "queue", "attachments": [str(attachment.id)]}, + format="json", + ) + self.assertEqual(queue_response.status_code, status.HTTP_400_BAD_REQUEST) + def test_create_conversation(self): conversation_id = str(uuid.uuid4()) diff --git a/ee/hogai/core/agent_modes/executables.py b/ee/hogai/core/agent_modes/executables.py index 8837c1b5cfce..69e30be4934c 100644 --- a/ee/hogai/core/agent_modes/executables.py +++ b/ee/hogai/core/agent_modes/executables.py @@ -1,7 +1,7 @@ import asyncio from collections.abc import Mapping, Sequence from typing import Any, Literal, TypeVar, cast -from uuid import uuid4 +from uuid import UUID, uuid4 from django.conf import settings @@ -35,6 +35,8 @@ from posthog.models import Team, User from posthog.sync import database_sync_to_async +from products.posthog_ai.backend.attachments import load_attachment_blocks + from ee.hogai.core.agent_modes.prompt_builder import AgentPromptBuilder from ee.hogai.core.agent_modes.prompts import ( ROOT_CONVERSATION_SUMMARY_PROMPT, @@ -145,8 +147,27 @@ async def arun(self, state: AssistantState, config: RunnableConfig) -> PartialAs messages_to_replace = updated_messages # Calculate the initial window. + conversation_messages = messages_to_replace or state.messages + configurable = config.get("configurable") or {} + conversation_id = configurable.get("thread_id") + human_messages_with_attachments = [ + message for message in conversation_messages if isinstance(message, HumanMessage) and message.attachments + ] + attachment_blocks = {} + if human_messages_with_attachments: + if not conversation_id: + raise ValueError("Conversation ID is required to hydrate attachments.") + attachment_blocks = await load_attachment_blocks( + human_messages_with_attachments, + team_id=self._team.id, + creator_id=self._user.id, + conversation_id=UUID(str(conversation_id)), + ) langchain_messages = self._construct_messages( - messages_to_replace or state.messages, state.root_conversation_start_id, state.root_tool_calls_count + conversation_messages, + state.root_conversation_start_id, + state.root_tool_calls_count, + attachment_blocks=attachment_blocks, ) window_id = state.root_conversation_start_id start_id = state.start_id @@ -181,7 +202,22 @@ async def arun(self, state: AssistantState, config: RunnableConfig) -> PartialAs messages_to_replace = insertion_result.messages # Update the window - langchain_messages = self._construct_messages(messages_to_replace, window_id, state.root_tool_calls_count) + attachment_blocks = await load_attachment_blocks( + [ + message + for message in messages_to_replace + if isinstance(message, HumanMessage) and message.attachments + ], + team_id=self._team.id, + creator_id=self._user.id, + conversation_id=UUID(str(conversation_id)), + ) + langchain_messages = self._construct_messages( + messages_to_replace, + window_id, + state.root_tool_calls_count, + attachment_blocks=attachment_blocks, + ) system_prompts = cast(list[BaseMessage], system_prompts) assert len(system_prompts) > 0 @@ -314,6 +350,7 @@ def _construct_messages( messages: Sequence[AssistantMessageUnion], window_start_id: str | None = None, tool_calls_count: int | None = None, + attachment_blocks: Mapping[int, list[dict[str, Any]]] | None = None, ) -> list[BaseMessage]: conversation_window = self._window_manager.get_messages_in_window(messages, window_start_id) @@ -321,7 +358,9 @@ def _construct_messages( tool_result_messages = self._get_tool_map(conversation_window) # Convert to Anthropic messages - history = self._convert_to_langchain_messages(conversation_window, tool_result_messages) + history = self._convert_to_langchain_messages( + conversation_window, tool_result_messages, attachment_blocks=attachment_blocks + ) # Force the agent to stop if the tool call limit is reached. history = self._add_limit_message_if_reached(history, tool_calls_count) @@ -359,9 +398,10 @@ def _convert_to_langchain_messages( self, messages: Sequence[AssistantMessageUnion], tool_result_messages: Mapping[str, AssistantToolCallMessage], + attachment_blocks: Mapping[int, list[dict[str, Any]]] | None = None, ) -> list[BaseMessage]: """Convert a conversation window to a list of Langchain messages.""" - return convert_to_anthropic_messages(messages, tool_result_messages) + return convert_to_anthropic_messages(messages, tool_result_messages, attachment_blocks) def _add_cache_control_to_last_message(self, messages: list[BaseMessage]) -> list[BaseMessage]: """Add cache control to the last message.""" diff --git a/ee/hogai/core/agent_modes/test/test_executables.py b/ee/hogai/core/agent_modes/test/test_executables.py index 602432a80c00..8bbd5c210c81 100644 --- a/ee/hogai/core/agent_modes/test/test_executables.py +++ b/ee/hogai/core/agent_modes/test/test_executables.py @@ -183,6 +183,29 @@ async def test_node_reconstructs_conversation(self, mock_model): ], ) + def test_node_reconstructs_multimodal_human_message_at_model_boundary(self): + node = _create_agent_node(self.team, self.user) + message = HumanMessage(content="Describe this") + image_block = { + "type": "image", + "source": {"type": "base64", "media_type": "image/png", "data": "cG5n"}, + } + result = node._construct_messages( + [message], + attachment_blocks={id(message): [image_block]}, + ) + self.assertEqual( + result, + [ + LangchainHumanMessage( + content=[ + image_block, + {"type": "text", "text": "Describe this", "cache_control": {"type": "ephemeral"}}, + ] + ) + ], + ) + @patch( "ee.hogai.core.agent_modes.executables.AgentExecutable._get_model", return_value=FakeChatAnthropic(responses=[]), diff --git a/ee/hogai/utils/anthropic.py b/ee/hogai/utils/anthropic.py index 8c41db57cc9f..fc316f2a39d8 100644 --- a/ee/hogai/utils/anthropic.py +++ b/ee/hogai/utils/anthropic.py @@ -37,8 +37,13 @@ def add_cache_control(message: BaseMessage, ttl: Literal["5m", "1h"] | None = No return message -def convert_human_message_to_anthropic_message(message: HumanMessage) -> messages.HumanMessage: - return messages.HumanMessage(content=[{"type": "text", "text": message.content}]) +def convert_human_message_to_anthropic_message( + message: HumanMessage, attachment_blocks: Mapping[int, list[dict[str, Any]]] | None = None +) -> messages.HumanMessage: + content: list[dict[str, Any]] = list((attachment_blocks or {}).get(id(message), [])) + if message.content: + content.append({"type": "text", "text": message.content}) + return messages.HumanMessage(content=content) def convert_context_message_to_anthropic_message(message: ContextMessage) -> messages.HumanMessage: @@ -105,10 +110,12 @@ def convert_failure_message_to_anthropic_message(message: FailureMessage) -> mes def convert_to_anthropic_message( - message: AssistantMessageUnion, tool_result_map: Mapping[str, AssistantToolCallMessage] + message: AssistantMessageUnion, + tool_result_map: Mapping[str, AssistantToolCallMessage], + attachment_blocks: Mapping[int, list[dict[str, Any]]] | None = None, ) -> list[messages.BaseMessage]: if isinstance(message, HumanMessage): - return [convert_human_message_to_anthropic_message(message)] + return [convert_human_message_to_anthropic_message(message, attachment_blocks)] if isinstance(message, ContextMessage): return [convert_context_message_to_anthropic_message(message)] elif isinstance(message, AssistantMessage): @@ -121,11 +128,12 @@ def convert_to_anthropic_message( def convert_to_anthropic_messages( conversation: Sequence[AssistantMessageUnion], tool_result_map: Mapping[str, AssistantToolCallMessage], + attachment_blocks: Mapping[int, list[dict[str, Any]]] | None = None, ) -> list[messages.BaseMessage]: history: list[messages.BaseMessage] = [] for message in conversation: try: - history.extend(convert_to_anthropic_message(message, tool_result_map)) + history.extend(convert_to_anthropic_message(message, tool_result_map, attachment_blocks)) except ValueError: continue return history diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index d37adb8ae909..395858d51c50 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -242,8 +242,10 @@ import type { HogFlowTemplate, } from 'products/workflows/frontend/Workflows/hogflows/types' -import { AgentMode } from '../queries/schema' +import { AgentMode, HumanMessage } from '../queries/schema' import type { MaxUIContext } from '../scenes/max/maxTypes' + +export type ConversationAttachment = NonNullable[number] import { AlertSimulationResult, AlertType, AlertTypeWrite } from './components/Alerts/types' import { ErrorTrackingFingerprint, @@ -6484,6 +6486,25 @@ const api = { }, conversations: { + attachments: { + upload(conversationId: string, image: File): Promise { + const data = new FormData() + data.append('image', image) + return new ApiRequest().conversation(conversationId).withAction('attachments').create({ data }) + }, + + delete(conversationId: string, attachmentId: string): Promise { + return new ApiRequest().conversation(conversationId).withAction(`attachments/${attachmentId}`).delete() + }, + + contentUrl(conversationId: string, attachmentId: string): string { + return new ApiRequest() + .conversation(conversationId) + .withAction(`attachments/${attachmentId}/content`) + .assembleFullUrl() + }, + }, + async stream( data: { /** The user message. Null content means we're resuming streaming or continuing previous generation. */ @@ -6494,6 +6515,7 @@ const api = { conversation?: string | null trace_id: string agent_mode?: AgentMode | null + attachments?: string[] resume_payload?: { action: 'approve' | 'reject' proposal_id: string diff --git a/frontend/src/queries/schema.json b/frontend/src/queries/schema.json index 18e8abe82883..8d770979b0bc 100644 --- a/frontend/src/queries/schema.json +++ b/frontend/src/queries/schema.json @@ -23688,6 +23688,12 @@ "HumanMessage": { "additionalProperties": false, "properties": { + "attachments": { + "items": { + "$ref": "#/definitions/HumanMessageAttachment" + }, + "type": "array" + }, "content": { "type": "string" }, @@ -23711,6 +23717,32 @@ "required": ["type", "content"], "type": "object" }, + "HumanMessageAttachment": { + "additionalProperties": false, + "properties": { + "content_type": { + "enum": ["image/png", "image/jpeg"], + "type": "string" + }, + "file_name": { + "type": "string" + }, + "height": { + "type": "number" + }, + "id": { + "type": "string" + }, + "size": { + "type": "number" + }, + "width": { + "type": "number" + } + }, + "required": ["id", "file_name", "content_type", "size", "width", "height"], + "type": "object" + }, "IQRDetectorConfig": { "additionalProperties": false, "properties": { diff --git a/frontend/src/queries/schema/schema-assistant-messages.ts b/frontend/src/queries/schema/schema-assistant-messages.ts index 4c197b77ea8b..a60ecad381d6 100644 --- a/frontend/src/queries/schema/schema-assistant-messages.ts +++ b/frontend/src/queries/schema/schema-assistant-messages.ts @@ -90,9 +90,19 @@ export interface BaseAssistantMessage { parent_tool_call_id?: string } +export interface HumanMessageAttachment { + id: string + file_name: string + content_type: 'image/png' | 'image/jpeg' + size: number + width: number + height: number +} + export interface HumanMessage extends BaseAssistantMessage { type: AssistantMessageType.Human content: string + attachments?: HumanMessageAttachment[] ui_context?: MaxUIContext trace_id?: string } diff --git a/frontend/src/scenes/max/Thread.tsx b/frontend/src/scenes/max/Thread.tsx index b8b7b7aa9e77..49dee8e5b298 100644 --- a/frontend/src/scenes/max/Thread.tsx +++ b/frontend/src/scenes/max/Thread.tsx @@ -32,6 +32,7 @@ import { Tooltip, } from '@posthog/lemon-ui' +import api from 'lib/api' import { InsightBreakdownSummary, PropertiesSummary, @@ -392,7 +393,13 @@ function Message({ }: MessageProps): JSX.Element | null { const { editInsightToolRegistered, registeredToolMap } = useValues(maxGlobalLogic) const { activeTabId, activeSceneId } = useValues(sceneLogic) - const { threadLoading, isSharedThread, pendingApprovalsData, resolvedApprovalStatuses } = useValues(maxThreadLogic) + const { + threadLoading, + isSharedThread, + pendingApprovalsData, + resolvedApprovalStatuses, + conversationId: threadConversationId, + } = useValues(maxThreadLogic) const { conversationId } = useValues(maxLogic) const groupType = message.type === 'human' ? 'human' : 'ai' @@ -473,6 +480,23 @@ function Message({ useCurrentPageContext={false} /> )} + {message.attachments && message.attachments.length > 0 && ( +
+ {message.attachments.map((attachment) => ( +
+ {attachment.file_name} +
+ ))} +
+ )} {maybeCommand ? (
{message.content}
- ) : ( - - )} + ) : message.content ? ( + + ) : null} ) } else if (isAssistantMessage(message)) { diff --git a/frontend/src/scenes/max/components/QuestionInput.test.tsx b/frontend/src/scenes/max/components/QuestionInput.test.tsx index 3ed3fa138e07..77700f1fcebe 100644 --- a/frontend/src/scenes/max/components/QuestionInput.test.tsx +++ b/frontend/src/scenes/max/components/QuestionInput.test.tsx @@ -3,6 +3,8 @@ import '@testing-library/jest-dom' import { cleanup, fireEvent, render, screen, waitFor } from '@testing-library/react' import { BindLogic, Provider } from 'kea' +import api from 'lib/api' + import { useMocks } from '~/mocks/jest' import { initKeaTests } from '~/test/init' @@ -26,6 +28,8 @@ describe('QuestionInput slash command autocomplete', () => { let threadLogicInstance: ReturnType beforeEach(() => { + Object.defineProperty(URL, 'createObjectURL', { configurable: true, value: jest.fn(() => 'blob:preview') }) + Object.defineProperty(URL, 'revokeObjectURL', { configurable: true, value: jest.fn() }) useMocks(maxMocks) initKeaTests() @@ -76,4 +80,131 @@ describe('QuestionInput slash command autocomplete', () => { fireEvent.change(input, { target: { value: '/' } }) await waitFor(() => expect(slashCommandItem()).toBeInTheDocument()) }) + + it.each([ + [ + 'picker', + (_input: HTMLTextAreaElement, fileInput: HTMLInputElement, file: File) => + fireEvent.change(fileInput, { target: { files: [file] } }), + ], + [ + 'drop', + (input: HTMLTextAreaElement, _fileInput: HTMLInputElement, file: File) => + fireEvent.drop(input.closest('label')!, { + dataTransfer: { files: [file], types: ['Files'] }, + }), + ], + [ + 'paste', + (input: HTMLTextAreaElement, _fileInput: HTMLInputElement, file: File) => + fireEvent.paste(input, { + clipboardData: { + items: [{ kind: 'file', type: file.type, getAsFile: () => file }], + }, + }), + ], + ])('uploads and previews an image from %s', async (_source, addFile) => { + jest.spyOn(api.conversations.attachments, 'upload').mockResolvedValue({ + id: 'attachment-id', + file_name: 'chart.png', + content_type: 'image/png', + size: 3, + width: 1, + height: 1, + }) + const input = screen.getByRole('textbox') as HTMLTextAreaElement + const fileInput = screen.getByLabelText('Choose PNG or JPEG images') as HTMLInputElement + addFile(input, fileInput, new File(['png'], 'chart.png', { type: 'image/png' })) + + await waitFor(() => expect(screen.getByAltText('chart.png')).toBeInTheDocument()) + expect(threadLogicInstance.values.pendingAttachments[0]?.attachment?.id).toBe('attachment-id') + }) + + it('removes an uploaded image and revokes its preview URL', async () => { + jest.spyOn(api.conversations.attachments, 'upload').mockResolvedValue({ + id: 'attachment-id', + file_name: 'chart.png', + content_type: 'image/png', + size: 3, + width: 1, + height: 1, + }) + jest.spyOn(api.conversations.attachments, 'delete').mockResolvedValue() + const fileInput = screen.getByLabelText('Choose PNG or JPEG images') + fireEvent.change(fileInput, { + target: { files: [new File(['png'], 'chart.png', { type: 'image/png' })] }, + }) + await waitFor(() => expect(screen.getByAltText('chart.png')).toBeInTheDocument()) + + fireEvent.click(screen.getByLabelText('Remove chart.png')) + await waitFor(() => expect(screen.queryByAltText('chart.png')).not.toBeInTheDocument()) + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:preview') + expect(api.conversations.attachments.delete).toHaveBeenCalledWith( + maxLogicInstance.values.frontendConversationId, + 'attachment-id' + ) + }) + + it('limits attachments to four files', async () => { + jest.spyOn(api.conversations.attachments, 'upload').mockImplementation(async (_conversationId, file) => ({ + id: file.name, + file_name: file.name, + content_type: 'image/png', + size: file.size, + width: 1, + height: 1, + })) + const files = Array.from( + { length: 5 }, + (_, index) => new File(['png'], `chart-${index}.png`, { type: 'image/png' }) + ) + fireEvent.change(screen.getByLabelText('Choose PNG or JPEG images'), { target: { files } }) + + await waitFor(() => expect(threadLogicInstance.values.pendingAttachments).toHaveLength(4)) + expect(screen.getByRole('alert')).toHaveTextContent('up to 4 images') + }) + + it('rejects unsupported files and guards image-only submit while uploading', async () => { + let finishUpload: ((value: any) => void) | undefined + jest.spyOn(api.conversations.attachments, 'upload').mockImplementation( + () => new Promise((resolve) => (finishUpload = resolve)) + ) + const fileInput = screen.getByLabelText('Choose PNG or JPEG images') + fireEvent.change(fileInput, { + target: { files: [new File(['svg'], 'image.svg', { type: 'image/svg+xml' })] }, + }) + expect(await screen.findByRole('alert')).toHaveTextContent('Only PNG and JPEG') + + fireEvent.change(fileInput, { + target: { + files: [new File([new Uint8Array(4 * 1024 * 1024 + 1)], 'large.png', { type: 'image/png' })], + }, + }) + expect(await screen.findByRole('alert')).toHaveTextContent('4 MiB or smaller') + + fireEvent.change(fileInput, { + target: { files: [new File(['png'], 'chart.png', { type: 'image/png' })] }, + }) + await waitFor(() => expect(threadLogicInstance.values.submissionDisabledReason).toContain('finish uploading')) + + finishUpload?.({ + id: 'attachment-id', + file_name: 'chart.png', + content_type: 'image/png', + size: 3, + width: 1, + height: 1, + }) + await waitFor(() => expect(threadLogicInstance.values.submissionDisabledReason).toBeUndefined()) + fireEvent.click(document.querySelector('[data-attr="max-send-message"]')!) + await waitFor(() => + expect(threadLogicInstance.values.threadRaw).toEqual([ + expect.objectContaining({ + type: 'human', + content: '', + attachments: [expect.objectContaining({ id: 'attachment-id' })], + }), + ]) + ) + }) }) diff --git a/frontend/src/scenes/max/components/QuestionInput.tsx b/frontend/src/scenes/max/components/QuestionInput.tsx index 4a5d19f4209d..e90b53a304ad 100644 --- a/frontend/src/scenes/max/components/QuestionInput.tsx +++ b/frontend/src/scenes/max/components/QuestionInput.tsx @@ -5,7 +5,7 @@ import { useActions, useValues } from 'kea' import posthog from 'posthog-js' import React, { ReactNode, useEffect, useMemo, useRef, useState } from 'react' -import { IconArrowRight, IconCheck, IconPencil, IconStopFilled, IconTrash, IconX } from '@posthog/icons' +import { IconArrowRight, IconCheck, IconImage, IconPencil, IconStopFilled, IconTrash, IconX } from '@posthog/icons' import { LemonButton, LemonSwitch, LemonTextArea, Spinner } from '@posthog/lemon-ui' import { useFeatureFlag } from 'lib/hooks/useFeatureFlag' @@ -163,9 +163,19 @@ export const QuestionInput = React.forwardRef(null) + const [isDraggingImages, setIsDraggingImages] = useState(false) + const fileInputRef = useRef(null) const displayQueuedMessages = useMemo(() => [...queuedMessages].reverse(), [queuedMessages]) const hasQuestion = question.trim().length > 0 + const hasReadyAttachment = pendingAttachments.some((attachment) => attachment.status === 'ready') + const hasInput = hasQuestion || hasReadyAttachment const isQueueingSubmission = queueingEnabled && threadLoading && hasQuestion const showStopButton = threadLoading && !isQueueingSubmission @@ -280,7 +294,76 @@ export const QuestionInput = React.forwardRef { + if (event.dataTransfer.types.includes('Files')) { + event.preventDefault() + setIsDraggingImages(true) + } + }} + onDragOver={(event) => { + if (event.dataTransfer.types.includes('Files')) { + event.preventDefault() + } + }} + onDragLeave={(event) => { + if (!event.currentTarget.contains(event.relatedTarget as Node | null)) { + setIsDraggingImages(false) + } + }} + onDrop={(event) => { + event.preventDefault() + setIsDraggingImages(false) + addPendingAttachments(Array.from(event.dataTransfer.files)) + }} > + {isDraggingImages && ( +
+ Drop PNG or JPEG images here +
+ )} + {pendingAttachments.length > 0 && ( +
+ {pendingAttachments.map((attachment) => ( +
+ {attachment.fileName} + {attachment.status === 'uploading' && ( + + )} + } + aria-label={`Remove ${attachment.fileName}`} + className="absolute right-0 top-0 bg-surface-primary" + onClick={(event) => { + event.preventDefault() + removePendingAttachment(attachment.localId) + }} + /> + {attachment.status === 'error' && ( + + Failed + + )} +
+ ))} +
+ )} + {attachmentInputError && ( +
+ {attachmentInputError} +
+ )} {handsFreeActive ? ( ) : ( @@ -291,7 +374,20 @@ export const QuestionInput = React.forwardRef -
+
) => { + const imageFiles = Array.from(event.clipboardData.items) + .filter((item) => item.kind === 'file' && item.type.startsWith('image/')) + .flatMap((item) => { + const file = item.getAsFile() + return file ? [file] : [] + }) + if (imageFiles.length > 0) { + addPendingAttachments(imageFiles) + } + }} + > {!question && (
setQuestion(value)} onPressEnter={() => { if ( - hasQuestion && + hasInput && !submissionDisabledReason && (!threadLoading || queueingEnabled) ) { @@ -411,6 +507,34 @@ export const QuestionInput = React.forwardRef + { + addPendingAttachments(Array.from(event.target.files ?? [])) + event.target.value = '' + }} + /> + {!handsFreeActive && ( + } + tooltip="Attach PNG or JPEG images" + aria-label="Attach images" + onClick={() => { + setAttachmentInputError(null) + fileInputRef.current?.click() + }} + disabledReason={ + pendingAttachments.length >= 4 ? 'You can attach up to 4 images' : undefined + } + /> + )} {!handsFreeActive && ( { if (threadLoading) { if (isQueueingSubmission) { diff --git a/frontend/src/scenes/max/maxLogic.tsx b/frontend/src/scenes/max/maxLogic.tsx index c42a47021a45..4aaa07b21bf9 100644 --- a/frontend/src/scenes/max/maxLogic.tsx +++ b/frontend/src/scenes/max/maxLogic.tsx @@ -18,7 +18,7 @@ import { sidePanelStateLogic } from '~/layout/navigation-3000/sidepanel/sidePane import { iconForType } from '~/layout/panel-layout/ProjectTree/defaultTree' import { actionsModel } from '~/models/actionsModel' import { productUrls } from '~/products' -import { AgentMode, RootAssistantMessage } from '~/queries/schema/schema-assistant-messages' +import { AgentMode, HumanMessage, RootAssistantMessage } from '~/queries/schema/schema-assistant-messages' import { Breadcrumb, Conversation, @@ -192,10 +192,16 @@ export const maxLogic = kea([ actions({ setQuestion: (question: string) => ({ question }), // update the form input - askMax: (prompt: string | null, addToThread: boolean = true, uiContext?: Partial) => ({ + askMax: ( + prompt: string | null, + addToThread: boolean = true, + uiContext?: Partial, + attachments?: NonNullable + ) => ({ prompt, addToThread, uiContext, + attachments, }), // used by maxThreadLogic to start a conversation scrollThreadToBottom: (behavior?: 'instant' | 'smooth') => ({ behavior }), openConversation: (conversationId: string) => ({ conversationId }), diff --git a/frontend/src/scenes/max/maxThreadLogic.test.ts b/frontend/src/scenes/max/maxThreadLogic.test.ts index 1a7e3b50f758..c3844aa0e129 100644 --- a/frontend/src/scenes/max/maxThreadLogic.test.ts +++ b/frontend/src/scenes/max/maxThreadLogic.test.ts @@ -997,6 +997,33 @@ describe('maxThreadLogic', () => { otherLogic.unmount() }) + it('preserves attachment metadata in the optimistic human message', async () => { + const streamSpy = mockStream() + const attachments = [ + { + id: 'attachment-id', + file_name: 'chart.png', + content_type: 'image/png' as const, + size: 3, + width: 1, + height: 1, + }, + ] + + await expectLogic(logic, () => { + logic.actions.askMax('', true, undefined, attachments) + }) + + expect(streamSpy).toHaveBeenCalledWith( + expect.objectContaining({ + content: '', + attachments: ['attachment-id'], + }), + expect.any(Object) + ) + expect((logic.values.threadRaw[0] as HumanMessage).attachments).toEqual(attachments) + }) + it('switches which thread processes askMax when activeThreadKey changes', async () => { const streamSpy = mockStream() diff --git a/frontend/src/scenes/max/maxThreadLogic.tsx b/frontend/src/scenes/max/maxThreadLogic.tsx index b8f6eaa8a783..e08d3c321863 100644 --- a/frontend/src/scenes/max/maxThreadLogic.tsx +++ b/frontend/src/scenes/max/maxThreadLogic.tsx @@ -18,7 +18,7 @@ import { router } from 'kea-router' import { subscriptions } from 'kea-subscriptions' import posthog from 'posthog-js' -import api, { ApiError } from 'lib/api' +import api, { ApiError, ConversationAttachment } from 'lib/api' import { JSONContent } from 'lib/components/RichContentEditor/types' import { FEATURE_FLAGS } from 'lib/constants' import { dayjs } from 'lib/dayjs' @@ -94,6 +94,9 @@ export const PENDING_AI_PROMPT_KEY = 'posthog_ai_pending_prompt' // dashboard context is included. Bounded so a stuck/failed load never blocks sending. export const MAX_DASHBOARD_CONTEXT_WAIT_MS = 8000 const DASHBOARD_CONTEXT_POLL_INTERVAL_MS = 100 +export const MAX_CONVERSATION_ATTACHMENTS = 4 +export const MAX_CONVERSATION_ATTACHMENT_BYTES = 4 * 1024 * 1024 +export const ACCEPTED_CONVERSATION_ATTACHMENT_TYPES = ['image/png', 'image/jpeg'] as const export type MessageStatus = 'loading' | 'completed' | 'error' @@ -101,6 +104,15 @@ export type ThreadMessage = RootAssistantMessage & { status: MessageStatus } +export type PendingConversationAttachment = { + localId: string + fileName: string + previewUrl: string + status: 'uploading' | 'ready' | 'error' + attachment?: ConversationAttachment + error?: string +} + const FAILURE_MESSAGE: FailureMessage & ThreadMessage = { type: AssistantMessageType.Failure, content: 'Oops! It looks like I’m having trouble answering this. Could you please try again?', @@ -198,6 +210,7 @@ export const maxThreadLogic = kea([ contextual_tools?: Record ui_context?: any resume_payload?: ResumePayload | null + attachments?: ConversationAttachment[] }, generationAttempt: number, addToThread: boolean = true @@ -281,6 +294,21 @@ export const maxThreadLogic = kea([ }), addPendingApprovalData: (approval: PendingApproval) => ({ approval }), loadPendingApprovalsData: (approvals: PendingApproval[]) => ({ approvals }), + addPendingAttachments: (files: File[]) => ({ files }), + startPendingAttachmentUpload: (attachment: PendingConversationAttachment, file: File) => ({ + attachment, + file, + }), + pendingAttachmentUploadSuccess: (localId: string, attachment: ConversationAttachment) => ({ + localId, + attachment, + }), + pendingAttachmentUploadFailure: (localId: string, error: string) => ({ localId, error }), + removePendingAttachment: (localId: string) => ({ localId }), + pendingAttachmentRemoved: (localId: string) => ({ localId }), + clearPendingAttachments: true, + pendingAttachmentsCleared: true, + setAttachmentInputError: (error: string | null) => ({ error }), }), reducers(({ props }) => ({ @@ -577,6 +605,30 @@ export const maxThreadLogic = kea([ resetThread: () => [], }, ], + pendingAttachments: [ + [] as PendingConversationAttachment[], + { + startPendingAttachmentUpload: (state, { attachment }) => [...state, attachment], + pendingAttachmentUploadSuccess: (state, { localId, attachment }) => + state.map((item) => + item.localId === localId ? { ...item, status: 'ready' as const, attachment } : item + ), + pendingAttachmentUploadFailure: (state, { localId, error }) => + state.map((item) => + item.localId === localId ? { ...item, status: 'error' as const, error } : item + ), + pendingAttachmentRemoved: (state, { localId }) => state.filter((item) => item.localId !== localId), + pendingAttachmentsCleared: () => [], + }, + ], + attachmentInputError: [ + null as string | null, + { + setAttachmentInputError: (_, { error }) => error, + addPendingAttachments: () => null, + pendingAttachmentsCleared: () => null, + }, + ], })), loaders(({ values }) => ({ @@ -624,10 +676,16 @@ export const maxThreadLogic = kea([ const traceId = uuid() actions.setTraceId(traceId) - if (generationAttempt === 0 && streamData.content && addToThread) { + if ( + generationAttempt === 0 && + streamData.content !== null && + (streamData.content.length > 0 || (streamData.attachments?.length ?? 0) > 0) && + addToThread + ) { const message: ThreadMessage = { type: AssistantMessageType.Human, content: streamData.content, + ...(streamData.attachments?.length ? { attachments: streamData.attachments } : {}), status: 'completed', trace_id: traceId, } @@ -643,6 +701,9 @@ export const maxThreadLogic = kea([ // Ensure we have valid data for the API call const apiData: any = { ...streamData } apiData.trace_id = traceId + if (streamData.attachments) { + apiData.attachments = streamData.attachments.map((attachment) => attachment.id) + } if (values.billingContext && values.featureFlags[FEATURE_FLAGS.MAX_BILLING_CONTEXT]) { apiData.billing_context = values.billingContext @@ -1003,7 +1064,88 @@ export const maxThreadLogic = kea([ actions.clearQueuedMessages() } }, - askMax: async ({ prompt, addToThread = true, uiContext }, breakpoint) => { + addPendingAttachments: ({ files }) => { + if (values.agentMode === AgentMode.Research) { + actions.setAttachmentInputError('Images are not supported in Research mode.') + return + } + if (values.isSandboxMode) { + actions.setAttachmentInputError('Images are not supported in Sandbox mode.') + return + } + + let availableSlots = MAX_CONVERSATION_ATTACHMENTS - values.pendingAttachments.length + if (files.length > availableSlots) { + actions.setAttachmentInputError(`You can attach up to ${MAX_CONVERSATION_ATTACHMENTS} images.`) + } + for (const file of files) { + if (availableSlots <= 0) { + break + } + if (!ACCEPTED_CONVERSATION_ATTACHMENT_TYPES.includes(file.type as 'image/png' | 'image/jpeg')) { + actions.setAttachmentInputError('Only PNG and JPEG images are supported.') + continue + } + if (file.size > MAX_CONVERSATION_ATTACHMENT_BYTES) { + actions.setAttachmentInputError('Images must be 4 MiB or smaller.') + continue + } + const localId = uuid() + const previewUrl = URL.createObjectURL(file) + cache.disposables.add(() => () => URL.revokeObjectURL(previewUrl), `attachment-preview-${localId}`, { + pauseOnPageHidden: false, + }) + actions.startPendingAttachmentUpload( + { + localId, + fileName: file.name, + previewUrl, + status: 'uploading', + }, + file + ) + availableSlots -= 1 + } + }, + startPendingAttachmentUpload: async ({ attachment, file }) => { + try { + const uploadedAttachment = await api.conversations.attachments.upload(values.conversationId, file) + if (!values.pendingAttachments.some((item) => item.localId === attachment.localId)) { + await api.conversations.attachments.delete(values.conversationId, uploadedAttachment.id) + return + } + actions.pendingAttachmentUploadSuccess(attachment.localId, uploadedAttachment) + } catch (error: any) { + posthog.captureException(error) + actions.pendingAttachmentUploadFailure( + attachment.localId, + error?.data?.image?.[0] || error?.data?.detail || 'Failed to upload image.' + ) + } + }, + removePendingAttachment: async ({ localId }) => { + const pendingAttachment = values.pendingAttachments.find((item) => item.localId === localId) + if (!pendingAttachment) { + return + } + cache.disposables.dispose(`attachment-preview-${localId}`) + actions.pendingAttachmentRemoved(localId) + if (pendingAttachment.attachment) { + try { + await api.conversations.attachments.delete(values.conversationId, pendingAttachment.attachment.id) + } catch (error) { + posthog.captureException(error) + lemonToast.error('Failed to delete image attachment.') + } + } + }, + clearPendingAttachments: () => { + for (const attachment of values.pendingAttachments) { + cache.disposables.dispose(`attachment-preview-${attachment.localId}`) + } + actions.pendingAttachmentsCleared() + }, + askMax: async ({ prompt, addToThread = true, uiContext, attachments: retryAttachments }, breakpoint) => { // Only process if this thread is the currently active one if (values.conversationId !== values.activeThreadKey) { return @@ -1059,20 +1201,45 @@ export const maxThreadLogic = kea([ values.billingContext && values.featureFlags[FEATURE_FLAGS.MAX_BILLING_CONTEXT] ? values.billingContext : undefined + const attachments = + retryAttachments ?? + values.pendingAttachments.flatMap((item) => + item.status === 'ready' && item.attachment ? [item.attachment] : [] + ) + + if (values.pendingAttachments.some((item) => item.status === 'uploading')) { + actions.setAttachmentInputError('Wait for images to finish uploading.') + return + } + if (values.pendingAttachments.some((item) => item.status === 'error')) { + actions.setAttachmentInputError('Remove failed images before sending.') + return + } + if (attachments.length > 0 && values.agentMode === AgentMode.Research) { + actions.setAttachmentInputError('Images are not supported in Research mode.') + return + } + if (attachments.length > 0 && values.isSandboxMode) { + actions.setAttachmentInputError('Images are not supported in Sandbox mode.') + return + } if ( values.queueingEnabled && values.threadLoading && addToThread && - typeof prompt === 'string' && - prompt.trim() !== '' + ((typeof prompt === 'string' && prompt.trim() !== '') || attachments.length > 0) ) { + if (attachments.length > 0) { + actions.setAttachmentInputError('Images cannot be added to queued messages.') + return + } if (values.queueIsFull) { lemonToast.error('You can only queue two messages at a time.') return } actions.enqueueQueuedMessage({ - content: prompt, + content: prompt ?? '', contextualTools, uiContext: mergedUiContext, billingContext, @@ -1129,6 +1296,7 @@ export const maxThreadLogic = kea([ // Clear the question actions.setQuestion('') + actions.clearPendingAttachments() // Drop #panel=max:… options so reload doesn't re-run auto-send from the hash if (props.panelId === SIDE_PANEL_PANEL_ID && sidePanelStateLogic.isMounted()) { sidePanelStateLogic.actions.setSidePanelOptions(null) @@ -1159,6 +1327,7 @@ export const maxThreadLogic = kea([ conversation: values.conversation?.id || values.conversationId, // Include auto-rejection payload if there was a pending approval resume_payload: autoRejectPayload, + attachments, }, 0, addToThread @@ -1211,7 +1380,7 @@ export const maxThreadLogic = kea([ retryLastMessage: () => { const lastMessage = values.threadRaw.filter(isHumanMessage).pop() as HumanMessage | undefined if (lastMessage) { - actions.askMax(lastMessage.content) + actions.askMax(lastMessage.content, true, undefined, lastMessage.attachments) } }, @@ -1720,14 +1889,44 @@ export const maxThreadLogic = kea([ ], submissionDisabledReason: [ - (s) => [s.contextDisabledReason, s.question, s.queueDisabledReason], - (contextDisabledReason, question, queueDisabledReason): string | undefined => { + (s) => [ + s.contextDisabledReason, + s.question, + s.queueDisabledReason, + s.pendingAttachments, + s.agentMode, + s.isSandboxMode, + ], + ( + contextDisabledReason, + question, + queueDisabledReason, + pendingAttachments, + agentMode, + isSandboxMode + ): string | undefined => { // Context-related reasons take precedence (form pending, streaming, etc.) if (contextDisabledReason) { return contextDisabledReason } - if (!question) { + if (pendingAttachments.some((attachment) => attachment.status === 'uploading')) { + return 'Wait for images to finish uploading' + } + + if (pendingAttachments.some((attachment) => attachment.status === 'error')) { + return 'Remove failed images before sending' + } + + if (pendingAttachments.length > 0 && agentMode === AgentMode.Research) { + return 'Images are not supported in Research mode' + } + + if (pendingAttachments.length > 0 && isSandboxMode) { + return 'Images are not supported in Sandbox mode' + } + + if (!question && pendingAttachments.length === 0) { return 'I need some input first' } diff --git a/posthog/schema.py b/posthog/schema.py index 53ba1d2bf98e..9541544f6b2a 100644 --- a/posthog/schema.py +++ b/posthog/schema.py @@ -56,6 +56,7 @@ ChartDisplayType as ChartDisplayType, ColorMode as ColorMode, Compare as Compare, + ContentType as ContentType, ConversionRateInputType as ConversionRateInputType, CoreEventCategory as CoreEventCategory, CorrelationType as CorrelationType, @@ -1501,6 +1502,18 @@ class HogQueryResponse(BaseModel): stdout: str | None = None +class HumanMessageAttachment(BaseModel): + model_config = ConfigDict( + extra="forbid", + ) + content_type: ContentType + file_name: str + height: float + id: str + size: float + width: float + + class InsightsThresholdBounds(BaseModel): model_config = ConfigDict( extra="forbid", @@ -25499,6 +25512,7 @@ class HumanMessage(BaseModel): model_config = ConfigDict( extra="forbid", ) + attachments: list[HumanMessageAttachment] | None = None content: str id: str | None = None parent_tool_call_id: str | None = None diff --git a/posthog/schema_enums.py b/posthog/schema_enums.py index 4b06c10d5f6b..0ab587b1b91e 100644 --- a/posthog/schema_enums.py +++ b/posthog/schema_enums.py @@ -2082,6 +2082,11 @@ class SessionsV2JoinMode(StrEnum): UUID = "uuid" +class ContentType(StrEnum): + IMAGE_PNG = "image/png" + IMAGE_JPEG = "image/jpeg" + + class InfinityValue(float, Enum): NUMBER_999999 = 999999 NUMBER__999999 = -999999 diff --git a/products/conversations/frontend/generated/api.schemas.ts b/products/conversations/frontend/generated/api.schemas.ts index 6f2d990ac0b4..a15775a7db4d 100644 --- a/products/conversations/frontend/generated/api.schemas.ts +++ b/products/conversations/frontend/generated/api.schemas.ts @@ -220,6 +220,11 @@ export interface MessageApi { agent_mode?: AgentModeEnumApi is_sandbox?: boolean resume_payload?: unknown + /** + * IDs of private PNG or JPEG attachments uploaded for this exact conversation. + * @maxItems 4 + */ + attachments?: string[] } export type ConversationApiMessagesItem = { [key: string]: unknown } @@ -287,6 +292,40 @@ export interface MessageMinimalApi { content: string } +export interface ConversationAttachmentUploadApi { + /** A PNG or JPEG image no larger than 4 MiB. */ + image: string +} + +/** + * * `image/png` - image/png + * * `image/jpeg` - image/jpeg + */ +export type ContentTypeEnumApi = (typeof ContentTypeEnumApi)[keyof typeof ContentTypeEnumApi] + +export const ContentTypeEnumApi = { + ImagePng: 'image/png', + ImageJpeg: 'image/jpeg', +} as const + +export interface ConversationAttachmentApi { + /** Attachment identifier. */ + readonly id: string + /** Sanitized display filename. */ + readonly file_name: string + /** Verified and re-encoded image content type. + * + * * `image/png` - image/png + * * `image/jpeg` - image/jpeg */ + readonly content_type: ContentTypeEnumApi + /** Sanitized image size in bytes. */ + readonly size: number + /** Decoded image width in pixels. */ + readonly width: number + /** Decoded image height in pixels. */ + readonly height: number +} + export type PatchedConversationApiMessagesItem = { [key: string]: unknown } export type PatchedConversationApiPendingApprovalsItem = { [key: string]: unknown } diff --git a/products/conversations/frontend/generated/api.ts b/products/conversations/frontend/generated/api.ts index 041bb8952868..7fd774507349 100644 --- a/products/conversations/frontend/generated/api.ts +++ b/products/conversations/frontend/generated/api.ts @@ -16,6 +16,8 @@ import type { ComposeTicketApi, ComposeTicketResponseApi, ConversationApi, + ConversationAttachmentApi, + ConversationAttachmentUploadApi, ConversationsListParams, ConversationsTicketsListParams, ConversationsTicketsMessagesListParams, @@ -158,6 +160,75 @@ export const conversationsAppendMessageCreate = async ( }) } +export const getConversationsAttachmentsCreateUrl = (projectId: string, conversation: string) => { + return `/api/projects/${projectId}/conversations/${conversation}/attachments/` +} + +/** + * Upload and sanitize a private PNG or JPEG attachment for this conversation. + */ +export const conversationsAttachmentsCreate = async ( + projectId: string, + conversation: string, + conversationAttachmentUploadApi: ConversationAttachmentUploadApi, + options?: RequestInit +): Promise => { + const formData = new FormData() + formData.append(`image`, conversationAttachmentUploadApi.image) + + return apiMutator(getConversationsAttachmentsCreateUrl(projectId, conversation), { + ...options, + method: 'POST', + body: formData, + }) +} + +export const getConversationsAttachmentsDestroyUrl = ( + projectId: string, + conversation: string, + attachmentId: string +) => { + return `/api/projects/${projectId}/conversations/${conversation}/attachments/${attachmentId}/` +} + +/** + * Delete a private conversation attachment. + */ +export const conversationsAttachmentsDestroy = async ( + projectId: string, + conversation: string, + attachmentId: string, + options?: RequestInit +): Promise => { + return apiMutator(getConversationsAttachmentsDestroyUrl(projectId, conversation, attachmentId), { + ...options, + method: 'DELETE', + }) +} + +export const getConversationsAttachmentsContentRetrieveUrl = ( + projectId: string, + conversation: string, + attachmentId: string +) => { + return `/api/projects/${projectId}/conversations/${conversation}/attachments/${attachmentId}/content/` +} + +/** + * Read private attachment content scoped to the authenticated user, team, and conversation. + */ +export const conversationsAttachmentsContentRetrieve = async ( + projectId: string, + conversation: string, + attachmentId: string, + options?: RequestInit +): Promise => { + return apiMutator(getConversationsAttachmentsContentRetrieveUrl(projectId, conversation, attachmentId), { + ...options, + method: 'GET', + }) +} + export const getConversationsCancelPartialUpdateUrl = (projectId: string, conversation: string) => { return `/api/projects/${projectId}/conversations/${conversation}/cancel/` } diff --git a/products/conversations/frontend/generated/api.zod.ts b/products/conversations/frontend/generated/api.zod.ts index a59774a53abe..b1b93f468668 100644 --- a/products/conversations/frontend/generated/api.zod.ts +++ b/products/conversations/frontend/generated/api.zod.ts @@ -18,6 +18,7 @@ import * as zod from 'zod' export const conversationsCreateBodyContentMax = 40000 export const conversationsCreateBodyIsSandboxDefault = false +export const conversationsCreateBodyAttachmentsMax = 4 export const ConversationsCreateBody = /* @__PURE__ */ zod .object({ @@ -50,6 +51,11 @@ export const ConversationsCreateBody = /* @__PURE__ */ zod ), is_sandbox: zod.boolean().default(conversationsCreateBodyIsSandboxDefault), resume_payload: zod.unknown().optional(), + attachments: zod + .array(zod.uuid()) + .max(conversationsCreateBodyAttachmentsMax) + .optional() + .describe('IDs of private PNG or JPEG attachments uploaded for this exact conversation.'), }) .describe('Serializer for appending a message to an existing conversation without triggering AI processing.') @@ -66,6 +72,13 @@ export const ConversationsAppendMessageCreateBody = /* @__PURE__ */ zod }) .describe('Serializer for appending a message to an existing conversation without triggering AI processing.') +/** + * Upload and sanitize a private PNG or JPEG attachment for this conversation. + */ +export const ConversationsAttachmentsCreateBody = /* @__PURE__ */ zod.object({ + image: zod.url().describe('A PNG or JPEG image no larger than 4 MiB.'), +}) + export const ConversationsCancelPartialUpdateBody = /* @__PURE__ */ zod.looseObject({}) export const ConversationsQueueCreateBody = /* @__PURE__ */ zod.looseObject({}) diff --git a/products/posthog_ai/backend/attachments.py b/products/posthog_ai/backend/attachments.py new file mode 100644 index 000000000000..746b43dd2fed --- /dev/null +++ b/products/posthog_ai/backend/attachments.py @@ -0,0 +1,212 @@ +import re +import base64 +from collections.abc import Sequence +from dataclasses import dataclass +from io import BytesIO +from pathlib import Path +from typing import TYPE_CHECKING, Any +from uuid import UUID, uuid4 + +from django.conf import settings +from django.core.files.uploadedfile import UploadedFile + +from PIL import Image, ImageOps, UnidentifiedImageError + +from posthog.models.user import User +from posthog.storage import object_storage +from posthog.sync import database_sync_to_async + +from products.posthog_ai.backend.models import ConversationAttachment + +if TYPE_CHECKING: + from posthog.schema import HumanMessage + +MAX_ATTACHMENT_BYTES = 4 * 1024 * 1024 +MAX_ATTACHMENTS_PER_MESSAGE = 4 +MAX_DECODED_PIXELS = 25_000_000 +ALLOWED_CONTENT_TYPES = frozenset({"image/png", "image/jpeg"}) +_FORMAT_TO_CONTENT_TYPE = {"PNG": "image/png", "JPEG": "image/jpeg"} +_FILENAME_UNSAFE = re.compile(r"[^A-Za-z0-9._ -]+") + + +class InvalidConversationAttachment(ValueError): + pass + + +@dataclass(frozen=True) +class ProcessedConversationAttachment: + content: bytes + content_type: str + file_name: str + width: int + height: int + + @property + def size(self) -> int: + return len(self.content) + + +def sanitize_attachment_filename(file_name: str, content_type: str) -> str: + base_name = Path(file_name.replace("\\", "/")).name + stem = Path(base_name).stem + safe_stem = _FILENAME_UNSAFE.sub("_", stem).strip(" ._")[:200] or "image" + extension = ".png" if content_type == "image/png" else ".jpg" + return f"{safe_stem}{extension}" + + +def process_conversation_attachment(upload: UploadedFile) -> ProcessedConversationAttachment: + claimed_content_type = (upload.content_type or "").lower() + if claimed_content_type == "image/jpg": + claimed_content_type = "image/jpeg" + if claimed_content_type not in ALLOWED_CONTENT_TYPES: + raise InvalidConversationAttachment("Only PNG and JPEG images are supported.") + if upload.size > MAX_ATTACHMENT_BYTES: + raise InvalidConversationAttachment("Images must be 4 MiB or smaller.") + + source = upload.read(MAX_ATTACHMENT_BYTES + 1) + if len(source) > MAX_ATTACHMENT_BYTES: + raise InvalidConversationAttachment("Images must be 4 MiB or smaller.") + + try: + with Image.open(BytesIO(source)) as image: + detected_content_type = _FORMAT_TO_CONTENT_TYPE.get(image.format or "") + width, height = image.size + if detected_content_type is None or detected_content_type != claimed_content_type: + raise InvalidConversationAttachment("The file contents do not match its PNG or JPEG type.") + if width <= 0 or height <= 0 or width * height > MAX_DECODED_PIXELS: + raise InvalidConversationAttachment("The decoded image is too large.") + image.verify() + + with Image.open(BytesIO(source)) as image: + image.load() + image = ImageOps.exif_transpose(image) + output = BytesIO() + if detected_content_type == "image/jpeg": + clean_image = image.convert("RGB") + clean_image.save(output, format="JPEG", quality=90, optimize=True) + else: + clean_image = image.convert("RGBA" if "A" in image.getbands() else "RGB") + clean_image.save(output, format="PNG", optimize=True) + except InvalidConversationAttachment: + raise + except (Image.DecompressionBombError, UnidentifiedImageError, OSError, ValueError) as error: + raise InvalidConversationAttachment("The uploaded file is not a valid PNG or JPEG image.") from error + + content = output.getvalue() + if len(content) > MAX_ATTACHMENT_BYTES: + raise InvalidConversationAttachment("The sanitized image exceeds the 4 MiB limit.") + + return ProcessedConversationAttachment( + content=content, + content_type=detected_content_type, + file_name=sanitize_attachment_filename(upload.name, detected_content_type), + width=width, + height=height, + ) + + +def save_conversation_attachment( + *, + processed: ProcessedConversationAttachment, + team_id: int, + creator: User, + conversation_id: UUID, +) -> ConversationAttachment: + if not settings.OBJECT_STORAGE_ENABLED: + raise InvalidConversationAttachment("Private object storage is required for image attachments.") + + extension = "png" if processed.content_type == "image/png" else "jpg" + object_path = f"posthog_ai/conversation_attachments/{uuid4().hex}.{extension}" + object_storage.write( + object_path, + processed.content, + extras={"ContentType": processed.content_type, "CacheControl": "private, max-age=3600"}, + ) + try: + return ConversationAttachment.objects.create( + team_id=team_id, + creator=creator, + conversation_id=conversation_id, + file_name=processed.file_name, + content_type=processed.content_type, + size=processed.size, + width=processed.width, + height=processed.height, + object_path=object_path, + ) + except Exception: + object_storage.delete(object_path) + raise + + +def attachment_to_ref(attachment: ConversationAttachment) -> dict[str, str | int]: + return { + "id": str(attachment.id), + "file_name": attachment.file_name, + "content_type": attachment.content_type, + "size": attachment.size, + "width": attachment.width, + "height": attachment.height, + } + + +@database_sync_to_async +def _load_scoped_attachments( + *, + attachment_ids: set[UUID], + team_id: int, + creator_id: int, + conversation_id: UUID, +) -> dict[UUID, ConversationAttachment]: + attachments = ConversationAttachment.objects.for_team(team_id).filter( + id__in=attachment_ids, + team_id=team_id, + creator_id=creator_id, + conversation_id=conversation_id, + ) + return {attachment.id: attachment for attachment in attachments} + + +async def load_attachment_blocks( + messages: Sequence["HumanMessage"], + *, + team_id: int, + creator_id: int, + conversation_id: UUID, +) -> dict[int, list[dict[str, Any]]]: + referenced_ids = {UUID(attachment.id) for message in messages for attachment in (message.attachments or [])} + if not referenced_ids: + return {} + + attachments_by_id = await _load_scoped_attachments( + attachment_ids=referenced_ids, + team_id=team_id, + creator_id=creator_id, + conversation_id=conversation_id, + ) + if len(attachments_by_id) != len(referenced_ids): + raise InvalidConversationAttachment("One or more conversation attachments are unavailable.") + + blocks_by_message: dict[int, list[dict[str, Any]]] = {} + for message in messages: + blocks: list[dict[str, Any]] = [] + for reference in message.attachments or []: + attachment = attachments_by_id.get(UUID(reference.id)) + if attachment is None: + raise InvalidConversationAttachment("Conversation attachment is unavailable.") + content = await database_sync_to_async(object_storage.read_bytes)(attachment.object_path) + if content is None: + raise InvalidConversationAttachment("Conversation attachment content is unavailable.") + blocks.append( + { + "type": "image", + "source": { + "type": "base64", + "media_type": attachment.content_type, + "data": base64.b64encode(content).decode("ascii"), + }, + } + ) + if blocks: + blocks_by_message[id(message)] = blocks + return blocks_by_message diff --git a/products/posthog_ai/backend/migrations/0004_conversation_attachments.py b/products/posthog_ai/backend/migrations/0004_conversation_attachments.py new file mode 100644 index 000000000000..7ce496406777 --- /dev/null +++ b/products/posthog_ai/backend/migrations/0004_conversation_attachments.py @@ -0,0 +1,61 @@ +# Generated by Django 5.2.14 on 2026-06-19 05:08 + +import django.db.models.deletion +import posthog.models.utils +from django.conf import settings +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("posthog", "1223_remove_ducklakecatalog_cross_account_fields"), + ("posthog_ai", "0003_conversation_topic"), + migrations.swappable_dependency(settings.AUTH_USER_MODEL), + ] + + operations = [ + migrations.CreateModel( + name="ConversationAttachment", + fields=[ + ( + "id", + models.UUIDField( + default=posthog.models.utils.UUIDT, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("conversation_id", models.UUIDField()), + ("file_name", models.CharField(max_length=255)), + ("content_type", models.CharField(max_length=32)), + ("size", models.PositiveIntegerField()), + ("width", models.PositiveIntegerField()), + ("height", models.PositiveIntegerField()), + ("object_path", models.CharField(max_length=255, unique=True)), + ( + "creator", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + to=settings.AUTH_USER_MODEL, + ), + ), + ( + "team", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="posthog.team" + ), + ), + ], + options={ + "db_table": "ee_conversationattachment", + "indexes": [ + models.Index( + fields=["team", "creator", "conversation_id"], + name="ee_conversa_team_id_c97e8e_idx", + ) + ], + }, + ), + ] diff --git a/products/posthog_ai/backend/migrations/max_migration.txt b/products/posthog_ai/backend/migrations/max_migration.txt index 5e8f520ca572..11851b33f6bb 100644 --- a/products/posthog_ai/backend/migrations/max_migration.txt +++ b/products/posthog_ai/backend/migrations/max_migration.txt @@ -1 +1 @@ -0003_conversation_topic +0004_conversation_attachments diff --git a/products/posthog_ai/backend/models/__init__.py b/products/posthog_ai/backend/models/__init__.py index b6f06a0fa916..381137c31cc8 100644 --- a/products/posthog_ai/backend/models/__init__.py +++ b/products/posthog_ai/backend/models/__init__.py @@ -2,6 +2,7 @@ from .assistant import ( AgentArtifact, Conversation, + ConversationAttachment, ConversationCheckpoint, ConversationCheckpointBlob, ConversationCheckpointWrite, @@ -12,6 +13,7 @@ "AgentArtifact", "AgentMemory", "Conversation", + "ConversationAttachment", "ConversationCheckpoint", "ConversationCheckpointBlob", "ConversationCheckpointWrite", diff --git a/products/posthog_ai/backend/models/assistant.py b/products/posthog_ai/backend/models/assistant.py index 484b57bf768f..5f7e69a44c20 100644 --- a/products/posthog_ai/backend/models/assistant.py +++ b/products/posthog_ai/backend/models/assistant.py @@ -5,6 +5,7 @@ from django.db import IntegrityError, models from django.utils import timezone +from posthog.models.scoping.root_mixin import TeamScopedRootMixin from posthog.models.team.team import Team from posthog.models.user import User from posthog.models.utils import CreatedMetaFields, DeletedMetaFields, UpdatedMetaFields, UUIDModel, UUIDTModel @@ -118,6 +119,24 @@ class Topic(models.TextChoices): ) +class ConversationAttachment(TeamScopedRootMixin, UUIDTModel): + team = models.ForeignKey(Team, on_delete=models.CASCADE) + creator = models.ForeignKey(User, on_delete=models.CASCADE) + conversation_id = models.UUIDField() + file_name = models.CharField(max_length=255) + content_type = models.CharField(max_length=32) + size = models.PositiveIntegerField() + width = models.PositiveIntegerField() + height = models.PositiveIntegerField() + object_path = models.CharField(max_length=255, unique=True) + + class Meta: + db_table = "ee_conversationattachment" + indexes = [ + models.Index(fields=["team", "creator", "conversation_id"]), + ] + + class ConversationCheckpoint(UUIDTModel): thread = models.ForeignKey(Conversation, on_delete=models.CASCADE, related_name="checkpoints") checkpoint_ns = models.TextField( diff --git a/products/posthog_ai/backend/tests/test_attachments.py b/products/posthog_ai/backend/tests/test_attachments.py new file mode 100644 index 000000000000..e292df2bd9e8 --- /dev/null +++ b/products/posthog_ai/backend/tests/test_attachments.py @@ -0,0 +1,94 @@ +import base64 +from uuid import uuid4 + +from posthog.test.base import BaseTest +from unittest.mock import patch + +from posthog.schema import HumanMessage, HumanMessageAttachment + +from posthog.sync import database_sync_to_async + +from products.posthog_ai.backend.attachments import InvalidConversationAttachment, load_attachment_blocks +from products.posthog_ai.backend.models import ConversationAttachment + + +class TestConversationAttachmentHydration(BaseTest): + @database_sync_to_async + def _create_attachment(self, conversation_id): + return ConversationAttachment.objects.for_team(self.team.id).create( + team=self.team, + creator=self.user, + conversation_id=conversation_id, + file_name="image.png", + content_type="image/png", + size=3, + width=1, + height=1, + object_path=f"private/{conversation_id}", + ) + + async def test_loads_scoped_attachment_as_transient_anthropic_block(self): + conversation_id = uuid4() + attachment = await self._create_attachment(conversation_id) + message = HumanMessage( + content="What is shown?", + attachments=[ + HumanMessageAttachment( + id=str(attachment.id), + file_name="client-name-is-ignored.png", + content_type="image/jpeg", + size=999, + width=999, + height=999, + ) + ], + ) + with patch("products.posthog_ai.backend.attachments.object_storage.read_bytes", return_value=b"png"): + blocks = await load_attachment_blocks( + [message], + team_id=self.team.id, + creator_id=self.user.id, + conversation_id=conversation_id, + ) + self.assertEqual( + blocks[id(message)], + [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": base64.b64encode(b"png").decode("ascii"), + }, + } + ], + ) + + async def test_rejects_user_and_conversation_mismatches(self): + conversation_id = uuid4() + attachment = await self._create_attachment(conversation_id) + message = HumanMessage( + content="", + attachments=[ + HumanMessageAttachment( + id=str(attachment.id), + file_name=attachment.file_name, + content_type="image/png", + size=attachment.size, + width=attachment.width, + height=attachment.height, + ) + ], + ) + for creator_id, scoped_conversation_id in [ + (self.user.id + 1, conversation_id), + (self.user.id, uuid4()), + ]: + with self.subTest(creator_id=creator_id, conversation_id=scoped_conversation_id): + with self.assertRaises(InvalidConversationAttachment): + await load_attachment_blocks( + [message], + team_id=self.team.id, + creator_id=creator_id, + conversation_id=scoped_conversation_id, + ) diff --git a/services/mcp/src/api/generated.ts b/services/mcp/src/api/generated.ts index d47789fe0433..192128c7c39b 100644 --- a/services/mcp/src/api/generated.ts +++ b/services/mcp/src/api/generated.ts @@ -12672,6 +12672,18 @@ export namespace Schemas { Base64: 'base64', } as const; + /** + * * `image/png` - image/png + * * `image/jpeg` - image/jpeg + */ + export type ContentTypeEnum = typeof ContentTypeEnum[keyof typeof ContentTypeEnum]; + + + export const ContentTypeEnum = { + ImagePng: 'image/png', + ImageJpeg: 'image/jpeg', + } as const; + export interface ContextGeneration { /** * ID of the Task currently generating this folder's CONTEXT.md, or null if none. @@ -12801,6 +12813,29 @@ export namespace Schemas { readonly pending_approvals: readonly ConversationPendingApprovalsItem[]; } + export interface ConversationAttachment { + /** Attachment identifier. */ + readonly id: string; + /** Sanitized display filename. */ + readonly file_name: string; + /** Verified and re-encoded image content type. + * + * * `image/png` - image/png + * * `image/jpeg` - image/jpeg */ + readonly content_type: ContentTypeEnum; + /** Sanitized image size in bytes. */ + readonly size: number; + /** Decoded image width in pixels. */ + readonly width: number; + /** Decoded image height in pixels. */ + readonly height: number; + } + + export interface ConversationAttachmentUpload { + /** A PNG or JPEG image no larger than 4 MiB. */ + image: string; + } + export interface ConversationMinimal { readonly id: string; readonly status: ConversationStatus; @@ -26718,6 +26753,11 @@ export namespace Schemas { agent_mode?: AgentModeEnum; is_sandbox?: boolean; resume_payload?: unknown; + /** + * IDs of private PNG or JPEG attachments uploaded for this exact conversation. + * @maxItems 4 + */ + attachments?: string[]; } export interface MessageCategory {