kaynine-core 0.1.0

Core agent loop, messages, events, policies, and provider abstractions for Kaynine
Documentation
use crate::error::{CredentialError, ProviderError, RunFailureReason};
use crate::ids::{ModelId, ProviderId};
use crate::message::{ContentBlock, FinishReason, ToolResultPayload, Usage};
use crate::tool::ToolSpec;
use async_trait::async_trait;
use futures_core::Stream;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::pin::Pin;
use tokio_util::sync::CancellationToken;

pub type ProviderStream = Pin<Box<dyn Stream<Item = Result<ProviderEvent, ProviderError>> + Send>>;

#[async_trait]
pub trait ModelProvider: Send + Sync {
    fn id(&self) -> ProviderId;

    async fn capabilities(
        &self,
        model: &ModelId,
        credentials: &dyn CredentialProvider,
        cancel: CancellationToken,
    ) -> Result<ModelCapabilities, ProviderError>;

    async fn stream(
        &self,
        request: ModelRequest,
        credentials: &dyn CredentialProvider,
        cancel: CancellationToken,
    ) -> Result<ProviderStream, ProviderError>;
}

#[async_trait]
pub trait TokenCounter: Send + Sync {
    async fn count_input(
        &self,
        request: &ModelRequest,
        credentials: &dyn CredentialProvider,
        cancel: CancellationToken,
    ) -> Result<TokenMeasurement, ProviderError>;
}

#[async_trait]
pub trait CredentialProvider: Send + Sync {
    async fn resolve(&self, request: CredentialRequest) -> Result<Credentials, CredentialError>;
}

#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct CredentialRequest {
    pub provider_id: ProviderId,
    pub model: ModelId,
}

#[derive(Clone, Debug, PartialEq, Eq, Default, Serialize, Deserialize)]
pub struct Credentials {
    pub bearer: Option<String>,
    pub headers: Vec<(String, String)>,
}

#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ModelRequest {
    pub model: ModelId,
    pub system_prompt: String,
    pub messages: Vec<ModelMessage>,
    pub tools: Vec<ToolSpec>,
    pub reasoning: ReasoningLevel,
    pub generation: GenerationOptions,
    pub provider_options: Value,
}

#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "role")]
pub enum ModelMessage {
    User {
        blocks: Vec<ContentBlock>,
    },
    Assistant {
        blocks: Vec<ContentBlock>,
    },
    ToolResults {
        results: Vec<ToolResultPayload>,
    },
    /// Compaction checkpoint (SPEC §4.4): projected to providers as a user
    /// text message carrying the covered prefix's summary.
    Summary {
        text: String,
    },
}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum ReasoningLevel {
    Off,
    Low,
    Medium,
    High,
}

#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct GenerationOptions {
    pub max_output_tokens: Option<u64>,
    pub temperature: Option<f64>,
}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelCapabilities {
    pub context_tokens: u32,
    pub max_output_tokens: u32,
    pub supports_tools: bool,
    pub supports_images: bool,
    pub supports_reasoning: bool,
}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum TokenMeasurementSource {
    ExactProvider,
    ExactLocal,
    UsageCalibrated,
    Heuristic,
}

impl std::fmt::Display for TokenMeasurementSource {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(f, "{self:?}")
    }
}

impl std::error::Error for TokenMeasurementSource {}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TokenMeasurement {
    pub input_tokens: u64,
    pub source: TokenMeasurementSource,
    pub safety_margin_tokens: u64,
}

#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ProviderEvent {
    ResponseStarted,
    TextDelta {
        block: u32,
        text: String,
    },
    ReasoningDelta {
        block: u32,
        text: String,
    },
    ToolCallStarted {
        block: u32,
        id: String,
        name: String,
    },
    ToolCallArgumentsDelta {
        block: u32,
        json: String,
    },
    UsageUpdated(Usage),
    ResponseCompleted {
        finish_reason: FinishReason,
    },
}

