use crate::pi::accumulator::Accumulator;
use crate::pi::converter::{self, Chunk};
use crate::pi::transport::PiTransport;
use crate::store::{now_ms, Msg, SessionMeta, SessionModel, Store};
use serde::Serialize;
use serde_json::json;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{mpsc, Mutex};
#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum SessionEvent {
Message { role: String, text: String },
Reasoning { text: String },
Thinking { on: bool },
Tool { call_id: String, name: String, input: Option<serde_json::Value>, result: Option<serde_json::Value>, is_error: bool },
UiRequest { request_id: serde_json::Value, method: String, title: Option<String>, message: Option<String>, options: Vec<String> },
}
#[derive(Clone, Serialize)]
pub struct ModelInfo {
pub provider: String,
pub model_id: String,
pub name: String,
}
fn parse_models(resp: &serde_json::Value) -> Vec<ModelInfo> {
let arr = match resp.get("data").and_then(|d| d.get("models")).and_then(|m| m.as_array()) {
Some(a) => a,
None => return Vec::new(),
};
arr.iter()
.filter_map(|m| {
let id = m.get("id").and_then(|v| v.as_str())?;
Some(ModelInfo {
provider: m
.get("provider")
.and_then(|v| v.as_str())
.unwrap_or("unknown")
.to_string(),
model_id: id.to_string(),
name: m.get("name").and_then(|v| v.as_str()).unwrap_or(id).to_string(),
})
})
.collect()
}
pub struct Session {
pub meta: Mutex<SessionMeta>,
messages: Mutex<Vec<Msg>>,
subscribers: Mutex<Vec<mpsc::UnboundedSender<SessionEvent>>>,
pi: Mutex<Option<Arc<PiTransport>>>,
current_model: Mutex<Option<(String, String)>>, pi_path: String,
cwd: PathBuf,
store: Store,
}
impl Session {
pub fn new(
meta: SessionMeta,
messages: Vec<Msg>,
pi_path: String,
cwd: PathBuf,
store: Store,
) -> Arc<Self> {
let current_model = meta.model.as_ref().map(|model| (model.provider.clone(), model.model_id.clone()));
Arc::new(Self {
meta: Mutex::new(meta),
messages: Mutex::new(messages),
subscribers: Mutex::new(Vec::new()),
pi: Mutex::new(None),
current_model: Mutex::new(current_model),
pi_path,
cwd,
store,
})
}
pub async fn messages(&self) -> Vec<Msg> {
self.messages.lock().await.clone()
}
pub async fn ensure_pi(self: &Arc<Self>) {
let mut g = self.pi.lock().await;
if g.is_some() {
return;
}
match PiTransport::spawn(&self.pi_path, &self.cwd, None) {
Ok(t) => {
let t = Arc::new(t);
let (saved_model, saved_thinking) = {
let meta = self.meta.lock().await;
(meta.model.clone(), meta.thinking_level.clone())
};
let restored_model = if let Some(model) = saved_model {
t.send_and_wait(
json!({ "type": "set_model", "provider": &model.provider, "modelId": &model.model_id }),
Duration::from_secs(15),
).await.is_ok().then_some((model.provider, model.model_id))
} else {
None
};
if let Some(model) = restored_model {
*self.current_model.lock().await = Some(model);
} else if let Ok(resp) = t
.send_and_wait(json!({ "type": "get_state" }), Duration::from_secs(30))
.await
{
if let Some(m) = resp.get("data").and_then(|d| d.get("model")) {
let id = m
.get("id")
.and_then(|v| v.as_str())
.or_else(|| m.get("modelId").and_then(|v| v.as_str()));
let prov = m.get("provider").and_then(|v| v.as_str()).map(String::from);
if let (Some(id), Some(prov)) = (id, prov) {
let model = SessionModel { provider: prov, model_id: id.to_string() };
*self.current_model.lock().await = Some((model.provider.clone(), model.model_id.clone()));
self.meta.lock().await.model = Some(model);
self.persist().await;
}
}
if saved_thinking.is_none() {
if let Some(level) = resp
.get("data")
.and_then(|data| data.get("thinkingLevel"))
.and_then(|value| value.as_str())
{
self.meta.lock().await.thinking_level = Some(level.to_string());
self.persist().await;
}
}
}
if let Some(level) = saved_thinking {
if t.send_and_wait(json!({ "type": "set_thinking_level", "level": level }), Duration::from_secs(15)).await.is_err() {
log::warn!("failed to restore thinking level for session {}", self.id().await);
}
}
tokio::spawn(pump(Arc::clone(self), Arc::clone(&t)));
*g = Some(t);
log::info!("session {} pi spawned", self.id().await);
}
Err(e) => log::warn!("pi spawn failed: {e}"),
}
}
pub async fn send_prompt(self: &Arc<Self>, text: String) {
self.ensure_pi().await;
let g = self.pi.lock().await;
if let Some(t) = g.as_ref() {
if let Err(e) = t.send(json!({ "type": "prompt", "message": text })) {
log::warn!("pi send: {e}");
}
}
}
pub async fn models(self: &Arc<Self>) -> (Option<(String, String)>, Vec<ModelInfo>) {
self.ensure_pi().await;
let avail = {
let g = self.pi.lock().await;
match g.as_ref() {
Some(t) => match t
.send_and_wait(json!({ "type": "get_available_models" }), Duration::from_secs(15))
.await
{
Ok(r) => parse_models(&r),
Err(_) => Vec::new(),
},
None => Vec::new(),
}
};
(self.current_model.lock().await.clone(), avail)
}
pub async fn set_model(self: &Arc<Self>, provider: String, model_id: String) -> bool {
self.ensure_pi().await;
let ok = {
let g = self.pi.lock().await;
if let Some(t) = g.as_ref() {
t.send_and_wait(
json!({ "type": "set_model", "provider": &provider, "modelId": &model_id }),
Duration::from_secs(15),
)
.await
.is_ok()
} else {
false
}
};
if ok {
*self.current_model.lock().await = Some((provider.clone(), model_id.clone()));
self.meta.lock().await.model = Some(SessionModel { provider, model_id });
self.persist().await;
}
ok
}
pub async fn set_thinking_level(self: &Arc<Self>, level: String) -> bool {
self.ensure_pi().await;
let ok = {
let g = self.pi.lock().await;
if let Some(t) = g.as_ref() {
t
.send_and_wait(json!({ "type": "set_thinking_level", "level": &level }), Duration::from_secs(15))
.await
.is_ok()
} else {
false
}
};
if ok {
self.meta.lock().await.thinking_level = Some(level);
self.persist().await;
}
ok
}
pub async fn rename(&self, title: String) {
self.meta.lock().await.title = title;
self.persist().await;
}
pub async fn abort(&self) {
let g = self.pi.lock().await;
if let Some(t) = g.as_ref() {
let _ = t.send(json!({ "type": "abort" })).ok();
}
drop(g);
self.push_thinking(false).await;
}
pub async fn meta(&self) -> SessionMeta {
self.meta.lock().await.clone()
}
pub async fn push_user(self: &Arc<Self>, text: &str) {
let first = self.messages.lock().await.is_empty();
if first {
let title: String = text.chars().take(40).collect();
self.meta.lock().await.title = title;
}
self.append(Msg {
role: "user".into(),
text: text.into(),
ts: now_ms(),
..Msg::default()
})
.await;
}
async fn append(&self, msg: Msg) {
let role = msg.role.clone();
let text = msg.text.clone();
self.messages.lock().await.push(msg);
self.touch().await;
self.persist().await;
self.broadcast(SessionEvent::Message { role, text }).await;
}
pub async fn push_assistant(self: &Arc<Self>, text: &str) {
self.append(Msg {
role: "assistant".into(),
text: text.into(),
ts: now_ms(),
..Msg::default()
})
.await;
}
pub async fn push_reasoning(&self, text: &str) {
self.messages.lock().await.push(Msg {
role: "reasoning".into(),
text: text.into(),
ts: now_ms(),
..Msg::default()
});
self.touch().await;
self.persist().await;
self.broadcast(SessionEvent::Reasoning {
text: text.into(),
})
.await;
}
pub async fn push_thinking(&self, on: bool) {
self.broadcast(SessionEvent::Thinking { on }).await;
}
pub async fn push_tool(&self, call_id: &str, name: &str, input: Option<serde_json::Value>, result: Option<serde_json::Value>, is_error: bool) {
{
let mut messages = self.messages.lock().await;
if let Some(message) = messages.iter_mut().rev().find(|message| {
message.role == "tool" && message.call_id.as_deref() == Some(call_id)
}) {
if input.is_some() {
message.input = input.clone();
}
if result.is_some() {
message.result = result.clone();
}
message.is_error = is_error;
} else {
messages.push(Msg {
role: "tool".into(),
text: String::new(),
ts: now_ms(),
call_id: Some(call_id.into()),
tool_name: Some(name.into()),
input: input.clone(),
result: result.clone(),
is_error,
});
}
}
self.touch().await;
self.persist().await;
self.broadcast(SessionEvent::Tool {
call_id: call_id.into(),
name: name.into(),
input,
result,
is_error,
})
.await;
}
pub async fn id(&self) -> String {
self.meta.lock().await.id.clone()
}
async fn touch(&self) {
self.meta.lock().await.updated_at = now_ms();
}
async fn persist(&self) {
let meta = self.meta.lock().await.clone();
let msgs = self.messages.lock().await.clone();
if let Err(e) = self.store.save(&meta, &msgs) {
log::warn!("persist failed: {e}");
}
}
async fn broadcast(&self, event: SessionEvent) {
let mut subs = self.subscribers.lock().await;
let mut dead = Vec::new();
for (i, tx) in subs.iter().enumerate() {
if tx.send(event.clone()).is_err() {
dead.push(i);
}
}
for i in dead.into_iter().rev() {
subs.remove(i);
}
}
pub async fn subscribe(&self) -> mpsc::UnboundedReceiver<SessionEvent> {
let (tx, rx) = mpsc::unbounded_channel();
self.subscribers.lock().await.push(tx);
rx
}
pub async fn respond_ui(&self, request_id: serde_json::Value, value: Option<serde_json::Value>, cancelled: bool) -> bool {
let pi = self.pi.lock().await.clone();
let Some(pi) = pi else { return false };
let mut response = json!({ "type": "extension_ui_response", "id": request_id });
if cancelled { response["cancelled"] = json!(true); }
else if let Some(value) = value { response["value"] = value; }
pi.send(response).is_ok()
}
}
async fn pump(session: Arc<Session>, transport: Arc<PiTransport>) {
let mut acc = Accumulator::new();
loop {
let event = {
let mut rx = transport.events_rx.lock().await;
rx.recv().await
};
let Some(event) = event else {
log::info!("session pi events closed");
break;
};
let t = event
.get("type")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
if t == "extension_ui_request" {
let request_id = event.get("id").cloned().unwrap_or(serde_json::Value::Null);
let method = event.get("method").and_then(|v| v.as_str()).unwrap_or("select").to_string();
let title = event.get("title").and_then(|v| v.as_str()).map(str::to_string);
let message = event.get("message").and_then(|v| v.as_str()).map(str::to_string);
let options = event.get("options").and_then(|v| v.as_array()).map(|values| values.iter().filter_map(|v| v.as_str().map(str::to_string)).collect()).unwrap_or_default();
session.broadcast(SessionEvent::UiRequest { request_id, method, title, message, options }).await;
continue;
}
for chunk in acc.handle(&event) {
emit(&session, chunk).await;
}
for chunk in converter::convert_event(&event) {
emit(&session, chunk).await;
}
match t.as_str() {
"agent_start" | "turn_start" => session.push_thinking(true).await,
"turn_end" | "agent_end" => session.push_thinking(false).await,
_ => {}
}
}
}
async fn emit(session: &Arc<Session>, chunk: Chunk) {
match chunk {
Chunk::Text(t) => session.push_assistant(&t).await,
Chunk::Reasoning(t) => session.push_reasoning(&t).await,
Chunk::ToolCall { id, name, input } => session.push_tool(&id, &name, Some(input), None, false).await,
Chunk::ToolResult {
id,
name,
result,
is_error,
} => session.push_tool(&id, &name, None, Some(result), is_error).await,
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn temp_store() -> Store {
let d = std::env::temp_dir().join(format!(
"xagent-pi-session-test-{}-{}",
std::process::id(),
uuid::Uuid::new_v4()
));
Store::new(d).unwrap()
}
fn meta(id: &str) -> SessionMeta {
SessionMeta {
id: id.to_string(),
project_id: "default".to_string(),
title: "New chat".to_string(),
created_at: 1,
updated_at: 1,
worktree: None,
branch: None,
model: None,
thinking_level: None,
}
}
#[test]
fn parse_models_reads_list() {
let resp = json!({
"data": {
"models": [
{ "id": "gpt-4o", "provider": "openai", "name": "GPT-4o" },
{ "id": "claude-3-7-sonnet", "provider": "anthropic", "name": "Claude 3.7 Sonnet" }
]
}
});
let models = parse_models(&resp);
assert_eq!(models.len(), 2);
assert_eq!(models[0].provider, "openai");
assert_eq!(models[0].model_id, "gpt-4o");
assert_eq!(models[0].name, "GPT-4o");
assert_eq!(models[1].provider, "anthropic");
}
#[test]
fn parse_models_defaults_missing_fields() {
let resp = json!({
"data": {
"models": [
{ "id": "only-id" },
{ "name": "no-id" }
]
}
});
let models = parse_models(&resp);
assert_eq!(models.len(), 1);
assert_eq!(models[0].provider, "unknown");
assert_eq!(models[0].model_id, "only-id");
assert_eq!(models[0].name, "only-id"); }
#[test]
fn parse_models_handles_malformed_responses() {
assert!(parse_models(&json!({})).is_empty());
assert!(parse_models(&json!({ "data": {} })).is_empty());
assert!(parse_models(&json!({ "data": { "models": "nope" } })).is_empty());
}
#[tokio::test]
async fn initializes_session_preferences_from_metadata() {
let store = temp_store();
let mut metadata = meta("preferences");
metadata.model = Some(SessionModel {
provider: "openai".to_string(),
model_id: "gpt-test".to_string(),
});
metadata.thinking_level = Some("high".to_string());
let session = Session::new(metadata, vec![], "pi".into(), PathBuf::from("."), store);
assert_eq!(
session.current_model.lock().await.as_ref(),
Some(&("openai".to_string(), "gpt-test".to_string()))
);
assert_eq!(session.meta().await.thinking_level.as_deref(), Some("high"));
}
#[tokio::test]
async fn first_message_sets_truncated_title() {
let store = temp_store();
let s = Session::new(meta("s1"), vec![], "pi".into(), PathBuf::from("."), store.clone());
let long = "x".repeat(100);
s.push_user(&long).await;
let m = s.meta().await;
assert_eq!(m.title, "x".repeat(40));
let (m2, msgs) = store.load("s1").unwrap();
assert_eq!(m2.title, m.title);
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0].role, "user");
assert_eq!(msgs[0].text, long);
}
#[tokio::test]
async fn existing_messages_keep_original_title() {
let store = temp_store();
let msgs = vec![Msg {
role: "user".into(),
text: "old".into(),
ts: 1,
..Msg::default()
}];
let s = Session::new(meta("s2"), msgs, "pi".into(), PathBuf::from("."), store);
s.push_user("second message").await;
let m = s.meta().await;
assert_eq!(m.title, "New chat");
}
#[tokio::test]
async fn push_user_broadcasts_message_event() {
let store = temp_store();
let s = Session::new(meta("s3"), vec![], "pi".into(), PathBuf::from("."), store);
let mut rx = s.subscribe().await;
s.push_user("hello world").await;
let ev = rx.recv().await.expect("subscriber should get an event");
match ev {
SessionEvent::Message { role, text } => {
assert_eq!(role, "user");
assert_eq!(text, "hello world");
}
other => panic!("unexpected event: {other:?}"),
}
}
#[tokio::test]
async fn push_assistant_broadcasts_and_appends() {
let store = temp_store();
let s = Session::new(meta("s4"), vec![], "pi".into(), PathBuf::from("."), store);
s.push_user("q").await;
let mut rx = s.subscribe().await;
s.push_assistant("answer").await;
let ev = rx.recv().await.unwrap();
match ev {
SessionEvent::Message { role, text } => {
assert_eq!(role, "assistant");
assert_eq!(text, "answer");
}
other => panic!("unexpected event: {other:?}"),
}
let msgs = s.messages().await;
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[1].role, "assistant");
}
#[tokio::test]
async fn tool_events_merge_and_persist() {
let store = temp_store();
let s = Session::new(meta("tools"), vec![], "pi".into(), PathBuf::from("."), store.clone());
s.push_tool("call-1", "read", Some(json!({ "path": "src/main.rs" })), None, false).await;
s.push_tool("call-1", "read", None, Some(json!({ "content": [{ "type": "text", "text": "hello" }] })), false).await;
let messages = s.messages().await;
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, "tool");
assert_eq!(messages[0].call_id.as_deref(), Some("call-1"));
assert!(messages[0].input.is_some());
assert!(messages[0].result.is_some());
let (_, stored) = store.load("tools").expect("tool event should be persisted");
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].call_id.as_deref(), Some("call-1"));
assert!(stored[0].result.is_some());
}
#[tokio::test]
async fn reasoning_is_persisted() {
let store = temp_store();
let s = Session::new(meta("reasoning"), vec![], "pi".into(), PathBuf::from("."), store.clone());
s.push_reasoning("先检查输入,再生成答案").await;
let (_, stored) = store.load("reasoning").expect("reasoning should be persisted");
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].role, "reasoning");
assert_eq!(stored[0].text, "先检查输入,再生成答案");
}
#[tokio::test]
async fn rename_updates_meta_and_disk() {
let store = temp_store();
let s = Session::new(meta("s5"), vec![], "pi".into(), PathBuf::from("."), store.clone());
s.rename("My new title".to_string()).await;
assert_eq!(s.meta().await.title, "My new title");
let (m, _) = store.load("s5").unwrap();
assert_eq!(m.title, "My new title");
}
#[tokio::test]
async fn push_thinking_broadcasts() {
let store = temp_store();
let s = Session::new(meta("s6"), vec![], "pi".into(), PathBuf::from("."), store);
let mut rx = s.subscribe().await;
s.push_thinking(true).await;
assert_eq!(rx.recv().await.unwrap(), SessionEvent::Thinking { on: true });
s.push_thinking(false).await;
assert_eq!(rx.recv().await.unwrap(), SessionEvent::Thinking { on: false });
}
}