hey 0.1.0

Minimal terminal AI coding agent: kernel loop + MCP/Skills self-evolution
Documentation
//! 测试公共模块:自写 tiny HTTP server(模拟各协议 wire)+ FakeProvider。
//!
//! 设计(docs/DESIGN.md §14):测试不碰真实网络,`cargo test` 全绿是里程碑验收项。
#![allow(dead_code)] // 每个测试 binary 单独编译本模块,未用方法属正常

use std::collections::HashMap;
use std::fmt::Write as _;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};

use async_trait::async_trait;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::net::{TcpListener, TcpStream};

use hey::llm::ir::{ChatRequest, Completion, ToolCall, Usage};
use hey::llm::{Delta, LlmError, Provider};

// ========== MockServer ==========

/// 捕获到的请求。
#[derive(Debug, Clone)]
pub struct MockRequest {
    pub method: String,
    pub path: String,
    pub headers: HashMap<String, String>,
    pub body: String,
}

impl MockRequest {
    pub fn header(&self, name: &str) -> Option<&str> {
        self.headers
            .iter()
            .find(|(k, _)| k.eq_ignore_ascii_case(name))
            .map(|(_, v)| v.as_str())
    }
}

/// 预设响应。
#[derive(Debug, Clone)]
pub struct MockResponse {
    pub status: u16,
    pub headers: Vec<(String, String)>,
    pub body: String,
}

impl MockResponse {
    pub fn sse(body: &str) -> Self {
        Self {
            status: 200,
            headers: vec![
                ("Content-Type".into(), "text/event-stream".into()),
                ("Cache-Control".into(), "no-cache".into()),
            ],
            body: body.to_string(),
        }
    }

    pub fn status(status: u16, body: &str) -> Self {
        Self {
            status,
            headers: vec![("Content-Type".into(), "application/json".into())],
            body: body.to_string(),
        }
    }
}

/// 自写 tiny HTTP/1.1 服务器:按请求序号循环返回预设响应,记录所有请求。
pub struct MockServer {
    pub addr: SocketAddr,
    pub requests: Arc<Mutex<Vec<MockRequest>>>,
    responses: Arc<Mutex<Vec<MockResponse>>>,
    counter: Arc<AtomicUsize>,
    shutdown: Option<tokio::sync::oneshot::Sender<()>>,
}

impl MockServer {
    pub async fn start(responses: Vec<MockResponse>) -> Self {
        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
        let addr = listener.local_addr().unwrap();
        let requests: Arc<Mutex<Vec<MockRequest>>> = Arc::new(Mutex::new(Vec::new()));
        let responses = Arc::new(Mutex::new(responses));
        let counter = Arc::new(AtomicUsize::new(0));
        let (tx, mut rx) = tokio::sync::oneshot::channel::<()>();

        let reqs = requests.clone();
        let resps = responses.clone();
        let cnt = counter.clone();
        tokio::spawn(async move {
            loop {
                tokio::select! {
                    _ = &mut rx => break,
                    accepted = listener.accept() => {
                        let Ok((stream, _)) = accepted else { continue };
                        let reqs = reqs.clone();
                        let resps = resps.clone();
                        let cnt = cnt.clone();
                        tokio::spawn(async move {
                            let _ = handle_conn(stream, reqs, resps, cnt).await;
                        });
                    }
                }
            }
        });

        Self {
            addr,
            requests,
            responses,
            counter,
            shutdown: Some(tx),
        }
    }

    /// 已收到的请求数。
    pub fn count(&self) -> usize {
        self.requests.lock().unwrap().len()
    }

    /// 第 i 个请求(按到达顺序)。
    pub fn request(&self, i: usize) -> Option<MockRequest> {
        self.requests.lock().unwrap().get(i).cloned()
    }
}

impl Drop for MockServer {
    fn drop(&mut self) {
        if let Some(tx) = self.shutdown.take() {
            let _ = tx.send(());
        }
    }
}