pub fn validate_request_capabilities(
    provider_id: &ProviderId,
    request: &ModelRequest,
    capabilities: &ModelCapabilities,
) -> Result<(), RunFailureReason> {
    let mismatch = |requirement: &str| RunFailureReason::ModelCapabilityMismatch {
        requirement: requirement.to_string(),
        provider: provider_id.clone(),
        model: request.model.clone(),
    };
    if !request.tools.is_empty() && !capabilities.supports_tools {
        return Err(mismatch("tools"));
    }
    let has_image = request.messages.iter().any(|m| match m {
        ModelMessage::User { blocks } | ModelMessage::Assistant { blocks } => blocks
            .iter()
            .any(|b| matches!(b, ContentBlock::Image { .. })),
        ModelMessage::ToolResults { .. } | ModelMessage::Summary { .. } => false,
    });
    if has_image && !capabilities.supports_images {
        return Err(mismatch("images"));
    }
    if request.reasoning != ReasoningLevel::Off && !capabilities.supports_reasoning {
        return Err(mismatch("reasoning"));
    }
    Ok(())
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::message::BinaryRef;
    use crate::tool::Concurrency;

    fn caps(
        supports_tools: bool,
        supports_images: bool,
        supports_reasoning: bool,
    ) -> ModelCapabilities {
        ModelCapabilities {
            context_tokens: 200_000,
            max_output_tokens: 8_192,
            supports_tools,
            supports_images,
            supports_reasoning,
        }
    }

    fn request_with_tools(n: usize, reasoning: ReasoningLevel, image: bool) -> ModelRequest {
        let mut blocks = vec![ContentBlock::Text { text: "q".into() }];
        if image {
            blocks.push(ContentBlock::Image {
                media_type: "image/png".into(),
                data: BinaryRef {
                    sha256: "h".into(),
                    byte_len: 1,
                },
            });
        }
        ModelRequest {
            model: ModelId::from("m"),
            system_prompt: String::new(),
            messages: vec![ModelMessage::User { blocks }],
            tools: (0..n)
                .map(|i| ToolSpec {
                    name: format!("t{i}"),
                    description: String::new(),
                    parameters_schema: serde_json::json!({"type": "object"}),
                    concurrency: Concurrency::Sequential,
                })
                .collect(),
            reasoning,
            generation: GenerationOptions::default(),
            provider_options: serde_json::Value::Null,
        }
    }

    #[test]
    fn rejects_tools_when_unsupported() {
        let err = validate_request_capabilities(
            &ProviderId::from(ProviderId::ANTHROPIC),
            &request_with_tools(1, ReasoningLevel::Off, false),
            &caps(false, true, true),
        )
        .unwrap_err();
        assert!(
            matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "tools")
        );
    }

    #[test]
    fn rejects_images_when_unsupported() {
        let err = validate_request_capabilities(
            &ProviderId::from(ProviderId::ANTHROPIC),
            &request_with_tools(0, ReasoningLevel::Off, true),
            &caps(true, false, true),
        )
        .unwrap_err();
        assert!(
            matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "images")
        );
    }

    #[test]
    fn rejects_reasoning_when_unsupported() {
        let err = validate_request_capabilities(
            &ProviderId::from(ProviderId::ANTHROPIC),
            &request_with_tools(0, ReasoningLevel::High, false),
            &caps(true, true, false),
        )
        .unwrap_err();
        assert!(
            matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "reasoning")
        );
    }

    #[test]
    fn accepts_matching_request() {
        assert!(validate_request_capabilities(
            &ProviderId::from(ProviderId::ANTHROPIC),
            &request_with_tools(2, ReasoningLevel::Medium, true),
            &caps(true, true, true),
        )
        .is_ok());
    }

    #[test]
    fn request_roundtrips_through_serde() {
        let req = request_with_tools(1, ReasoningLevel::Off, false);
        let json = serde_json::to_string(&req).unwrap();
        let back: ModelRequest = serde_json::from_str(&json).unwrap();
        assert_eq!(back, req);
    }

    #[test]
    fn summary_message_roundtrips_through_serde() {
        let message = ModelMessage::Summary {
            text: "摘要".into(),
        };
        let json = serde_json::to_string(&message).unwrap();
        assert!(json.contains(r#""role""#));
        let back: ModelMessage = serde_json::from_str(&json).unwrap();
        assert_eq!(back, message);
    }
}