use crate::{CompletionRequest, InputImage, LlmError, Result};
use af_context::RequestContext;
use async_trait::async_trait;
use serde_json::{json, Value};
pub const MAX_INPUT_IMAGES: usize = 8;
pub fn select_images(messages: &mut [crate::ChatMessage], limit: usize) -> Result<()> {
let limit = limit.min(MAX_INPUT_IMAGES);
for message in messages.iter() {
if !message.images.is_empty() && message.role != crate::Role::User {
return Err(LlmError::InvalidInput(
"images require a user message".into(),
));
}
for image in &message.images {
image.validate()?;
}
}
let Some(current) = messages
.iter()
.rposition(|message| message.role == crate::Role::User)
else {
return Ok(());
};
if messages[current].images.len() > limit {
return Err(LlmError::InvalidInput(format!(
"current input exceeds the {limit} image limit"
)));
}
let mut remaining = limit - messages[current].images.len();
for message in messages[..current].iter_mut().rev() {
let keep = remaining.min(message.images.len());
message.images.drain(..message.images.len() - keep);
remaining -= keep;
}
Ok(())
}
#[async_trait]
pub trait ImageResolver: Send + Sync {
async fn resolve(&self, context: &RequestContext, image: &InputImage) -> Result<String>;
}
pub(crate) async fn provider_request(
request: &CompletionRequest,
resolver: Option<&dyn ImageResolver>,
) -> Result<Value> {
let mut value = serde_json::to_value(request)?;
let count = request
.messages
.iter()
.map(|m| m.images.len())
.sum::<usize>();
if count == 0 {
return Ok(value);
}
if count > 8 {
return Err(LlmError::InvalidInput(
"at most eight images per request".into(),
));
}
let context = request
.context
.as_ref()
.ok_or_else(|| LlmError::InvalidInput("image caller context is required".into()))?;
context
.validate()
.map_err(|_| LlmError::InvalidInput("invalid image caller context".into()))?;
let resolver = resolver
.ok_or_else(|| LlmError::InvalidInput("image resolver is not configured".into()))?;
for (message, wire) in request.messages.iter().zip(
value["messages"]
.as_array_mut()
.expect("serialized messages"),
) {
wire.as_object_mut()
.expect("serialized message")
.remove("images");
if message.images.is_empty() {
continue;
}
if message.role != crate::Role::User {
return Err(LlmError::InvalidInput(
"images require a user message".into(),
));
}
let mut parts = Vec::new();
if let Some(text) = &message.content {
parts.push(json!({"type":"text","text":text}));
}
for image in &message.images {
image.validate()?;
let url = resolver.resolve(context, image).await.map_err(|_| {
LlmError::InvalidInput("image authorization or resolution failed".into())
})?;
let parsed = reqwest::Url::parse(&url)
.map_err(|_| LlmError::InvalidInput("invalid image projection".into()))?;
if parsed.scheme() != "https"
|| parsed.host_str().is_none()
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| parsed.fragment().is_some()
{
return Err(LlmError::InvalidInput(
"image projection requires credential-free HTTPS authority".into(),
));
}
parts.push(json!({"type":"image_url","image_url":{"url":url}}));
}
wire["content"] = Value::Array(parts);
}
Ok(value)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a5_current_images_win_over_bounded_recent_history() {
let mut messages = (0..12)
.map(|index| {
let mut message = crate::ChatMessage::user(format!("turn {index}"));
message.images.push(InputImage {
asset_id: format!("asset-{index}").parse().unwrap(),
media_type: "image/png".into(),
});
message
})
.collect::<Vec<_>>();
let original = messages.clone();
select_images(&mut messages, 4).unwrap();
assert_eq!(
messages
.iter()
.flat_map(|m| &m.images)
.map(|i| i.asset_id.as_str())
.collect::<Vec<_>>(),
["asset-8", "asset-9", "asset-10", "asset-11"]
);
assert_eq!(original.iter().flat_map(|m| &m.images).count(), 12);
messages[11].images = vec![original[11].images[0].clone(); 5];
assert!(select_images(&mut messages, 4).is_err());
assert_eq!(messages[11].images.len(), 5);
messages[11].images.clear();
select_images(&mut messages, 0).unwrap();
assert!(messages.iter().all(|m| m.images.is_empty()));
}
struct FreshResolver(std::sync::atomic::AtomicUsize);
#[async_trait]
impl ImageResolver for FreshResolver {
async fn resolve(&self, _: &RequestContext, _: &InputImage) -> Result<String> {
let attempt = self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if attempt >= 2 {
return Err(LlmError::InvalidInput("revoked".into()));
}
Ok(format!("https://assets.example/image?signature={attempt}"))
}
}
#[tokio::test]
async fn a5_each_attempt_reauthorizes_and_refreshes_only_the_ephemeral_projection() {
let mut message = crate::ChatMessage::user("current");
message.images.push(InputImage {
asset_id: "asset".parse().unwrap(),
media_type: "image/png".into(),
});
let mut request = CompletionRequest::new("vision", vec![message]);
request.context = Some(RequestContext {
tenant_id: "tenant".parse().unwrap(),
subject_id: "owner".parse().unwrap(),
request_id: "request".parse().unwrap(),
roles: Default::default(),
entitlements: Default::default(),
locale: "en".into(),
});
let durable = serde_json::to_value(&request).unwrap();
let resolver = FreshResolver(std::sync::atomic::AtomicUsize::new(0));
let first = provider_request(&request, Some(&resolver)).await.unwrap();
let second = provider_request(&request, Some(&resolver)).await.unwrap();
assert_ne!(first, second);
assert!(provider_request(&request, Some(&resolver)).await.is_err());
assert_eq!(serde_json::to_value(&request).unwrap(), durable);
assert!(!durable.to_string().contains("signature"));
}
struct Resolver;
#[async_trait]
impl ImageResolver for Resolver {
async fn resolve(&self, context: &RequestContext, _: &InputImage) -> Result<String> {
if context.subject_id.as_str() != "owner" {
return Err(LlmError::InvalidInput("secret denial".into()));
}
Ok("https://assets.example/image?signature=secret".into())
}
}
#[tokio::test]
async fn image_wire_is_governed_and_durable_request_has_no_signed_url() {
let mut message = crate::ChatMessage::user("Describe");
message.images.push(InputImage {
asset_id: "asset".parse().unwrap(),
media_type: "image/png".into(),
});
let mut request = CompletionRequest::new("vision", vec![message]);
assert!(provider_request(&request, Some(&Resolver)).await.is_err());
request.context = Some(RequestContext {
tenant_id: "tenant".parse().unwrap(),
subject_id: "owner".parse().unwrap(),
request_id: "request".parse().unwrap(),
roles: Default::default(),
entitlements: Default::default(),
locale: "en".into(),
});
assert!(provider_request(&request, None).await.is_err());
let wire = provider_request(&request, Some(&Resolver)).await.unwrap();
assert_eq!(wire["messages"][0]["content"][1]["type"], "image_url");
assert!(wire["messages"][0].get("images").is_none());
assert!(wire.get("context").is_none());
assert!(!serde_json::to_string(&request)
.unwrap()
.contains("signature"));
request.context.as_mut().unwrap().subject_id = "revoked".parse().unwrap();
let error = provider_request(&request, Some(&Resolver))
.await
.unwrap_err();
assert!(!error.to_string().contains("secret"));
request.context.as_mut().unwrap().subject_id = "owner".parse().unwrap();
request.messages[0].images[0].media_type = "text/plain".into();
assert!(provider_request(&request, Some(&Resolver)).await.is_err());
}
struct Projection(&'static str);
#[async_trait]
impl ImageResolver for Projection {
async fn resolve(&self, _: &RequestContext, _: &InputImage) -> Result<String> {
Ok(self.0.into())
}
}
#[tokio::test]
async fn rejects_untrusted_references_and_invalid_projections() {
let mut request = CompletionRequest::new("vision", vec![crate::ChatMessage::user("image")]);
request.context = Some(RequestContext {
tenant_id: "tenant".parse().unwrap(),
subject_id: "owner".parse().unwrap(),
request_id: "request".parse().unwrap(),
roles: Default::default(),
entitlements: Default::default(),
locale: "en".into(),
});
let image = InputImage {
asset_id: "asset".parse().unwrap(),
media_type: "image/png".into(),
};
request.messages[0].images = vec![image.clone(); 9];
assert!(provider_request(&request, Some(&Resolver)).await.is_err());
request.messages[0].images = vec![image];
request.messages[0].role = crate::Role::Assistant;
assert!(provider_request(&request, Some(&Resolver)).await.is_err());
request.messages[0].role = crate::Role::User;
for id in [
"https://assets.example/private?token=secret",
"data:image/png;base64,AAAA",
"asset/secret",
] {
request.messages[0].images[0].asset_id = id.parse().unwrap();
assert!(provider_request(&request, Some(&Resolver)).await.is_err());
}
request.messages[0].images[0].asset_id = "asset".parse().unwrap();
for url in [
"not a url",
"http://assets.example/image",
"https://user:password@assets.example/image",
"https://assets.example/image#secret",
] {
assert!(provider_request(&request, Some(&Projection(url)))
.await
.is_err());
}
request.messages[0].content = None;
assert!(provider_request(&request, Some(&Resolver)).await.is_ok());
request.context.as_mut().unwrap().locale.clear();
assert!(provider_request(&request, Some(&Resolver)).await.is_ok());
}
}