use std::cell::{Cell, RefCell};
use std::collections::BTreeMap;
use std::rc::Rc;
use apiplant_ai::{ChatReply, ChatRequest, Done, Event, Message, Role, ToolCall, ToolDefinition};
use apiplant_auth::Principal;
use apiplant_core::{Access, Agent, Scope};
use futures_util::StreamExt;
use ntex::util::Bytes;
use ntex::web::types::{Json, Path, State};
use ntex::web::{HttpRequest, HttpResponse};
use serde::Deserialize;
use serde_json::{json, Map, Value};
use uuid::Uuid;
use crate::response::{db_error, error, ok};
use crate::sse;
use crate::state::AppState;
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct Body {
message: String,
messages: Vec<Message>,
thread_id: Option<Uuid>,
title: Option<String>,
stream: Option<bool>,
}
struct Caller {
principal: Option<Principal>,
active_org: Option<Uuid>,
}
pub async fn chat(
req: HttpRequest,
state: State<AppState>,
path: Path<String>,
body: Json<Body>,
) -> HttpResponse {
let name = path.into_inner();
let Some(agent) = state.app.agents.get(&name).cloned() else {
return error(404, format!("unknown ai agent `{name}`"));
};
let Some(ai) = state
.agent_ais
.get(&name)
.cloned()
.or_else(|| state.ai.clone())
else {
return error(404, "this app has no ai assistant");
};
let caller = match admit(&state, &req, &agent).await {
Ok(caller) => caller,
Err(response) => return response,
};
let body = body.into_inner();
let prepared = match prepare(&state, &agent, &caller, body).await {
Ok(prepared) => prepared,
Err(response) => return response,
};
if prepared.stream == Some(false) {
return match chat_with_tools(&state, &prepared, &ai).await {
Ok(reply) => {
if let Err(response) = persist_assistant(&state, &prepared, &reply).await {
return response;
}
ok(&reply_json(reply, prepared.thread_id))
}
Err(e) => refused(e),
};
}
if !prepared.agent.tools.is_empty() {
return match chat_with_tools(&state, &prepared, &ai).await {
Ok(reply) => {
if let Err(response) = persist_assistant(&state, &prepared, &reply).await {
return response;
}
let mut response = HttpResponse::Ok();
sse::headers(&mut response);
let thread_id = prepared.thread_id;
let events = futures_util::stream::iter([
Ok::<Bytes, sse::Never>(sse::delta(&reply.text)),
Ok(sse::done(&done_json(reply.done, thread_id))),
]);
response.streaming(Box::pin(events))
}
Err(e) => refused(e),
};
}
let stream = match ai.stream(&prepared.request).await {
Ok(stream) => stream,
Err(e) => return refused(e),
};
let text = Rc::new(RefCell::new(String::new()));
let done = Rc::new(RefCell::new(Done::default()));
let reply_model =
prepared
.request
.model
.clone()
.unwrap_or_else(|| ai.model().to_string());
let ended = Rc::new(Cell::new(false));
let failed = Rc::new(Cell::new(false));
let text_events = text.clone();
let done_events = done.clone();
let ended_events = ended.clone();
let failed_events = failed.clone();
let events = stream
.map(move |event| -> Result<Bytes, sse::Never> {
Ok(match event {
Ok(Event::Delta(chunk)) => {
text_events.borrow_mut().push_str(&chunk);
sse::delta(&chunk)
}
Ok(Event::Reasoning(chunk)) => sse::event("reasoning", &json!({ "text": chunk })),
Ok(Event::Done(agent_done)) => {
*done_events.borrow_mut() = agent_done;
ended_events.set(true);
Bytes::new()
}
Err(e) => {
failed_events.set(true);
sse::failure(&e.to_string())
}
})
})
.chain(futures_util::stream::once(async move {
let reply = ChatReply {
text: text.borrow().clone(),
provider: ai.provider().as_str().to_string(),
model: reply_model,
done: done.borrow().clone(),
tool_calls: Vec::new(),
};
let frame = match persist_assistant(&state, &prepared, &reply).await {
Ok(()) => sse::done(&done_json(reply.done.clone(), prepared.thread_id)),
Err(_) => {
let mut bytes = sse::failure("agent history could not be saved").to_vec();
bytes.extend_from_slice(
sse::done(&done_json(reply.done.clone(), prepared.thread_id)).as_ref(),
);
Bytes::from(bytes)
}
};
Ok::<Bytes, sse::Never>(frame)
}));
let mut response = HttpResponse::Ok();
sse::headers(&mut response);
response.streaming(Box::pin(events))
}
struct Prepared {
request: ChatRequest,
thread_id: Option<Uuid>,
owner_id: Option<Uuid>,
active_org: Option<Uuid>,
stream: Option<bool>,
agent: Agent,
}
async fn prepare(
state: &State<AppState>,
agent: &Agent,
caller: &Caller,
body: Body,
) -> Result<Prepared, HttpResponse> {
let message = body.message.trim().to_string();
if message.is_empty() {
return Err(error(400, "ask a message"));
}
if body.messages.iter().any(|entry| entry.role == Role::System) {
return Err(error(
400,
"agent conversations cannot include caller-supplied system messages",
));
}
let mut messages = if agent.meta.storage.enabled {
if !body.messages.is_empty() {
return Err(error(
400,
"stored agents keep their own history; use `thread_id` instead of `messages`",
));
}
if let Some(thread_id) = body.thread_id {
load_messages(state, agent, caller, thread_id).await?
} else {
Vec::new()
}
} else {
if body.thread_id.is_some() {
return Err(error(400, "this agent does not persist history"));
}
body.messages
};
messages.push(Message::user(message.clone()));
let owner_id = caller.principal.as_ref().map(|principal| principal.user_id);
let thread_id = if agent.meta.storage.enabled {
let owner = owner_id.expect("stored agents require an authenticated caller");
Some(
ensure_thread(
state,
agent,
caller.active_org,
owner,
body.thread_id,
body.title.as_deref(),
&message,
)
.await?,
)
} else {
None
};
if let Some(thread_id) = thread_id {
insert_message(
state,
agent,
caller.active_org,
owner_id.expect("stored agents require an owner"),
thread_id,
"user",
&message,
None,
)
.await?;
}
let mut request = ChatRequest {
messages,
model: agent.meta.model.clone(),
system: (!agent.meta.system.trim().is_empty()).then(|| {
render_system_prompt(
&agent.meta.system,
&agent_context(agent, caller, owner_id, thread_id),
)
}),
temperature: agent.meta.temperature,
max_tokens: agent.meta.max_tokens,
tools: agent
.tools
.iter()
.map(|tool| ToolDefinition {
name: tool.name.clone(),
description: tool.description.clone(),
input_schema: tool.input_schema.clone(),
})
.collect(),
};
if request.model.as_deref().is_some_and(|model| model.trim().is_empty()) {
request.model = None;
}
Ok(Prepared {
request,
thread_id,
owner_id,
active_org: caller.active_org,
stream: body.stream,
agent: agent.clone(),
})
}
fn agent_context(
agent: &Agent,
caller: &Caller,
owner_id: Option<Uuid>,
thread_id: Option<Uuid>,
) -> BTreeMap<String, String> {
let mut values = BTreeMap::new();
values.insert("agent_name".to_string(), agent.meta.name.clone());
values.insert("agent_description".to_string(), agent.meta.description.clone());
values.insert(
"authenticated".to_string(),
caller.principal.is_some().to_string(),
);
values.insert(
"user_id".to_string(),
owner_id.map(|id| id.to_string()).unwrap_or_default(),
);
values.insert(
"principal_id".to_string(),
owner_id.map(|id| id.to_string()).unwrap_or_default(),
);
values.insert(
"organization_id".to_string(),
caller
.active_org
.map(|id| id.to_string())
.unwrap_or_default(),
);
values.insert(
"thread_id".to_string(),
thread_id.map(|id| id.to_string()).unwrap_or_default(),
);
values
}
fn render_system_prompt(template: &str, context: &BTreeMap<String, String>) -> String {
let mut rendered = template.to_string();
for (key, value) in context {
rendered = rendered.replace(&format!("{{{{{key}}}}}"), value);
rendered = rendered.replace(&format!("{{{{ {key} }}}}"), value);
}
rendered
}
async fn admit(
state: &State<AppState>,
req: &HttpRequest,
agent: &Agent,
) -> Result<Caller, HttpResponse> {
let principal = state.resolve_principal(req).await;
if agent.meta.scope == Scope::Organization {
if matches!(agent.permissions.chat, Access::Private) {
return Err(error(404, format!("unknown ai agent `{}`", agent.meta.name)));
}
let Some(principal) = principal else {
return Err(error(401, "authentication required"));
};
let active_org = state.active_org(req, &Some(principal.clone()));
let Some(org) = active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
let Some(membership) = principal.membership(org) else {
return Err(error(403, "you are not a member of this organisation"));
};
if let Access::Role(role) = &agent.permissions.chat {
if !membership.has_role(role) {
return Err(error(
403,
format!("requires the `{role}` role in this organisation"),
));
}
}
return Ok(Caller {
principal: Some(principal),
active_org: Some(org),
});
}
match &agent.permissions.chat {
Access::Public => Ok(Caller {
principal,
active_org: None,
}),
Access::Private => Err(error(404, format!("unknown ai agent `{}`", agent.meta.name))),
Access::Authenticated => match principal {
Some(principal) => Ok(Caller {
principal: Some(principal),
active_org: None,
}),
None => Err(error(401, "authentication required")),
},
Access::Member => match principal {
Some(principal) if state.active_org(req, &Some(principal.clone())).is_some() => Ok(Caller {
active_org: state.active_org(req, &Some(principal.clone())),
principal: Some(principal),
}),
Some(_) => Err(error(
403,
"select an organisation with the X-Organization header",
)),
None => Err(error(401, "authentication required")),
},
Access::Role(role) => match principal {
Some(principal) => {
let active_org = state.active_org(req, &Some(principal.clone()));
let Some(org) = active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
if !principal.has_role_in(org, role) {
return Err(error(403, format!("requires the `{role}` role in this organisation")));
}
Ok(Caller {
principal: Some(principal),
active_org,
})
}
None => Err(error(401, "authentication required")),
},
Access::Owner => Err(error(500, "internal error")),
}
}
async fn chat_with_tools(
state: &State<AppState>,
prepared: &Prepared,
ai: &apiplant_ai::Ai,
) -> Result<ChatReply, apiplant_ai::AiError> {
let mut request = prepared.request.clone();
for _ in 0..8 {
let reply = ai.complete(&request).await?;
if reply.tool_calls.is_empty() {
return Ok(reply);
}
let calls = reply.tool_calls.clone();
persist_tool_calls(state, prepared, &calls).await.map_err(|response| {
apiplant_ai::AiError::Transport(format!(
"agent history could not be saved: {}",
response.status()
))
})?;
request.messages.push(Message::assistant_tool_calls(calls.clone()));
for call in calls {
let output = invoke_tool(state, prepared, &call).await?;
persist_tool_result(state, prepared, &call, &output).await.map_err(|response| {
apiplant_ai::AiError::Transport(format!(
"agent history could not be saved: {}",
response.status()
))
})?;
request.messages.push(Message::tool_result(call.id, output));
}
}
Err(apiplant_ai::AiError::Request(
"agent exceeded the tool-call limit".to_string(),
))
}
async fn invoke_tool(
state: &State<AppState>,
prepared: &Prepared,
call: &ToolCall,
) -> Result<String, apiplant_ai::AiError> {
let Some(tool) = prepared
.agent
.tools
.iter()
.find(|tool| tool.name == call.name)
else {
return Ok(json!({ "error": format!("unknown tool `{}`", call.name) }).to_string());
};
let Some(loaded) = state.functions.get(&tool.function) else {
tracing::warn!(
agent = %prepared.agent.meta.name,
tool = %tool.name,
function = %tool.function,
"agent tool names a missing function"
);
return Ok(json!({ "error": format!("tool function `{}` is not loaded", tool.function) }).to_string());
};
let functions = state.functions.clone();
let function = tool.function.clone();
let config_json = loaded.config_json.clone();
let principal_id = prepared
.owner_id
.map(|id| id.to_string())
.unwrap_or_default();
let input = call.input.to_string();
let bridge = crate::functions::HostBridge::new(
state.db.clone(),
tokio::runtime::Handle::current(),
config_json,
principal_id,
)
.with_services(
state.mailer.clone(),
state.cache.clone(),
state.payments.clone(),
state.ai.clone(),
);
tokio::task::spawn_blocking(move || {
let f = functions.get(&function).expect("checked above");
f.invoke(bridge, &input)
})
.await
.map_err(|_| apiplant_ai::AiError::Transport("tool function task panicked".to_string()))?
.map_err(|message| {
tracing::warn!(
agent = %prepared.agent.meta.name,
tool = %tool.name,
function = %tool.function,
error = %message,
"agent tool function failed"
);
apiplant_ai::AiError::Request(format!("tool `{}` failed: {message}", tool.name))
})
}
async fn ensure_thread(
state: &State<AppState>,
agent: &Agent,
active_org: Option<Uuid>,
owner_id: Uuid,
thread_id: Option<Uuid>,
explicit_title: Option<&str>,
message: &str,
) -> Result<Uuid, HttpResponse> {
if let Some(thread_id) = thread_id {
let rows = load_thread_rows(state, agent, active_org, owner_id, thread_id).await?;
return rows
.first()
.and_then(|row| row.get("id"))
.and_then(Value::as_str)
.and_then(|value| Uuid::parse_str(value).ok())
.ok_or_else(|| error(404, "unknown thread"));
}
let resource_name = agent.thread_resource_name();
let Some(resource) = state.app.resources.get(&resource_name) else {
return Err(error(500, "agent history resource is missing"));
};
let mut data = Map::new();
data.insert("owner_id".into(), json!(owner_id));
let title = explicit_title
.map(str::trim)
.filter(|title| !title.is_empty())
.map(str::to_string)
.unwrap_or_else(|| title_from(message));
if !title.is_empty() {
data.insert("title".into(), json!(title));
}
if agent.meta.scope == Scope::Organization {
let Some(org) = active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
data.insert("organization_id".into(), json!(org));
}
let row = state.db.create(resource, &data).await.map_err(db_error)?;
row.get("id")
.and_then(Value::as_str)
.and_then(|value| Uuid::parse_str(value).ok())
.ok_or_else(|| error(500, "created thread is missing an id"))
}
async fn load_messages(
state: &State<AppState>,
agent: &Agent,
caller: &Caller,
thread_id: Uuid,
) -> Result<Vec<Message>, HttpResponse> {
load_thread_rows(
state,
agent,
caller.active_org,
caller.principal.as_ref().map(|principal| principal.user_id).unwrap_or_default(),
thread_id,
)
.await?;
let Some(table) = state.table(&agent.message_resource_name()) else {
return Err(error(500, "agent history resource is missing"));
};
let sql = format!(
"SELECT role, content, tool_call_id, tool_name, tool_input FROM {table} \
WHERE thread_id = $1::uuid ORDER BY created_at ASC, id ASC"
);
let rows = state
.db
.raw_json(&sql, &[Value::String(thread_id.to_string())])
.await
.map_err(|e| {
tracing::error!(error = %e, "failed to load agent history");
error(500, "internal error")
})?;
let mut messages = Vec::new();
for row in rows.as_array().map(Vec::as_slice).unwrap_or_default() {
let raw_role = row.get("role").and_then(Value::as_str).unwrap_or_default();
let role = match raw_role {
"assistant" | "tool_call" => Role::Assistant,
"system" => Role::System,
"tool" | "tool_result" => Role::Tool,
_ => Role::User,
};
let content = row
.get("content")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string();
if raw_role == "tool_call" {
messages.push(Message::assistant_tool_calls(vec![ToolCall {
id: row
.get("tool_call_id")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
name: row
.get("tool_name")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
input: row.get("tool_input").cloned().unwrap_or_else(|| json!({})),
}]));
continue;
}
messages.push(Message {
role,
content,
tool_call_id: row
.get("tool_call_id")
.and_then(Value::as_str)
.map(str::to_string),
..Message::default()
});
}
Ok(messages)
}
async fn load_thread_rows(
state: &State<AppState>,
agent: &Agent,
active_org: Option<Uuid>,
owner_id: Uuid,
thread_id: Uuid,
) -> Result<Vec<Value>, HttpResponse> {
let Some(table) = state.table(&agent.thread_resource_name()) else {
return Err(error(500, "agent history resource is missing"));
};
let mut sql =
format!("SELECT id::text AS id FROM {table} WHERE id = $1::uuid AND owner_id = $2::uuid");
let mut params = vec![
Value::String(thread_id.to_string()),
Value::String(owner_id.to_string()),
];
if agent.meta.scope == Scope::Organization {
let Some(org) = active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
sql.push_str(" AND organization_id = $3::uuid");
params.push(Value::String(org.to_string()));
}
if matches!(agent.permissions.history, Access::Private) {
return Err(error(404, "unknown thread"));
}
let rows = state.db.raw_json(&sql, ¶ms).await.map_err(|e| {
tracing::error!(error = %e, "failed to load agent thread");
error(500, "internal error")
})?;
let rows = rows.as_array().cloned().unwrap_or_default();
if rows.is_empty() {
return Err(error(404, "unknown thread"));
}
Ok(rows)
}
async fn persist_assistant(
state: &State<AppState>,
prepared: &Prepared,
reply: &ChatReply,
) -> Result<(), HttpResponse> {
let Some(thread_id) = prepared.thread_id else {
return Ok(());
};
let Some(owner_id) = prepared.owner_id else {
return Ok(());
};
if reply.text.trim().is_empty() {
return Ok(());
}
insert_message(
state,
&prepared.agent,
prepared.active_org,
owner_id,
thread_id,
"assistant",
&reply.text,
Some(reply),
)
.await
}
async fn persist_tool_calls(
state: &State<AppState>,
prepared: &Prepared,
calls: &[ToolCall],
) -> Result<(), HttpResponse> {
for call in calls {
insert_tool_event(state, prepared, "tool_call", call, None).await?;
}
Ok(())
}
async fn persist_tool_result(
state: &State<AppState>,
prepared: &Prepared,
call: &ToolCall,
output: &str,
) -> Result<(), HttpResponse> {
insert_tool_event(state, prepared, "tool_result", call, Some(output)).await
}
async fn insert_tool_event(
state: &State<AppState>,
prepared: &Prepared,
role: &str,
call: &ToolCall,
output: Option<&str>,
) -> Result<(), HttpResponse> {
let Some(thread_id) = prepared.thread_id else {
return Ok(());
};
let Some(owner_id) = prepared.owner_id else {
return Ok(());
};
let resource_name = prepared.agent.message_resource_name();
let Some(resource) = state.app.resources.get(&resource_name) else {
return Err(error(500, "agent history resource is missing"));
};
let mut data = Map::new();
data.insert("thread_id".into(), json!(thread_id));
data.insert("owner_id".into(), json!(owner_id));
data.insert("role".into(), json!(role));
data.insert(
"content".into(),
json!(output.map(str::to_string).unwrap_or_else(|| call.input.to_string())),
);
data.insert("tool_call_id".into(), json!(call.id));
data.insert("tool_name".into(), json!(call.name));
data.insert("tool_input".into(), call.input.clone());
if let Some(output) = output {
data.insert(
"tool_output".into(),
serde_json::from_str(output).unwrap_or_else(|_| json!(output)),
);
}
if prepared.agent.meta.scope == Scope::Organization {
let Some(org) = prepared.active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
data.insert("organization_id".into(), json!(org));
}
state.db.create(resource, &data).await.map_err(db_error)?;
Ok(())
}
async fn insert_message(
state: &State<AppState>,
agent: &Agent,
active_org: Option<Uuid>,
owner_id: Uuid,
thread_id: Uuid,
role: &str,
content: &str,
reply: Option<&ChatReply>,
) -> Result<(), HttpResponse> {
let resource_name = agent.message_resource_name();
let Some(resource) = state.app.resources.get(&resource_name) else {
return Err(error(500, "agent history resource is missing"));
};
let mut data = Map::new();
data.insert("thread_id".into(), json!(thread_id));
data.insert("owner_id".into(), json!(owner_id));
data.insert("role".into(), json!(role));
data.insert("content".into(), json!(content));
if agent.meta.scope == Scope::Organization {
let Some(org) = active_org else {
return Err(error(
403,
"select an organisation with the X-Organization header",
));
};
data.insert("organization_id".into(), json!(org));
}
if let Some(reply) = reply {
if !reply.provider.is_empty() {
data.insert("provider".into(), json!(reply.provider));
}
if !reply.model.is_empty() {
data.insert("model".into(), json!(reply.model));
}
if !reply.done.finish_reason.is_empty() {
data.insert("finish_reason".into(), json!(reply.done.finish_reason));
}
if let Some(tokens) = reply.done.input_tokens {
data.insert("input_tokens".into(), json!(tokens as i64));
}
if let Some(tokens) = reply.done.output_tokens {
data.insert("output_tokens".into(), json!(tokens as i64));
}
}
state.db.create(resource, &data).await.map_err(db_error)?;
Ok(())
}
fn title_from(message: &str) -> String {
let single = message.split_whitespace().collect::<Vec<_>>().join(" ");
let short = single.chars().take(80).collect::<String>();
if single.chars().count() > 80 {
format!("{short}…")
} else {
short
}
}
fn reply_json(reply: ChatReply, thread_id: Option<Uuid>) -> Value {
let mut value = serde_json::to_value(reply).unwrap_or_else(|_| json!({}));
if let Some(map) = value.as_object_mut() {
if let Some(thread_id) = thread_id {
map.insert("thread_id".into(), json!(thread_id));
}
}
value
}
fn done_json(done: Done, thread_id: Option<Uuid>) -> Value {
let mut value = serde_json::to_value(done).unwrap_or_else(|_| json!({}));
if let Some(map) = value.as_object_mut() {
if let Some(thread_id) = thread_id {
map.insert("thread_id".into(), json!(thread_id));
}
}
value
}
fn refused(e: apiplant_ai::AiError) -> HttpResponse {
match e {
apiplant_ai::AiError::Request(message) => error(400, message),
apiplant_ai::AiError::Provider {
provider,
status,
body,
} => {
tracing::warn!(provider, status, body = %body, "the ai provider refused a request");
HttpResponse::BadGateway().json(&json!({
"error": format!("the ai provider refused this request: {body}"),
"provider": provider,
"provider_status": status,
}))
}
other => {
tracing::error!(error = %other, "ai request failed");
HttpResponse::BadGateway().json(&json!({ "error": other.to_string() }))
}
}
}