pub mod error;
pub use crate::error::TypeError;
use pyo3::prelude::*;
use schemars::JsonSchema;
use serde::de::Error;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::any::type_name;
use std::fmt;
use std::fmt::Display;
use std::path::{Path, PathBuf};
use tracing::error;
pub mod anthropic;
pub mod common;
pub mod google;
pub mod openai;
pub mod prompt;
pub mod spec;
pub mod tools;
pub mod traits;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
#[pyclass(from_py_object)]
pub enum Model {
Undefined,
}
impl Model {
pub fn as_str(&self) -> &str {
match self {
Model::Undefined => "undefined",
}
}
pub fn from_string(s: &str) -> Result<Self, TypeError> {
match s.to_lowercase().as_str() {
"undefined" => Ok(Model::Undefined),
_ => Err(TypeError::UnknownModelError(s.to_string())),
}
}
}
pub enum Common {
Undefined,
}
impl Common {
pub fn as_str(&self) -> &str {
match self {
Common::Undefined => "undefined",
}
}
}
impl Display for Common {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
#[pyclass(from_py_object, eq, eq_int)]
pub enum Provider {
OpenAI,
Gemini,
Google,
Vertex,
Anthropic,
GoogleAdk,
Undefined, }
impl Provider {
pub fn from_string(s: &str) -> Result<Self, TypeError> {
match s.to_lowercase().as_str() {
"openai" => Ok(Provider::OpenAI),
"gemini" => Ok(Provider::Gemini),
"google" => Ok(Provider::Google),
"vertex" => Ok(Provider::Vertex),
"anthropic" => Ok(Provider::Anthropic),
"google_adk" => Ok(Provider::GoogleAdk),
"undefined" => Ok(Provider::Undefined), _ => Err(TypeError::UnknownProviderError(s.to_string())),
}
}
pub fn extract_provider(provider: &Bound<'_, PyAny>) -> Result<Provider, TypeError> {
match provider.is_instance_of::<Provider>() {
true => Ok(provider.extract::<Provider>().inspect_err(|e| {
error!("Failed to extract provider: {}", e);
})?),
false => {
let provider = provider.extract::<String>().unwrap();
Ok(Provider::from_string(&provider).inspect_err(|e| {
error!("Failed to convert string to provider: {}", e);
})?)
}
}
}
pub fn as_str(&self) -> &str {
match self {
Provider::OpenAI => "openai",
Provider::Gemini => "gemini",
Provider::Vertex => "vertex",
Provider::Google => "google",
Provider::Anthropic => "anthropic",
Provider::GoogleAdk => "google_adk",
Provider::Undefined => "undefined", }
}
}
impl Display for Provider {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[pyclass(from_py_object, eq, eq_int)]
#[derive(Debug, PartialEq, Clone)]
pub enum SaveName {
Prompt,
}
#[pymethods]
impl SaveName {
#[staticmethod]
pub fn from_string(s: &str) -> Option<Self> {
match s {
"prompt" => Some(SaveName::Prompt),
_ => None,
}
}
pub fn as_string(&self) -> &str {
match self {
SaveName::Prompt => "prompt",
}
}
pub fn __str__(&self) -> String {
self.to_string()
}
}
impl Display for SaveName {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "{}", self.as_string())
}
}
impl AsRef<Path> for SaveName {
fn as_ref(&self) -> &Path {
match self {
SaveName::Prompt => Path::new("prompt"),
}
}
}
impl From<SaveName> for PathBuf {
fn from(save_name: SaveName) -> Self {
PathBuf::from(save_name.as_ref())
}
}
pub trait StructuredOutput: for<'de> serde::Deserialize<'de> + JsonSchema {
fn type_name() -> &'static str {
type_name::<Self>().rsplit("::").next().unwrap_or("Unknown")
}
fn model_validate_json_value(value: &Value) -> Result<Self, serde_json::Error> {
match &value {
Value::String(json_str) => Self::model_validate_json_str(json_str),
Value::Object(_) => {
serde_json::from_value(value.clone())
}
_ => {
Err(Error::custom("Expected a JSON string or object"))
}
}
}
fn model_validate_json_str(value: &str) -> Result<Self, serde_json::Error> {
serde_json::from_str(value)
}
fn get_structured_output_schema() -> Value {
let schema = ::schemars::schema_for!(Self);
schema.into()
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
#[pyclass(from_py_object)]
pub enum SettingsType {
GoogleChat,
OpenAIChat,
ModelSettings,
Anthropic,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_provider_google_adk_round_trip() {
let p = Provider::from_string("google_adk").unwrap();
assert_eq!(p, Provider::GoogleAdk);
assert_eq!(p.as_str(), "google_adk");
}
#[test]
fn test_provider_all_variants_round_trip() {
for (s, variant) in [
("openai", Provider::OpenAI),
("gemini", Provider::Gemini),
("google", Provider::Google),
("vertex", Provider::Vertex),
("anthropic", Provider::Anthropic),
("google_adk", Provider::GoogleAdk),
("undefined", Provider::Undefined),
] {
let parsed = Provider::from_string(s).unwrap();
assert_eq!(parsed, variant);
assert_eq!(parsed.as_str(), s);
}
}
#[test]
fn test_provider_unknown_string_errors() {
assert!(Provider::from_string("not_a_provider").is_err());
}
}