use std::path::{Path, PathBuf};
use std::sync::mpsc::{Receiver, Sender, channel};
use std::sync::{Mutex, MutexGuard};
use std::thread;
use crate::sdk::errors::{LimitError, NETWORK_ERROR_REASON, RunError, contains_http_status};
use crate::sdk::tool::{PiToolAdapter, SeherTool};
use crate::sdk::util::encode_session_id;
static PI_ENV_MUTEX: Mutex<()> = Mutex::new(());
struct PiEnvGuard {
_lock: MutexGuard<'static, ()>,
saved: Vec<(String, Option<String>)>,
}
impl PiEnvGuard {
fn acquire(env: &indexmap::IndexMap<String, String>) -> Self {
let lock: MutexGuard<'static, ()> = PI_ENV_MUTEX
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let saved: Vec<(String, Option<String>)> = env
.keys()
.map(|k| (k.clone(), std::env::var(k).ok()))
.collect();
for (k, v) in env {
unsafe { std::env::set_var(k, v) };
}
Self { _lock: lock, saved }
}
}
impl Drop for PiEnvGuard {
fn drop(&mut self) {
for (k, old_val) in &self.saved {
match old_val {
Some(v) => unsafe { std::env::set_var(k, v) },
None => unsafe { std::env::remove_var(k) },
}
}
}
}
const SKILLS_DIR: &str = ".agents/skills";
const PI_LIMIT_TOKENS: &[&str] = &[
"ratelimit",
"rate-limit",
"rate-limited",
"usagelimit",
"usage-limit",
"usage-limited",
"quota",
];
const PI_LIMIT_PHRASES: &[&str] = &["rate limit", "usage limit", "too many requests"];
fn is_pi_limit(msg: &str) -> bool {
if contains_http_status(msg, 429) {
return true;
}
let lower = msg.to_lowercase();
if PI_LIMIT_PHRASES.iter().any(|p| lower.contains(p)) {
return true;
}
lower
.split(|c: char| {
c.is_whitespace()
|| matches!(
c,
'(' | ')'
| '['
| ']'
| '{'
| '}'
| ','
| ';'
| ':'
| '.'
| '\''
| '"'
| '/'
| '\\'
| '!'
| '?'
)
})
.filter(|t| !t.is_empty())
.any(|t| PI_LIMIT_TOKENS.contains(&t))
}
#[derive(Debug)]
pub enum StreamChunk {
Delta(String),
Done(String),
Session(String),
Limit(LimitError),
Error(String),
}
#[derive(Clone, Default)]
pub struct PiRunnerOptions {
pub provider: Option<String>,
pub model: Option<String>,
pub api_key: Option<String>,
pub thinking: Option<String>,
pub system_prompt: Option<String>,
pub working_directory: Option<PathBuf>,
pub env: indexmap::IndexMap<String, String>,
pub tools: Vec<SeherTool>,
}
impl std::fmt::Debug for PiRunnerOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PiRunnerOptions")
.field("provider", &self.provider)
.field("model", &self.model)
.field("api_key", &self.api_key.as_ref().map(|_| "***"))
.field("thinking", &self.thinking)
.field("system_prompt", &self.system_prompt)
.field("working_directory", &self.working_directory)
.field("env", &self.env.keys().collect::<Vec<_>>())
.field(
"tools",
&self
.tools
.iter()
.map(|t| t.name.as_str())
.collect::<Vec<_>>(),
)
.finish()
}
}
#[derive(Debug, Clone)]
pub struct PiRunOutput {
pub text: String,
pub session_id: String,
}
fn encode_cwd_dir(cwd: &Path) -> String {
cwd.to_string_lossy()
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '-' {
c
} else {
'-'
}
})
.collect()
}
#[must_use]
pub fn pi_session_path(working_directory: Option<&Path>, id: &str) -> PathBuf {
let cwd = working_directory.map_or_else(
|| std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")),
Path::to_path_buf,
);
let cwd = std::fs::canonicalize(&cwd).unwrap_or(cwd);
let base = dirs::data_dir()
.unwrap_or_else(|| {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".local")
.join("share")
})
.join("seher")
.join("pi-sessions");
base.join(encode_cwd_dir(&cwd))
.join(format!("{}.jsonl", encode_session_id(id)))
}
fn seed_session_file(
path: &Path,
id: &str,
working_directory: Option<&Path>,
) -> std::io::Result<()> {
let mut header = pi::session::SessionHeader::new();
header.id = id.to_string();
if let Some(cwd) = working_directory {
header.cwd = cwd.display().to_string();
}
let line = serde_json::to_string(&header).map_err(std::io::Error::other)?;
std::fs::write(path, format!("{line}\n"))
}
#[must_use]
fn load_hardcoded_skills(working_directory: Option<&Path>) -> Option<String> {
let home = dirs::home_dir()?;
let skills_dir = home.join(SKILLS_DIR);
if !skills_dir.is_dir() {
return None;
}
let cwd = working_directory.map_or_else(
|| std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")),
Path::to_path_buf,
);
let options = pi::resources::LoadSkillsOptions {
cwd,
agent_dir: home.join(".agents"),
skill_paths: vec![skills_dir],
include_defaults: false,
};
let result = pi::resources::load_skills(options);
if result.skills.is_empty() {
return None;
}
Some(pi::resources::format_skills_for_prompt(&result.skills))
}
pub struct PiRunner {
opts: PiRunnerOptions,
}
impl PiRunner {
#[must_use]
pub fn new(opts: PiRunnerOptions) -> Self {
Self { opts }
}
#[must_use]
pub fn stream(&self, prompt: String, resume: Option<String>) -> Receiver<StreamChunk> {
let (tx, rx) = channel();
let opts = self.opts.clone();
thread::spawn(move || run_on_thread(&opts, &prompt, resume.as_deref(), &tx));
rx
}
pub fn run(&self, prompt: String, resume: Option<String>) -> Result<PiRunOutput, RunError> {
let rx = self.stream(prompt, resume);
let mut buffered = String::new();
let mut session_id = String::new();
loop {
match rx.recv() {
Ok(StreamChunk::Delta(d)) => buffered.push_str(&d),
Ok(StreamChunk::Session(id)) => session_id = id,
Ok(StreamChunk::Done(text)) => {
return Ok(PiRunOutput {
text: if text.is_empty() { buffered } else { text },
session_id,
});
}
Ok(StreamChunk::Limit(error)) => {
return Err(RunError::Limit {
error,
partial: buffered,
});
}
Ok(StreamChunk::Error(msg)) => {
return Err(RunError::Other {
message: msg,
partial: buffered,
});
}
Err(_) => {
return Err(RunError::Other {
message: "pi runner channel closed".to_string(),
partial: buffered,
});
}
}
}
}
}
fn validate_tool_names(tools: &[SeherTool]) -> Result<(), String> {
let mut seen = std::collections::HashSet::new();
for tool in tools {
if pi::sdk::BUILTIN_TOOL_NAMES.contains(&tool.name.as_str()) {
return Err(format!(
"custom tool '{}' collides with a pi built-in tool ({})",
tool.name,
pi::sdk::BUILTIN_TOOL_NAMES.join(", ")
));
}
if !seen.insert(tool.name.as_str()) {
return Err(format!("duplicate custom tool name '{}'", tool.name));
}
}
Ok(())
}
#[must_use]
pub fn split_thinking_suffix(model: &str) -> (&str, Option<&str>) {
match model.rsplit_once(':') {
Some((m, t))
if t.parse::<pi::model::ThinkingLevel>().is_ok() || t.eq_ignore_ascii_case("max") =>
{
(m, Some(t))
}
_ => (model, None),
}
}
#[must_use]
pub fn split_model_ref(
fallback_provider: &str,
model_id: &str,
) -> (String, String, Option<String>) {
let (provider, model_rest) = match model_id.split_once('/') {
Some((p, m)) => (p.to_string(), m),
None => (fallback_provider.to_string(), model_id),
};
let (model, thinking) = split_thinking_suffix(model_rest);
(provider, model.to_string(), thinking.map(str::to_string))
}
fn parse_thinking(thinking: Option<&str>) -> Result<Option<pi::model::ThinkingLevel>, String> {
thinking
.map(|t| {
t.parse::<pi::model::ThinkingLevel>()
.map_err(|e| format!("invalid thinking level '{t}': {e}"))
})
.transpose()
}
fn run_on_thread(
opts: &PiRunnerOptions,
prompt: &str,
resume: Option<&str>,
tx: &Sender<StreamChunk>,
) {
use pi::model::AssistantMessageEvent;
use pi::sdk::{AgentEvent, SessionOptions, create_agent_session};
if let Err(msg) = validate_tool_names(&opts.tools) {
let _ = tx.send(StreamChunk::Error(msg));
return;
}
let thinking = match parse_thinking(opts.thinking.as_deref()) {
Ok(t) => t,
Err(msg) => {
let _ = tx.send(StreamChunk::Error(msg));
return;
}
};
for k in opts.env.keys() {
if k.contains('=') || k.contains('\0') {
let _ = tx.send(StreamChunk::Error(format!(
"invalid env key '{k}': must not contain '=' or NUL"
)));
return;
}
}
let session_id = resume.map_or_else(|| uuid::Uuid::new_v4().to_string(), str::to_string);
let session_path = pi_session_path(opts.working_directory.as_deref(), &session_id);
if resume.is_none() {
let created = session_path
.parent()
.map_or(Ok(()), std::fs::create_dir_all)
.and_then(|()| {
seed_session_file(
&session_path,
&session_id,
opts.working_directory.as_deref(),
)
});
if let Err(e) = created {
let _ = tx.send(StreamChunk::Error(format!(
"failed to create session file {}: {e}",
session_path.display()
)));
return;
}
}
let _ = tx.send(StreamChunk::Session(session_id.clone()));
let prompt_text = prompt.to_string();
let skills_appendix = load_hardcoded_skills(opts.working_directory.as_deref());
let provider_label = opts.provider.clone().unwrap_or_else(|| "pi".to_string());
let _env_guard = (!opts.env.is_empty()).then(|| PiEnvGuard::acquire(&opts.env));
let outcome: Result<(), CloseOutcome> = futures::executor::block_on(async {
let session_opts = SessionOptions {
provider: opts.provider.clone(),
model: opts.model.clone(),
api_key: opts.api_key.clone(),
thinking,
system_prompt: opts.system_prompt.clone(),
append_system_prompt: skills_appendix,
working_directory: opts.working_directory.clone(),
no_session: false,
session_path: Some(session_path.clone()),
..Default::default()
};
let mut handle = create_agent_session(session_opts)
.await
.map_err(|e| CloseOutcome::Error(format!("create_agent_session failed: {e}")))?;
if !opts.tools.is_empty() {
let custom: Vec<Box<dyn pi::tools::Tool>> = opts
.tools
.iter()
.cloned()
.map(|t| Box::new(PiToolAdapter::new(t)) as Box<dyn pi::tools::Tool>)
.collect();
handle.session_mut().agent.extend_tools(custom);
}
let txd = tx.clone();
let assistant = handle
.prompt(&prompt_text, move |ev: AgentEvent| {
if let AgentEvent::MessageUpdate {
assistant_message_event,
..
} = ev
&& let AssistantMessageEvent::TextDelta { delta, .. } = assistant_message_event
{
let _ = txd.send(StreamChunk::Delta(delta));
}
})
.await
.map_err(|e| classify_pi_error(&provider_label, &e.to_string()))?;
check_trailing_assistant_error(&provider_label, &assistant)
});
match outcome {
Ok(()) => {
let _ = tx.send(StreamChunk::Done(String::new()));
}
Err(CloseOutcome::Limit(e)) => {
let _ = tx.send(StreamChunk::Limit(e));
}
Err(CloseOutcome::Error(msg)) => {
let _ = tx.send(StreamChunk::Error(msg));
}
}
}
enum CloseOutcome {
Limit(LimitError),
Error(String),
}
fn trailing_assistant_error(assistant: &pi::model::AssistantMessage) -> Option<String> {
if matches!(assistant.stop_reason, pi::model::StopReason::Error) {
let error = assistant
.error_message
.clone()
.unwrap_or_else(|| "pi: assistant turn ended with stopReason error".to_string());
Some(if error == NETWORK_ERROR_REASON {
format!("provider error: {error}")
} else {
error
})
} else {
None
}
}
fn check_trailing_assistant_error(
provider: &str,
assistant: &pi::model::AssistantMessage,
) -> Result<(), CloseOutcome> {
match trailing_assistant_error(assistant) {
Some(msg) => Err(classify_pi_error(provider, &msg)),
None => Ok(()),
}
}
fn classify_pi_error(provider: &str, msg: &str) -> CloseOutcome {
if is_pi_limit(msg) {
CloseOutcome::Limit(LimitError {
provider: provider.to_string(),
reset_at: None,
})
} else {
CloseOutcome::Error(msg.to_string())
}
}
#[cfg(test)]
#[expect(clippy::expect_used, reason = "tests may panic on unexpected fixtures")]
mod tests {
use super::*;
#[test]
fn detects_common_limit_phrases() {
assert!(is_pi_limit("Rate limit exceeded"));
assert!(is_pi_limit("usage-limit"));
assert!(is_pi_limit("HTTP 429 Too Many Requests"));
assert!(is_pi_limit("Quota exceeded for the day"));
}
#[test]
fn rejects_unrelated_messages() {
assert!(!is_pi_limit("unexpected end of stream"));
assert!(!is_pi_limit("connection refused"));
}
#[test]
fn trailing_network_marker_text_stays_an_ordinary_error() {
let assistant = pi::model::AssistantMessage {
stop_reason: pi::model::StopReason::Error,
error_message: Some(NETWORK_ERROR_REASON.to_string()),
..Default::default()
};
let msg = trailing_assistant_error(&assistant).expect("error stop must be detected");
let close = classify_pi_error("opencode-go", &msg);
let CloseOutcome::Error(msg) = close else {
panic!("network marker text must remain an ordinary error");
};
let error = RunError::Other {
message: msg,
partial: "partial".to_string(),
};
assert!(matches!(
error,
RunError::Other { message, partial }
if message == "provider error: network_error" && partial == "partial"
));
}
#[test]
fn trailing_assistant_error_with_limit_message_classifies_as_limit() {
let assistant = pi::model::AssistantMessage {
stop_reason: pi::model::StopReason::Error,
error_message: Some("429 Weekly usage limit reached. Resets in 1 day.".to_string()),
..Default::default()
};
let msg = trailing_assistant_error(&assistant).expect("error stop must be detected");
assert!(matches!(
classify_pi_error("opencode-go", &msg),
CloseOutcome::Limit(_)
));
}
#[test]
fn trailing_assistant_error_with_other_message_classifies_as_error() {
let assistant = pi::model::AssistantMessage {
stop_reason: pi::model::StopReason::Error,
error_message: Some("500 internal server error".to_string()),
..Default::default()
};
let msg = trailing_assistant_error(&assistant).expect("error stop must be detected");
assert!(matches!(
classify_pi_error("opencode-go", &msg),
CloseOutcome::Error(m) if m.contains("500")
));
}
#[test]
fn trailing_assistant_error_ignores_normal_stop() {
let assistant = pi::model::AssistantMessage::default();
assert!(trailing_assistant_error(&assistant).is_none());
}
#[test]
fn parse_thinking_accepts_known_levels() {
assert_eq!(parse_thinking(None), Ok(None));
assert_eq!(
parse_thinking(Some("high")),
Ok(Some(pi::model::ThinkingLevel::High))
);
assert_eq!(
parse_thinking(Some("off")),
Ok(Some(pi::model::ThinkingLevel::Off))
);
}
#[test]
fn split_thinking_suffix_extracts_known_levels() {
assert_eq!(
split_thinking_suffix("opus-4.7:high"),
("opus-4.7", Some("high"))
);
assert_eq!(
split_thinking_suffix("opus-4.7:med"),
("opus-4.7", Some("med"))
);
assert_eq!(
split_thinking_suffix("llama-3.1:free:low"),
("llama-3.1:free", Some("low"))
);
}
#[test]
fn split_thinking_suffix_keeps_unrecognized_suffix() {
assert_eq!(
split_thinking_suffix("meta-llama/llama-3.1-8b-instruct:free"),
("meta-llama/llama-3.1-8b-instruct:free", None)
);
assert_eq!(split_thinking_suffix("opus-4.7"), ("opus-4.7", None));
assert_eq!(split_thinking_suffix("opus-4.7:"), ("opus-4.7:", None));
}
#[test]
fn split_model_ref_extracts_provider_model_and_thinking() {
assert_eq!(
split_model_ref("codex", "openai-codex/gpt-5.5:xhigh"),
(
"openai-codex".to_string(),
"gpt-5.5".to_string(),
Some("xhigh".to_string())
)
);
}
#[test]
fn split_model_ref_keeps_extra_slashes_in_model() {
assert_eq!(
split_model_ref("openrouter", "openrouter/moonshotai/kimi-k2.6"),
(
"openrouter".to_string(),
"moonshotai/kimi-k2.6".to_string(),
None
)
);
}
#[test]
fn split_model_ref_uses_fallback_provider_when_no_slash() {
assert_eq!(
split_model_ref("anthropic", "claude-sonnet-4-5"),
(
"anthropic".to_string(),
"claude-sonnet-4-5".to_string(),
None
)
);
}
#[test]
fn split_model_ref_does_not_strip_non_thinking_colon_suffix() {
assert_eq!(
split_model_ref("openrouter", "openrouter/meta-llama/llama-3-8b:free"),
(
"openrouter".to_string(),
"meta-llama/llama-3-8b:free".to_string(),
None
)
);
}
#[test]
fn parse_thinking_rejects_unknown_level() {
let err = parse_thinking(Some("turbo")).expect_err("'turbo' must not parse");
assert!(err.contains("invalid thinking level 'turbo'"), "{err}");
}
#[test]
fn rejects_substring_false_positives() {
assert!(!is_pi_limit("Read 5429 bytes before EOF"));
assert!(!is_pi_limit("loaded squotahelper module"));
assert!(!is_pi_limit("status 429 returned"));
assert!(!is_pi_limit("request 429 of 1000 completed"));
}
#[test]
fn detects_http_429_context() {
assert!(is_pi_limit("Kimi API error (HTTP 429): server busy"));
assert!(is_pi_limit(
"Anthropic API error (HTTP 429): {\"error\":{\"type\":\"rate_limit_error\"}}"
));
}
#[test]
fn rejects_http_status_with_trailing_digits() {
assert!(!is_pi_limit("Kimi API error (HTTP 4290): oversized"));
}
#[test]
#[expect(clippy::unwrap_used, reason = "test panics on unexpected fixtures")]
fn seeded_session_file_is_openable_by_pi() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("seeded.jsonl");
let id = "11111111-2222-3333-4444-555555555555";
seed_session_file(&path, id, Some(dir.path())).unwrap();
let session =
futures::executor::block_on(pi::session::Session::open(&path.to_string_lossy()))
.unwrap();
assert_eq!(session.header.id, id);
assert_eq!(session.header.cwd, dir.path().display().to_string());
assert!(session.entries.is_empty());
}
#[test]
fn pi_runner_options_default_has_no_tools() {
assert!(PiRunnerOptions::default().tools.is_empty());
}
#[test]
fn pi_runner_options_debug_masks_api_key_and_lists_tool_names() {
let opts = PiRunnerOptions {
api_key: Some("sk-secret".to_string()),
tools: vec![SeherTool::new(
"echo",
"Echo",
serde_json::json!({"type": "object"}),
std::sync::Arc::new(|_| Ok(String::new())),
)],
..Default::default()
};
let dbg = format!("{opts:?}");
assert!(!dbg.contains("sk-secret"), "got: {dbg}");
assert!(dbg.contains("***"), "got: {dbg}");
assert!(dbg.contains("echo"), "got: {dbg}");
}
fn named_tool(name: &str) -> SeherTool {
SeherTool::new(
name,
"test tool",
serde_json::json!({"type": "object"}),
std::sync::Arc::new(|_| Ok(String::new())),
)
}
#[test]
fn validate_tool_names_accepts_unique_custom_names() {
let tools = vec![named_tool("alpha"), named_tool("beta")];
assert!(validate_tool_names(&tools).is_ok());
assert!(validate_tool_names(&[]).is_ok());
}
#[test]
fn validate_tool_names_rejects_builtin_collision() {
let err = validate_tool_names(&[named_tool("read")]).expect_err("should reject");
assert!(err.contains("read"), "got: {err}");
assert!(err.contains("built-in"), "got: {err}");
}
#[test]
fn validate_tool_names_rejects_duplicates() {
let err = validate_tool_names(&[named_tool("alpha"), named_tool("alpha")])
.expect_err("should reject");
assert!(err.contains("duplicate"), "got: {err}");
assert!(err.contains("alpha"), "got: {err}");
}
#[test]
fn stream_emits_error_on_builtin_tool_collision() {
let runner = PiRunner::new(PiRunnerOptions {
tools: vec![named_tool("bash")],
..Default::default()
});
let rx = runner.stream("hi".to_string(), None);
match rx.recv().expect("one chunk") {
StreamChunk::Error(msg) => assert!(msg.contains("bash"), "got: {msg}"),
other => panic!("expected Error chunk, got {other:?}"),
}
}
#[test]
fn pi_session_path_is_deterministic_for_same_cwd_and_id() {
let dir = std::env::temp_dir();
let a = pi_session_path(Some(&dir), "abc");
let b = pi_session_path(Some(&dir), "abc");
assert_eq!(a, b);
assert!(a.to_string_lossy().ends_with("abc.jsonl"));
}
#[test]
fn pi_session_path_canonicalizes_symlinked_cwd() {
let dir = std::env::temp_dir();
let canonical = std::fs::canonicalize(&dir).unwrap_or_else(|_| dir.clone());
assert_eq!(
pi_session_path(Some(&dir), "abc"),
pi_session_path(Some(&canonical), "abc"),
);
}
#[test]
fn pi_session_path_sanitizes_session_id() {
let dir = std::env::temp_dir();
let path = pi_session_path(Some(&dir), "../etc/passwd");
let file_name = path.file_name().expect("file name").to_string_lossy();
assert!(
!path.to_string_lossy().contains("../"),
"path must not contain traversal: {path:?}"
);
assert_eq!(file_name, "---etc-passwd.jsonl");
}
}