use std::sync::mpsc::{Receiver, Sender, channel};
use std::thread;
use crate::sdk::errors::{LimitError, RunError};
const PI_LIMIT_TOKENS: &[&str] = &[
"429",
"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 {
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),
Limit(LimitError),
Error(String),
}
#[derive(Clone, Default)]
pub struct PiRunnerOptions {
pub provider: Option<String>,
pub model: Option<String>,
pub api_key: Option<String>,
pub system_prompt: Option<String>,
}
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("system_prompt", &self.system_prompt)
.finish()
}
}
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) -> Receiver<StreamChunk> {
let (tx, rx) = channel();
let opts = self.opts.clone();
thread::spawn(move || run_on_thread(&opts, &prompt, &tx));
rx
}
pub fn run(&self, prompt: String) -> Result<String, RunError> {
let rx = self.stream(prompt);
let mut buffered = String::new();
loop {
match rx.recv() {
Ok(StreamChunk::Delta(d)) => buffered.push_str(&d),
Ok(StreamChunk::Done(text)) => {
return Ok(if text.is_empty() { buffered } else { text });
}
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 run_on_thread(opts: &PiRunnerOptions, prompt: &str, tx: &Sender<StreamChunk>) {
use pi::model::AssistantMessageEvent;
use pi::sdk::{AgentEvent, SessionOptions, create_agent_session};
let prompt_text = match opts.system_prompt.as_deref() {
Some(sys) => format!("{sys}\n\n{prompt}"),
None => prompt.to_string(),
};
let provider_label = opts.provider.clone().unwrap_or_else(|| "pi".to_string());
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(),
no_session: true,
..Default::default()
};
let mut handle = create_agent_session(session_opts)
.await
.map_err(|e| CloseOutcome::Error(format!("create_agent_session failed: {e}")))?;
let txd = tx.clone();
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()))?;
Ok(())
});
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 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)]
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 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"));
}
}