use serde::{Deserialize, Serialize};
use crate::Json;
pub const ANNOTATED_LLM_REQUEST_SCHEMA: &str = "nemo.relay.AnnotatedLlmRequest@2";
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct AnnotatedLlmRequest {
#[serde(default)]
pub messages: Vec<Message>,
#[serde(skip_serializing_if = "Option::is_none")]
pub instructions: Option<MessageContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub params: Option<GenerationParams>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<ToolDefinition>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
pub store: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub previous_response_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub truncation: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_tier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parallel_tool_calls: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tool_calls: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_logprobs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_specific: Option<ApiSpecificRequest>,
#[serde(flatten)]
pub extra: serde_json::Map<String, Json>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "role", rename_all = "lowercase")]
pub enum Message {
System {
content: MessageContent,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
User {
content: MessageContent,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
Developer {
content: MessageContent,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
Assistant {
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<MessageContent>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
Tool {
content: MessageContent,
tool_call_id: String,
},
Function {
content: Option<String>,
name: String,
},
#[serde(rename = "tool_call")]
ToolCallItem {
#[serde(skip_serializing_if = "Option::is_none")]
id: Option<String>,
call_id: String,
name: String,
arguments: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
#[serde(rename = "tool_result")]
ToolResultItem {
#[serde(skip_serializing_if = "Option::is_none")]
id: Option<String>,
call_id: String,
output: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
#[serde(rename = "provider_native")]
ProviderNative {
provider: String,
kind: String,
value: Json,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Parts(Vec<ContentPart>),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ContentPart {
Text {
text: String,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
ImageUrl {
image_url: OpenAiImageUrl,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
Image {
image: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
Audio {
audio: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
File {
file: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
Refusal {
refusal: String,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
ToolUse {
id: String,
name: String,
input: Json,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
ToolResult {
tool_use_id: String,
content: Json,
#[serde(skip_serializing_if = "Option::is_none")]
is_error: Option<bool>,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
ProviderNative {
provider: String,
kind: String,
value: Json,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct OpenAiImageUrl {
pub url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub detail: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolCall {
pub id: String,
#[serde(rename = "type")]
pub call_type: String,
pub function: FunctionCall,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct FunctionCall {
pub name: String,
pub arguments: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ToolDefinition {
Function {
function: FunctionDefinition,
#[serde(default, flatten)]
extra: serde_json::Map<String, Json>,
},
ProviderNative {
provider: String,
kind: String,
value: Json,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct FunctionDefinition {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
pub strict: Option<bool>,
#[serde(default, flatten)]
pub extra: serde_json::Map<String, Json>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ToolChoice {
Auto,
None,
Required,
#[serde(untagged)]
Specific(ToolChoiceFunction),
#[serde(untagged)]
ProviderNative(ProviderNativeComponent),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ProviderNativeComponent {
pub provider: String,
pub kind: String,
pub value: Json,
}
#[allow(
clippy::large_enum_variant,
reason = "provider wire-schema fields stay directly mutable on each public variant"
)]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "api")]
pub enum ApiSpecificRequest {
#[serde(rename = "anthropic_messages")]
AnthropicMessages {
#[serde(skip_serializing_if = "Option::is_none")]
cache_control: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
container: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
inference_geo: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
output_config: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
thinking: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
top_k: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
user_profile_id: Option<String>,
},
#[serde(rename = "openai_chat")]
OpenAIChat {
#[serde(skip_serializing_if = "Option::is_none")]
audio: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
frequency_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
function_call: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
functions: Option<Vec<Json>>,
#[serde(skip_serializing_if = "Option::is_none")]
logit_bias: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
logprobs: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
modalities: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
moderation: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
n: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
prediction: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
presence_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_options: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_retention: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
reasoning_effort: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
response_format: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
safety_identifier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
seed: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
stream_options: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
verbosity: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
web_search_options: Option<Json>,
},
#[serde(rename = "openai_responses")]
OpenAIResponses {
#[serde(skip_serializing_if = "Option::is_none")]
background: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
context_management: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
conversation: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
moderation: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_options: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_retention: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
safety_identifier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
stream_options: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
text: Option<Json>,
},
#[serde(rename = "oci_genai")]
OCIGenAI {
#[serde(skip_serializing_if = "Option::is_none")]
compartment_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
serving_mode: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
api_format: Option<String>,
},
#[serde(rename = "custom")]
Custom {
api_name: String,
data: Json,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolChoiceFunction {
#[serde(rename = "type")]
pub choice_type: String,
pub function: ToolChoiceFunctionName,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolChoiceFunctionName {
pub name: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub struct GenerationParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop: Option<Vec<String>>,
}
impl AnnotatedLlmRequest {
pub fn system_prompt(&self) -> Option<&str> {
if let Some(text) = self.instructions.as_ref().and_then(first_content_text) {
return Some(text);
}
self.messages.iter().find_map(|m| match m {
Message::System { content, .. } | Message::Developer { content, .. } => {
first_content_text(content)
}
_ => None,
})
}
pub fn last_user_message(&self) -> Option<&str> {
self.messages.iter().rev().find_map(|m| match m {
Message::User { content, .. } => first_content_text(content),
Message::ProviderNative {
provider, value, ..
} if provider == "openai_responses"
&& value.get("role").and_then(Json::as_str) == Some("user") =>
{
native_message_text(value)
}
_ => None,
})
}
pub fn has_tool_calls(&self) -> bool {
self.messages.iter().any(|m| match m {
Message::Assistant {
tool_calls: Some(calls),
content,
..
} => !calls.is_empty() || content.as_ref().is_some_and(content_has_tool_use),
Message::Assistant {
content: Some(content),
..
} => content_has_tool_use(content),
Message::ToolCallItem { .. } => true,
Message::ProviderNative { value, .. } => matches!(
value.get("type").and_then(Json::as_str),
Some("function_call" | "custom_tool_call" | "tool_use")
),
_ => false,
})
}
}
fn first_content_text(content: &MessageContent) -> Option<&str> {
match content {
MessageContent::Text(text) => Some(text.as_str()),
MessageContent::Parts(parts) => parts.iter().find_map(|part| match part {
ContentPart::Text { text, .. } => Some(text.as_str()),
ContentPart::ProviderNative { value, .. } => value
.get("text")
.and_then(Json::as_str)
.or_else(|| value.get("refusal").and_then(Json::as_str)),
ContentPart::ImageUrl { .. }
| ContentPart::Image { .. }
| ContentPart::Audio { .. }
| ContentPart::File { .. }
| ContentPart::Refusal { .. }
| ContentPart::ToolUse { .. }
| ContentPart::ToolResult { .. } => None,
}),
}
}
fn content_has_tool_use(content: &MessageContent) -> bool {
match content {
MessageContent::Text(_) => false,
MessageContent::Parts(parts) => parts.iter().any(|part| match part {
ContentPart::ToolUse { .. } => true,
ContentPart::ProviderNative { value, .. } => matches!(
value.get("type").and_then(Json::as_str),
Some("tool_use" | "mcp_tool_use" | "server_tool_use")
),
_ => false,
}),
}
}
fn native_message_text(value: &Json) -> Option<&str> {
match value.get("content")? {
Json::String(text) => Some(text.as_str()),
Json::Array(parts) => parts.iter().find_map(|part| {
part.get("text")
.and_then(Json::as_str)
.or_else(|| part.get("refusal").and_then(Json::as_str))
}),
_ => None,
}
}