diff --git a/crates/multimodal/src/lib.rs b/crates/multimodal/src/lib.rs index 31baf8e254..0994c7462d 100644 --- a/crates/multimodal/src/lib.rs +++ b/crates/multimodal/src/lib.rs @@ -11,7 +11,7 @@ pub use media::{ImageFetchConfig, MediaConnector, MediaConnectorConfig, MediaSou pub use registry::{ModelMetadata, ModelProcessorSpec, ModelRegistry}; pub use tracker::{AsyncMultiModalTracker, TrackerOutput}; pub use types::{ - ChatContentPart, FieldLayout, ImageDetail, ImageFrame, ImageSize, ImageSource, Modality, + FieldLayout, ImageDetail, ImageFrame, ImageSize, ImageSource, MediaContentPart, Modality, MultiModalData, MultiModalUUIDs, PlaceholderRange, PromptReplacement, TokenId, TrackedMedia, }; // Re-export vision processing components diff --git a/crates/multimodal/src/tracker.rs b/crates/multimodal/src/tracker.rs index 0d1c1eec1e..b44e7e90aa 100644 --- a/crates/multimodal/src/tracker.rs +++ b/crates/multimodal/src/tracker.rs @@ -6,7 +6,7 @@ use super::{ error::{MultiModalError, MultiModalResult}, media::{ImageFetchConfig, MediaConnector, MediaSource}, types::{ - ChatContentPart, ImageDetail, Modality, MultiModalData, MultiModalUUIDs, TrackedMedia, + ImageDetail, MediaContentPart, Modality, MultiModalData, MultiModalUUIDs, TrackedMedia, }, }; @@ -33,17 +33,17 @@ impl AsyncMultiModalTracker { } } - pub fn push_part(&mut self, part: ChatContentPart) -> MultiModalResult<()> { + pub fn push_part(&mut self, part: MediaContentPart) -> MultiModalResult<()> { match part { - ChatContentPart::Text { .. } => {} - ChatContentPart::ImageUrl { url, detail, uuid } => { + MediaContentPart::Text { .. } => {} + MediaContentPart::ImageUrl { url, detail, uuid } => { let source = match url::Url::parse(&url) { Ok(parsed) if parsed.scheme() == "data" => MediaSource::DataUrl(url), _ => MediaSource::Url(url), }; self.enqueue_image(source, detail.unwrap_or_default(), uuid); } - ChatContentPart::ImageData { + MediaContentPart::ImageData { data, mime_type: _, uuid, @@ -55,7 +55,7 @@ impl AsyncMultiModalTracker { uuid, ); } - ChatContentPart::ImageEmbeds { .. } => { + MediaContentPart::ImageEmbeds { .. } => { return Err(MultiModalError::UnsupportedContent("image_embeds")); } } diff --git a/crates/multimodal/src/types.rs b/crates/multimodal/src/types.rs index b4f4831a01..9ecd93a434 100644 --- a/crates/multimodal/src/types.rs +++ b/crates/multimodal/src/types.rs @@ -38,7 +38,7 @@ pub enum ImageDetail { /// A normalized content part understood by the tracker. #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] -pub enum ChatContentPart { +pub enum MediaContentPart { Text { text: String, }, diff --git a/crates/multimodal/tests/multimodal_tracker_test.rs b/crates/multimodal/tests/multimodal_tracker_test.rs index f2625edcd0..dea2b4f256 100644 --- a/crates/multimodal/tests/multimodal_tracker_test.rs +++ b/crates/multimodal/tests/multimodal_tracker_test.rs @@ -2,8 +2,8 @@ use std::{path::PathBuf, sync::Arc, time::Duration}; use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine}; use llm_multimodal::{ - AsyncMultiModalTracker, ChatContentPart, ImageFetchConfig, ImageSource, MediaConnector, - MediaConnectorConfig, MediaSource, Modality, + AsyncMultiModalTracker, ImageFetchConfig, ImageSource, MediaConnector, MediaConnectorConfig, + MediaContentPart, MediaSource, Modality, }; use reqwest::Client; use tempfile::tempdir; @@ -104,12 +104,12 @@ async fn tracker_fetches_images_and_records_uuids() { let mut tracker = AsyncMultiModalTracker::new(connector); tracker - .push_part(ChatContentPart::Text { + .push_part(MediaContentPart::Text { text: "before".into(), }) .expect("text part"); tracker - .push_part(ChatContentPart::ImageData { + .push_part(MediaContentPart::ImageData { data: tiny_png_bytes(), mime_type: Some("image/png".into()), uuid: Some("img-1".into()), @@ -117,7 +117,7 @@ async fn tracker_fetches_images_and_records_uuids() { }) .expect("image part"); tracker - .push_part(ChatContentPart::Text { + .push_part(MediaContentPart::Text { text: "after".into(), }) .expect("text part"); diff --git a/model_gateway/src/routers/grpc/multimodal.rs b/model_gateway/src/routers/grpc/multimodal.rs index 2c6271b8e9..008d4225dc 100644 --- a/model_gateway/src/routers/grpc/multimodal.rs +++ b/model_gateway/src/routers/grpc/multimodal.rs @@ -1,23 +1,29 @@ -//! Multimodal processing integration for gRPC chat pipeline. +//! Multimodal processing integration for gRPC pipeline (chat + messages). //! //! This module bridges the `llm-multimodal` crate with the gRPC router pipeline, //! handling the full processing chain: extract content parts → fetch images → //! preprocess pixels → expand placeholder tokens → build proto MultimodalInputs. +//! +//! Both the chat completion pipeline and the Messages API pipeline share the same +//! processing core (`process_multimodal_parts`). Only the detection and extraction +//! functions differ because they work with different input types (`ChatMessage` vs +//! `InputMessage`). use std::{collections::HashMap, path::Path, sync::Arc}; use anyhow::{Context, Result}; use dashmap::DashMap; use llm_multimodal::{ - AsyncMultiModalTracker, ChatContentPart, FieldLayout, ImageDetail, ImageFrame, - ImageProcessorRegistry, MediaConnector, MediaConnectorConfig, Modality, ModelMetadata, - ModelRegistry, ModelSpecificValue, PlaceholderRange, PreProcessorConfig, PreprocessedImages, + AsyncMultiModalTracker, FieldLayout, ImageDetail, ImageFrame, ImageProcessorRegistry, + MediaConnector, MediaConnectorConfig, MediaContentPart, Modality, ModelMetadata, ModelRegistry, + ModelSpecificValue, PlaceholderRange, PreProcessorConfig, PreprocessedImages, PromptReplacement, TrackedMedia, TrackerOutput, }; use llm_tokenizer::TokenizerTrait; use openai_protocol::{ chat::{ChatMessage, MessageContent}, common::ContentPart, + messages::{ImageSource, InputContent, InputContentBlock, InputMessage, Role}, }; use tracing::{debug, warn}; @@ -165,8 +171,8 @@ pub(crate) fn has_multimodal_content(messages: &[ChatMessage]) -> bool { } /// Extract multimodal content parts from OpenAI chat messages, -/// converting protocol `ContentPart` to multimodal crate `ChatContentPart`. -fn extract_content_parts(messages: &[ChatMessage]) -> Vec { +/// converting protocol `ContentPart` to multimodal crate `MediaContentPart`. +fn extract_content_parts(messages: &[ChatMessage]) -> Vec { let mut parts = Vec::new(); for msg in messages { @@ -182,14 +188,14 @@ fn extract_content_parts(messages: &[ChatMessage]) -> Vec { match part { ContentPart::ImageUrl { image_url } => { let detail = image_url.detail.as_deref().and_then(parse_detail); - parts.push(ChatContentPart::ImageUrl { + parts.push(MediaContentPart::ImageUrl { url: image_url.url.clone(), detail, uuid: None, }); } ContentPart::Text { text } => { - parts.push(ChatContentPart::Text { text: text.clone() }); + parts.push(MediaContentPart::Text { text: text.clone() }); } ContentPart::VideoUrl { .. } => {} // Skip VideoUrl for now } @@ -210,6 +216,96 @@ fn parse_detail(detail: &str) -> Option { } } +// --------------------------------------------------------------------------- +// Messages API multimodal detection and extraction +// --------------------------------------------------------------------------- + +/// Check if any messages in a Messages API request contain multimodal content. +pub(crate) fn has_multimodal_content_messages(messages: &[InputMessage]) -> bool { + messages.iter().any(|msg| { + if msg.role != Role::User { + return false; + } + match &msg.content { + InputContent::Blocks(blocks) => blocks + .iter() + .any(|block| matches!(block, InputContentBlock::Image(_))), + InputContent::String(_) => false, + } + }) +} + +/// Extract multimodal content parts from Messages API input messages, +/// converting `InputContentBlock::Image` to multimodal crate `MediaContentPart`. +fn extract_content_parts_messages(messages: &[InputMessage]) -> Vec { + let mut parts = Vec::new(); + + for msg in messages { + if msg.role != Role::User { + continue; + } + let blocks = match &msg.content { + InputContent::Blocks(blocks) => blocks, + InputContent::String(_) => continue, + }; + + for block in blocks { + match block { + InputContentBlock::Image(image_block) => match &image_block.source { + ImageSource::Base64 { media_type, data } => { + // Convert base64 to data URL for the media connector + let data_url = format!("data:{media_type};base64,{data}"); + parts.push(MediaContentPart::ImageUrl { + url: data_url, + detail: None, + uuid: None, + }); + } + ImageSource::Url { url } => { + parts.push(MediaContentPart::ImageUrl { + url: url.clone(), + detail: None, + uuid: None, + }); + } + }, + InputContentBlock::Text(text_block) => { + parts.push(MediaContentPart::Text { + text: text_block.text.clone(), + }); + } + _ => {} + } + } + } + + parts +} + +/// Process multimodal content from Messages API input messages. +/// +/// Entry point for the messages preparation stage. Extracts image content parts +/// from `InputMessage`, then delegates to the shared processing core. +pub(crate) async fn process_multimodal_messages( + messages: &[InputMessage], + model_id: &str, + tokenizer: &dyn TokenizerTrait, + token_ids: Vec, + components: &MultimodalComponents, + tokenizer_source: &str, +) -> Result { + let content_parts = extract_content_parts_messages(messages); + process_multimodal_parts( + content_parts, + model_id, + tokenizer, + token_ids, + components, + tokenizer_source, + ) + .await +} + /// Process multimodal content: fetch images, preprocess pixels, expand tokens, collect hashes. /// /// Single entry point called from preparation.rs. Handles the full pipeline: @@ -222,8 +318,30 @@ pub(crate) async fn process_multimodal( components: &MultimodalComponents, tokenizer_source: &str, ) -> Result { - // Step 1: Fetch images let content_parts = extract_content_parts(messages); + process_multimodal_parts( + content_parts, + model_id, + tokenizer, + token_ids, + components, + tokenizer_source, + ) + .await +} + +/// Shared multimodal processing core. +/// +/// Takes pre-extracted `MediaContentPart`s (from either chat or messages pipeline) +/// and runs the full processing chain: fetch → preprocess → expand → build intermediate. +async fn process_multimodal_parts( + content_parts: Vec, + model_id: &str, + tokenizer: &dyn TokenizerTrait, + token_ids: Vec, + components: &MultimodalComponents, + tokenizer_source: &str, +) -> Result { let mut tracker = AsyncMultiModalTracker::new(components.media_connector.clone()); for part in content_parts { @@ -674,12 +792,12 @@ mod tests { assert_eq!(parts.len(), 2); match &parts[0] { - ChatContentPart::Text { text } => assert_eq!(text, "Describe this:"), + MediaContentPart::Text { text } => assert_eq!(text, "Describe this:"), _ => panic!("Expected Text part"), } match &parts[1] { - ChatContentPart::ImageUrl { url, detail, .. } => { + MediaContentPart::ImageUrl { url, detail, .. } => { assert_eq!(url, "https://example.com/image.jpg"); assert_eq!(*detail, Some(ImageDetail::High)); } diff --git a/model_gateway/src/routers/grpc/regular/stages/messages/preparation.rs b/model_gateway/src/routers/grpc/regular/stages/messages/preparation.rs index ec32fffead..a58b831399 100644 --- a/model_gateway/src/routers/grpc/regular/stages/messages/preparation.rs +++ b/model_gateway/src/routers/grpc/regular/stages/messages/preparation.rs @@ -3,13 +3,14 @@ use async_trait::async_trait; use axum::response::Response; use openai_protocol::{common::StringOrArray, messages::CreateMessageRequest}; -use tracing::error; +use tracing::{debug, error}; use crate::routers::{ error, grpc::{ common::stages::PipelineStage, context::{PreparationOutput, RequestContext}, + multimodal, utils::{self, message_utils}, }, }; @@ -35,10 +36,6 @@ impl PipelineStage for MessagePreparationStage { } impl MessagePreparationStage { - #[expect( - clippy::unused_async, - reason = "multimodal processing will add .await in follow-up PR" - )] async fn prepare_messages( &self, ctx: &mut RequestContext, @@ -99,9 +96,75 @@ impl MessagePreparationStage { } }; - let token_ids = encoding.token_ids().to_vec(); - - // Multimodal processing — postponed (see design doc appendix) + let mut token_ids = encoding.token_ids().to_vec(); + + // Step 3.5: Multimodal processing (fetch + preprocess + expand tokens + hash) + let mut multimodal_intermediate = None; + if multimodal::has_multimodal_content_messages(&request.messages) { + if let Some(mm_components) = ctx.components.multimodal.as_ref() { + let model_id = ctx.input.model_id.as_deref().unwrap_or(&request.model); + let tokenizer_source = ctx + .components + .tokenizer_registry + .get_by_name(model_id) + .or_else(|| ctx.components.tokenizer_registry.get_by_id(model_id)) + .map(|e| e.source) + .unwrap_or_default(); + + if tokenizer_source.is_empty() { + error!( + function = "MessagePreparationStage::execute", + model = %model_id, + "Tokenizer source path not found for multimodal processing" + ); + return Err(error::bad_request( + "multimodal_config_missing", + format!("Tokenizer source path not found for model: {model_id}"), + )); + } + + match multimodal::process_multimodal_messages( + &request.messages, + model_id, + &*tokenizer, + token_ids, + mm_components, + &tokenizer_source, + ) + .await + { + Ok(output) => { + debug!( + function = "MessagePreparationStage::execute", + expanded_tokens = output.expanded_token_ids.len(), + "Multimodal processing complete" + ); + token_ids = output.expanded_token_ids; + multimodal_intermediate = Some(output.intermediate); + } + Err(e) => { + error!( + function = "MessagePreparationStage::execute", + error = %e, + "Multimodal processing failed" + ); + return Err(error::bad_request( + "multimodal_processing_failed", + format!("Multimodal processing failed: {e}"), + )); + } + } + } else { + error!( + function = "MessagePreparationStage::execute", + "Multimodal content detected but multimodal components not initialized" + ); + return Err(error::bad_request( + "multimodal_not_supported", + "Multimodal content detected but multimodal processing is not available", + )); + } + } // Step 4: Build tool constraints if tools present let tool_call_constraint = if filtered_tools.is_empty() { @@ -135,6 +198,9 @@ impl MessagePreparationStage { false, // no_stop_trim default ); + let mut processed_messages = processed_messages; + processed_messages.multimodal_intermediate = multimodal_intermediate; + // Store results in context ctx.state.preparation = Some(PreparationOutput { original_text: Some(processed_messages.text.clone()), diff --git a/model_gateway/src/routers/grpc/regular/stages/messages/request_building.rs b/model_gateway/src/routers/grpc/regular/stages/messages/request_building.rs index df98517944..6cc1311bfe 100644 --- a/model_gateway/src/routers/grpc/regular/stages/messages/request_building.rs +++ b/model_gateway/src/routers/grpc/regular/stages/messages/request_building.rs @@ -10,6 +10,7 @@ use crate::routers::{ grpc::{ common::stages::{helpers, PipelineStage}, context::{ClientSelection, RequestContext}, + multimodal::assemble_multimodal_data, proto_wrapper::ProtoRequest, }, }; @@ -74,13 +75,18 @@ impl PipelineStage for MessageRequestBuildingStage { ) })?; + // Assemble backend-specific multimodal data now that the backend is known + let multimodal_data = processed_messages + .multimodal_intermediate + .map(|intermediate| assemble_multimodal_data(intermediate, builder_client)); + let mut proto_request = builder_client .build_messages_request( request_id, &messages_request, processed_messages.text, prep.token_ids, - None, // multimodal data — postponed + multimodal_data, prep.tool_constraints, ) .map_err(|e| {