async fn handle_conn(
    mut stream: TcpStream,
    reqs: Arc<Mutex<Vec<MockRequest>>>,
    resps: Arc<Mutex<Vec<MockResponse>>>,
    counter: Arc<AtomicUsize>,
) -> std::io::Result<()> {
    let mut reader = BufReader::new(&mut stream);
    let mut head = String::new();
    // 读请求头(直到空行)
    loop {
        let mut line = String::new();
        let n = reader.read_line(&mut line).await?;
        if n == 0 {
            return Ok(());
        }
        head.push_str(&line);
        if line == "\r\n" || line == "\n" {
            break;
        }
    }

    // 解析请求行 + 头
    let mut lines = head.lines();
    let request_line = lines.next().unwrap_or_default().to_string();
    let mut parts = request_line.split_whitespace();
    let method = parts.next().unwrap_or("").to_string();
    let path = parts.next().unwrap_or("").to_string();
    let mut headers = HashMap::new();
    let mut content_length = 0usize;
    for l in lines {
        if let Some((k, v)) = l.split_once(':') {
            headers.insert(k.trim().to_string(), v.trim().to_string());
            if k.trim().eq_ignore_ascii_case("content-length") {
                content_length = v.trim().parse().unwrap_or(0);
            }
        }
    }

    // 读 body
    let mut body = String::new();
    if content_length > 0 {
        let mut buf = vec![0u8; content_length];
        reader.read_exact(&mut buf).await?;
        body = String::from_utf8_lossy(&buf).to_string();
    }

    // 记录请求
    reqs.lock().unwrap().push(MockRequest {
        method,
        path,
        headers,
        body,
    });

    // 取响应(按序号循环,最后一条兜底)
    let idx = counter.fetch_add(1, Ordering::SeqCst);
    let resp = {
        let r = resps.lock().unwrap();
        r.get(idx).or_else(|| r.last()).cloned()
    };
    let Some(resp) = resp else {
        return Ok(());
    };

    let reason = match resp.status {
        200 => "OK",
        400 => "Bad Request",
        401 => "Unauthorized",
        402 => "Payment Required",
        429 => "Too Many Requests",
        500 => "Internal Server Error",
        503 => "Service Unavailable",
        _ => "OK",
    };
    let mut out = format!("HTTP/1.1 {} {}\r\n", resp.status, reason);
    for (k, v) in &resp.headers {
        let _ = write!(out, "{k}: {v}\r\n");
    }
    let _ = write!(out, "Content-Length: {}\r\n", resp.body.len());
    out.push_str("\r\n");
    out.push_str(&resp.body);
    stream.write_all(out.as_bytes()).await?;
    stream.flush().await?;
    Ok(())
}

// ========== FakeProvider ==========

/// 一轮预设响应。
#[derive(Debug, Clone)]
pub struct FakeTurn {
    pub text: &'static str,
    pub tool_calls: Vec<ToolCall>,
}

/// 内存 Provider:按调用次数播放脚本(超出循环最后一条)。
pub struct FakeProvider {
    pub model: String,
    pub script: Vec<FakeTurn>,
    pub calls: Arc<AtomicUsize>,
}

impl FakeProvider {
    pub fn new(model: &str, script: Vec<FakeTurn>) -> Self {
        Self {
            model: model.to_string(),
            script,
            calls: Arc::new(AtomicUsize::new(0)),
        }
    }

    pub fn call_count(&self) -> usize {
        self.calls.load(Ordering::SeqCst)
    }
}

#[async_trait]
impl Provider for FakeProvider {
    fn model(&self) -> &str {
        &self.model
    }

    async fn stream(
        &self,
        _req: &ChatRequest,
        on_delta: &mut (dyn FnMut(Delta) + Send),
    ) -> Result<Completion, LlmError> {
        let i = self.calls.fetch_add(1, Ordering::SeqCst);
        let turn = &self.script[i.min(self.script.len() - 1)];
        if !turn.text.is_empty() {
            on_delta(Delta::Text(turn.text.to_string()));
        }
        Ok(Completion {
            text: turn.text.to_string(),
            thinking: String::new(),
            tool_calls: turn.tool_calls.clone(),
            usage: Usage {
                input: 10,
                output: turn.text.len() as u64,
            },
        })
    }
}

/// 可注入错误的 Provider:脚本为 `Result<FakeTurn, LlmError>`,按调用次数播放。
/// 用于 turn 级重试测试(agent.rs 的「整轮无部分响应才重试」)。
pub struct FlakyProvider {
    pub model: String,
    pub script: Vec<Result<FakeTurn, LlmError>>,
    pub calls: Arc<AtomicUsize>,
    /// 出错前先吐一段 delta(模拟流中断:有部分响应)。
    pub partial_before_err: bool,
}

impl FlakyProvider {
    pub fn new(model: &str, script: Vec<Result<FakeTurn, LlmError>>) -> Self {
        Self {
            model: model.to_string(),
            script,
            calls: Arc::new(AtomicUsize::new(0)),
            partial_before_err: false,
        }
    }

    pub fn call_count(&self) -> usize {
        self.calls.load(Ordering::SeqCst)
    }
}

#[async_trait]
impl Provider for FlakyProvider {
    fn model(&self) -> &str {
        &self.model
    }

    async fn stream(
        &self,
        _req: &ChatRequest,
        on_delta: &mut (dyn FnMut(Delta) + Send),
    ) -> Result<Completion, LlmError> {
        let i = self.calls.fetch_add(1, Ordering::SeqCst);
        let step = &self.script[i.min(self.script.len() - 1)];
        match step {
            Ok(turn) => {
                if !turn.text.is_empty() {
                    on_delta(Delta::Text(turn.text.to_string()));
                }
                Ok(Completion {
                    text: turn.text.to_string(),
                    thinking: String::new(),
                    tool_calls: turn.tool_calls.clone(),
                    usage: Usage::default(),
                })
            }
            Err(e) => {
                if self.partial_before_err {
                    on_delta(Delta::Text("partial".to_string()));
                }
                Err(e.clone())
            }
        }
    }
}

/// 工具调用便捷构造。
pub fn tool_call(id: &str, name: &str, args: &str) -> ToolCall {
    ToolCall {
        id: id.to_string(),
        name: name.to_string(),
        arguments: args.to_string(),
    }
}