use futures::StreamExt;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
use crate::error::LlmError;
use crate::llm::{LlmBackend, LlmRequest, LlmResponse, LlmToolSchema, OnDelta};
use crate::state::{ChatMessage, ToolCallRecord};
const STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
pub struct OpenAiLlmBackend {
api_key: String,
base_url: String,
client: Client,
}
impl OpenAiLlmBackend {
pub fn new(api_key: impl Into<String>) -> Self {
Self::with_base_url(api_key, "https://api.openai.com")
}
pub fn with_base_url(api_key: impl Into<String>, base_url: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
base_url: base_url.into(),
client: Client::builder()
.timeout(Duration::from_secs(45))
.build()
.unwrap_or_else(|_| Client::new()),
}
}
}
#[derive(Serialize)]
struct OaRequest<'a> {
model: &'a str,
messages: Vec<OaMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
tools: Option<Vec<OaTool<'a>>>,
tool_choice: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
stream_options: Option<serde_json::Value>,
}
#[derive(Serialize)]
#[serde(tag = "role", rename_all = "snake_case")]
enum OaMessage {
System {
content: String,
},
User {
content: String,
},
Assistant {
content: Option<String>,
#[serde(skip_serializing_if = "Vec::is_empty")]
tool_calls: Vec<OaToolCallEmit>,
},
Tool {
tool_call_id: String,
content: String,
},
}
#[derive(Serialize)]
struct OaToolCallEmit {
id: String,
#[serde(rename = "type")]
typ: &'static str,
function: OaToolFn,
}
#[derive(Serialize)]
struct OaToolFn {
name: String,
arguments: String,
}
#[derive(Serialize)]
struct OaTool<'a> {
#[serde(rename = "type")]
typ: &'static str,
function: OaToolDef<'a>,
}
#[derive(Serialize)]
struct OaToolDef<'a> {
name: String,
description: &'a str,
parameters: &'a serde_json::Value,
}
#[derive(Deserialize)]
struct OaResponse {
choices: Vec<OaChoice>,
usage: OaUsage,
}
#[derive(Deserialize)]
struct OaChoice {
message: OaMessageIn,
}
#[derive(Deserialize)]
struct OaMessageIn {
content: Option<String>,
tool_calls: Option<Vec<OaToolCallIn>>,
}
#[derive(Deserialize)]
struct OaToolCallIn {
id: String,
function: OaToolFnIn,
}
#[derive(Deserialize)]
struct OaToolFnIn {
name: String,
arguments: String, }
#[derive(Deserialize)]
struct OaUsage {
prompt_tokens: u32,
completion_tokens: u32,
}
impl LlmBackend for OpenAiLlmBackend {
fn complete<'a>(
&'a self,
req: LlmRequest,
) -> Pin<Box<dyn Future<Output = Result<LlmResponse, LlmError>> + Send + 'a>> {
Box::pin(async move {
let messages = build_messages(&req);
let tools: Option<Vec<OaTool<'_>>> = if req.tools.is_empty() {
None
} else {
Some(req.tools.iter().map(build_tool).collect())
};
let body = OaRequest {
model: &req.provider.model,
messages,
tools,
tool_choice: "auto",
stream: None,
stream_options: None,
};
let url = format!("{}/v1/chat/completions", self.base_url);
let resp = self
.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&body)
.send()
.await
.map_err(|e| LlmError::Transport(e.to_string()))?;
let status = resp.status();
if status.is_server_error() {
return Err(LlmError::ServiceUnavailable);
}
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(LlmError::BadRequest(format!("{status}: {text}")));
}
let oa: OaResponse = resp
.json()
.await
.map_err(|e| LlmError::Decode(e.to_string()))?;
let choice = oa
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::Decode("no choices".into()))?;
let tool_calls = choice
.message
.tool_calls
.unwrap_or_default()
.into_iter()
.map(|c| build_tool_call_record(c.id, &c.function.name, &c.function.arguments))
.collect();
Ok(LlmResponse {
content: choice.message.content,
tool_calls,
tokens_in: oa.usage.prompt_tokens,
tokens_out: oa.usage.completion_tokens,
})
})
}
fn complete_streaming<'a>(
&'a self,
req: LlmRequest,
on_delta: OnDelta,
) -> Pin<Box<dyn Future<Output = Result<LlmResponse, LlmError>> + Send + 'a>> {
Box::pin(async move {
let messages = build_messages(&req);
let tools: Option<Vec<OaTool<'_>>> = if req.tools.is_empty() {
None
} else {
Some(req.tools.iter().map(build_tool).collect())
};
let body = OaRequest {
model: &req.provider.model,
messages,
tools,
tool_choice: "auto",
stream: Some(true),
stream_options: Some(serde_json::json!({ "include_usage": true })),
};
let url = format!("{}/v1/chat/completions", self.base_url);
let resp = self
.client
.post(&url)
.bearer_auth(&self.api_key)
.json(&body)
.send()
.await
.map_err(|e| LlmError::Transport(e.to_string()))?;
let status = resp.status();
if status.is_server_error() {
return Err(LlmError::ServiceUnavailable);
}
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(LlmError::BadRequest(format!("{status}: {text}")));
}
let mut acc = StreamAccumulator::new();
let mut line_buf = SseLineBuffer::new();
let mut stream = resp.bytes_stream();
while let Some(chunk) = tokio::time::timeout(STREAM_IDLE_TIMEOUT, stream.next())
.await
.map_err(|_| LlmError::Transport("stream idle timeout".into()))?
{
let bytes = chunk.map_err(|e| LlmError::Transport(e.to_string()))?;
for line in line_buf.push_bytes(&bytes) {
if let Some(text) = acc.push_line(&line)? {
on_delta(&text);
}
}
}
if let Some(line) = line_buf.take_remainder()
&& let Some(text) = acc.push_line(&line)?
{
on_delta(&text);
}
Ok(acc.finish())
})
}
}
struct SseLineBuffer {
buf: Vec<u8>,
}
impl SseLineBuffer {
fn new() -> Self {
Self { buf: Vec::new() }
}
fn push_bytes(&mut self, bytes: &[u8]) -> Vec<String> {
self.buf.extend_from_slice(bytes);
let mut lines = Vec::new();
while let Some(idx) = self.buf.iter().position(|&b| b == b'\n') {
let line_bytes: Vec<u8> = self.buf.drain(..=idx).collect();
let line = String::from_utf8_lossy(&line_bytes);
lines.push(line.trim_end_matches(['\r', '\n']).to_string());
}
lines
}
fn take_remainder(&mut self) -> Option<String> {
if self.buf.is_empty() {
return None;
}
let line = String::from_utf8_lossy(&self.buf);
let line = line.trim_end_matches(['\r', '\n']).to_string();
self.buf.clear();
Some(line)
}
}
struct StreamAccumulator {
content: String,
tool_calls: BTreeMap<u32, ToolCallFrag>,
tokens_in: u32,
tokens_out: u32,
}
#[derive(Default)]
struct ToolCallFrag {
id: String,
name: String,
arguments: String,
}
impl StreamAccumulator {
fn new() -> Self {
Self {
content: String::new(),
tool_calls: BTreeMap::new(),
tokens_in: 0,
tokens_out: 0,
}
}
fn push_line(&mut self, line: &str) -> Result<Option<String>, LlmError> {
let line = line.trim();
if line.is_empty() {
return Ok(None);
}
let payload = match line.strip_prefix("data:") {
Some(rest) => rest.trim(),
None => return Ok(None), };
if payload == "[DONE]" {
return Ok(None);
}
let chunk: OaStreamChunk =
serde_json::from_str(payload).map_err(|e| LlmError::Decode(e.to_string()))?;
if let Some(usage) = chunk.usage {
self.tokens_in = usage.prompt_tokens;
self.tokens_out = usage.completion_tokens;
}
let mut emit: Option<String> = None;
for choice in chunk.choices {
let delta = choice.delta;
if let Some(text) = delta.content
&& !text.is_empty()
{
self.content.push_str(&text);
emit = Some(text);
}
for tc in delta.tool_calls.unwrap_or_default() {
let frag = self.tool_calls.entry(tc.index).or_default();
if let Some(id) = tc.id {
frag.id = id;
}
if let Some(func) = tc.function {
if let Some(name) = func.name {
frag.name = name;
}
if let Some(args) = func.arguments {
frag.arguments.push_str(&args);
}
}
}
}
Ok(emit)
}
fn finish(self) -> LlmResponse {
let tool_calls = self
.tool_calls
.into_values()
.map(|frag| build_tool_call_record(frag.id, &frag.name, &frag.arguments))
.collect();
let content = if self.content.is_empty() {
None
} else {
Some(self.content)
};
LlmResponse {
content,
tool_calls,
tokens_in: self.tokens_in,
tokens_out: self.tokens_out,
}
}
}
#[derive(Deserialize)]
struct OaStreamChunk {
#[serde(default)]
choices: Vec<OaStreamChoice>,
#[serde(default)]
usage: Option<OaUsage>,
}
#[derive(Deserialize)]
struct OaStreamChoice {
delta: OaStreamDelta,
}
#[derive(Deserialize, Default)]
struct OaStreamDelta {
#[serde(default)]
content: Option<String>,
#[serde(default)]
tool_calls: Option<Vec<OaStreamToolCall>>,
}
#[derive(Deserialize)]
struct OaStreamToolCall {
index: u32,
#[serde(default)]
id: Option<String>,
#[serde(default)]
function: Option<OaStreamToolFn>,
}
#[derive(Deserialize)]
struct OaStreamToolFn {
#[serde(default)]
name: Option<String>,
#[serde(default)]
arguments: Option<String>,
}
fn build_tool_call_record(call_id: String, name: &str, raw_args: &str) -> ToolCallRecord {
let args: serde_json::Value = serde_json::from_str(raw_args).unwrap_or(serde_json::json!({}));
let (extension_id, tool_name) = split_tool_name(name);
ToolCallRecord {
call_id,
extension_id,
tool_name,
args,
}
}
pub fn encode_tool_name(extension_id: &str, tool_name: &str) -> String {
if extension_id.is_empty() {
return tool_name.to_string();
}
format!("{}_FN_{}", extension_id.replace('.', "_DOT_"), tool_name)
}
pub fn split_tool_name(name: &str) -> (String, String) {
match name.split_once("_FN_") {
Some((ext, tool)) => (ext.replace("_DOT_", "."), tool.to_string()),
None => (String::new(), name.to_string()),
}
}
fn build_messages(req: &LlmRequest) -> Vec<OaMessage> {
let mut out: Vec<OaMessage> = Vec::with_capacity(req.history.len() + 1);
out.push(OaMessage::System {
content: req.system_prompt.clone(),
});
for m in &req.history {
match m {
ChatMessage::System { content } => out.push(OaMessage::System {
content: content.clone(),
}),
ChatMessage::User { content } => out.push(OaMessage::User {
content: content.clone(),
}),
ChatMessage::Assistant {
content,
tool_calls,
} => {
let calls = tool_calls
.iter()
.map(|tc| OaToolCallEmit {
id: tc.call_id.clone(),
typ: "function",
function: OaToolFn {
name: encode_tool_name(&tc.extension_id, &tc.tool_name),
arguments: tc.args.to_string(),
},
})
.collect();
out.push(OaMessage::Assistant {
content: Some(content.clone()),
tool_calls: calls,
});
}
ChatMessage::Tool { call_id, content } => {
out.push(OaMessage::Tool {
tool_call_id: call_id.clone(),
content: content.to_string(),
});
}
}
}
out
}
fn build_tool(t: &LlmToolSchema) -> OaTool<'_> {
OaTool {
typ: "function",
function: OaToolDef {
name: encode_tool_name(&t.extension_id, &t.tool_name),
description: &t.description,
parameters: &t.parameters,
},
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn tool_name_round_trips_and_is_openai_safe() {
for (ext, tool) in [
("greentic.telco-x-tools", "tx_resolve_prefix"),
("greentic.tavily", "tavily_search"),
("http", "fetch"),
] {
let encoded = encode_tool_name(ext, tool);
assert!(
!encoded.contains('.'),
"encoded name must be OpenAI-safe (no dots): {encoded}"
);
assert_eq!(split_tool_name(&encoded), (ext.into(), tool.into()));
}
assert_eq!(encode_tool_name("", "toolname-no-ext"), "toolname-no-ext");
assert_eq!(
split_tool_name("toolname-no-ext"),
(String::new(), "toolname-no-ext".into())
);
}
#[test]
fn line_buffer_reassembles_multibyte_char_split_across_chunks() {
let full = "data: {\"choices\":[{\"delta\":{\"content\":\"héllo\"}}]}\n";
let bytes = full.as_bytes();
let e_pos = full.find('é').expect("contains é");
let split = e_pos + 1; let (a, b) = bytes.split_at(split);
let mut lb = SseLineBuffer::new();
let mut lines = lb.push_bytes(a);
lines.extend(lb.push_bytes(b));
assert_eq!(lines.len(), 1, "one complete line expected");
assert!(
!lines[0].contains('\u{FFFD}'),
"line must not contain U+FFFD: {:?}",
lines[0]
);
let mut acc = StreamAccumulator::new();
let delta = acc.push_line(&lines[0]).expect("parse");
assert_eq!(delta.as_deref(), Some("héllo"));
}
#[test]
fn line_buffer_reassembles_emoji_split_across_chunks() {
let full = "data: {\"choices\":[{\"delta\":{\"content\":\"hi 😀\"}}]}\n";
let bytes = full.as_bytes();
let emoji_pos = full.find('😀').expect("contains emoji");
let (a, b) = bytes.split_at(emoji_pos + 2); let mut lb = SseLineBuffer::new();
let mut lines = lb.push_bytes(a);
lines.extend(lb.push_bytes(b));
assert_eq!(lines.len(), 1);
let mut acc = StreamAccumulator::new();
let delta = acc.push_line(&lines[0]).expect("parse");
assert_eq!(delta.as_deref(), Some("hi 😀"));
}
#[test]
fn line_buffer_yields_remainder_without_trailing_newline() {
let mut lb = SseLineBuffer::new();
let lines = lb.push_bytes(b"data: {\"choices\":[{\"delta\":{\"content\":\"x\"}}]}");
assert!(lines.is_empty(), "no newline yet: no complete line");
let rem = lb.take_remainder().expect("remainder present");
let mut acc = StreamAccumulator::new();
assert_eq!(acc.push_line(&rem).expect("parse").as_deref(), Some("x"));
assert!(lb.take_remainder().is_none(), "remainder consumed");
}
#[test]
fn line_buffer_splits_two_events_in_one_chunk() {
let mut lb = SseLineBuffer::new();
let chunk = "data: {\"choices\":[{\"delta\":{\"content\":\"A\"}}]}\n\
data: {\"choices\":[{\"delta\":{\"content\":\"B\"}}]}\n";
let lines = lb.push_bytes(chunk.as_bytes());
assert_eq!(lines.len(), 2);
let mut acc = StreamAccumulator::new();
let d0 = acc.push_line(&lines[0]).expect("parse");
let d1 = acc.push_line(&lines[1]).expect("parse");
assert_eq!(d0.as_deref(), Some("A"));
assert_eq!(d1.as_deref(), Some("B"));
}
#[test]
fn push_line_handles_crlf_line_ending() {
let mut acc = StreamAccumulator::new();
let line = "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\r";
let delta = acc.push_line(line).expect("parse");
assert_eq!(delta.as_deref(), Some("hi"));
}
#[test]
fn push_line_ignores_sse_comment_keepalive() {
let mut acc = StreamAccumulator::new();
assert_eq!(acc.push_line(": keep-alive").expect("no error"), None);
assert_eq!(acc.push_line(":").expect("no error"), None);
assert_eq!(acc.push_line("").expect("no error"), None);
}
#[test]
fn accumulates_stream_lines_into_response() {
let lines = [
r#"data: {"choices":[{"delta":{"content":"Hel"}}]}"#,
r#"data: {"choices":[{"delta":{"content":"lo"}}]}"#,
r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"c1","function":{"name":"kb_FN_lookup","arguments":"{\"q\":"}}]}}]}"#,
r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"x\"}"}}]}}]}"#,
r#"data: {"usage":{"prompt_tokens":7,"completion_tokens":9},"choices":[]}"#,
"data: [DONE]",
];
let mut acc = StreamAccumulator::new();
let mut deltas = Vec::new();
for l in lines {
if let Some(text) = acc.push_line(l).expect("parse") {
deltas.push(text);
}
}
let resp = acc.finish();
assert_eq!(deltas, vec!["Hel".to_string(), "lo".to_string()]);
assert_eq!(resp.content.as_deref(), Some("Hello"));
assert_eq!(resp.tool_calls.len(), 1);
assert_eq!(resp.tool_calls[0].tool_name, "lookup");
assert_eq!(resp.tool_calls[0].args, serde_json::json!({"q":"x"}));
assert_eq!(resp.tokens_in, 7);
assert_eq!(resp.tokens_out, 9);
}
#[test]
fn assistant_with_empty_tool_calls_omits_field() {
use crate::config::LlmProviderRef;
use crate::state::ChatMessage;
let req = LlmRequest {
system_prompt: "sys".into(),
history: vec![ChatMessage::Assistant {
content: "hi".into(),
tool_calls: vec![],
}],
tools: vec![],
provider: LlmProviderRef {
provider: "openai".into(),
model: "gpt-4o".into(),
credential_ref: None,
},
};
let value = serde_json::to_value(build_messages(&req)).unwrap();
assert_eq!(value[1]["role"], "assistant");
assert!(
value[1].get("tool_calls").is_none(),
"empty tool_calls must be omitted, got: {}",
value[1]
);
}
}