use std::borrow::Borrow;
use std::fmt;
use std::str::FromStr;
use std::sync::Arc;
use serde::{Deserialize, Serialize};
use thiserror::Error;
macro_rules! arc_str_newtype {
(
$(#[$struct_doc:meta])*
struct $name:ident;
new_doc: $(#[$new_doc:meta])*
as_str_doc: $(#[$as_str_doc:meta])*
default_doc: $(#[$default_doc:meta])*
) => {
$(#[$struct_doc])*
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(transparent)]
pub struct $name(Arc<str>);
impl $name {
$(#[$new_doc])*
#[must_use]
pub fn new(s: impl Into<Arc<str>>) -> Self {
Self(s.into())
}
$(#[$as_str_doc])*
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl Default for $name {
$(#[$default_doc])*
fn default() -> Self {
Self(Arc::from(""))
}
}
impl fmt::Display for $name {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl AsRef<str> for $name {
fn as_ref(&self) -> &str {
&self.0
}
}
impl Borrow<str> for $name {
fn borrow(&self) -> &str {
&self.0
}
}
impl From<&str> for $name {
fn from(s: &str) -> Self {
Self(Arc::from(s))
}
}
impl From<String> for $name {
fn from(s: String) -> Self {
Self(Arc::from(s.as_str()))
}
}
impl FromStr for $name {
type Err = std::convert::Infallible;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(Self::from(s))
}
}
impl PartialEq<str> for $name {
fn eq(&self, other: &str) -> bool {
self.0.as_ref() == other
}
}
impl PartialEq<&str> for $name {
fn eq(&self, other: &&str) -> bool {
self.0.as_ref() == *other
}
}
impl PartialEq<String> for $name {
fn eq(&self, other: &String) -> bool {
self.0.as_ref() == other.as_str()
}
}
impl PartialEq<$name> for str {
fn eq(&self, other: &$name) -> bool {
self == other.0.as_ref()
}
}
impl PartialEq<$name> for String {
fn eq(&self, other: &$name) -> bool {
self.as_str() == other.0.as_ref()
}
}
};
}
arc_str_newtype!(
struct ToolName;
new_doc:
as_str_doc:
default_doc:
);
arc_str_newtype!(
struct ProviderName;
new_doc:
as_str_doc:
default_doc:
);
impl ProviderName {
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
#[must_use]
pub fn as_non_empty(&self) -> Option<&str> {
if self.0.is_empty() {
None
} else {
Some(&self.0)
}
}
}
arc_str_newtype!(
struct SkillName;
new_doc:
as_str_doc:
default_doc:
);
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(transparent)]
pub struct SessionId(String);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
pub enum SessionIdError {
#[error("session id must not be empty")]
Empty,
#[error("session id must not contain path separators, '..', or NUL bytes")]
UnsafeCharacters,
}
impl SessionId {
pub fn new(s: impl Into<String>) -> Self {
let s = s.into();
debug_assert!(!s.is_empty(), "SessionId must not be empty");
Self(s)
}
pub fn try_new(s: impl Into<String>) -> Result<Self, SessionIdError> {
let s = s.into();
if s.is_empty() {
return Err(SessionIdError::Empty);
}
if s.contains('/') || s.contains('\\') || s.contains("..") || s.contains('\0') {
return Err(SessionIdError::UnsafeCharacters);
}
Ok(Self(s))
}
#[must_use]
pub fn generate() -> Self {
Self(uuid::Uuid::new_v4().to_string())
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl Default for SessionId {
fn default() -> Self {
Self::generate()
}
}
impl fmt::Display for SessionId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl AsRef<str> for SessionId {
fn as_ref(&self) -> &str {
&self.0
}
}
impl std::ops::Deref for SessionId {
type Target = str;
fn deref(&self) -> &str {
&self.0
}
}
impl From<String> for SessionId {
fn from(s: String) -> Self {
Self::new(s)
}
}
impl From<&str> for SessionId {
fn from(s: &str) -> Self {
Self::new(s)
}
}
impl From<uuid::Uuid> for SessionId {
fn from(u: uuid::Uuid) -> Self {
Self(u.to_string())
}
}
impl FromStr for SessionId {
type Err = std::convert::Infallible;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(Self::new(s))
}
}
impl PartialEq<str> for SessionId {
fn eq(&self, other: &str) -> bool {
self.0 == other
}
}
impl PartialEq<&str> for SessionId {
fn eq(&self, other: &&str) -> bool {
self.0 == *other
}
}
impl PartialEq<String> for SessionId {
fn eq(&self, other: &String) -> bool {
self.0 == *other
}
}
impl PartialEq<SessionId> for str {
fn eq(&self, other: &SessionId) -> bool {
self == other.0
}
}
impl PartialEq<SessionId> for String {
fn eq(&self, other: &SessionId) -> bool {
*self == other.0
}
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct ToolDefinition {
pub name: ToolName,
pub description: String,
pub parameters: serde_json::Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_schema: Option<serde_json::Value>,
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StopHint {
MaxTokens,
MaxTurnRequests,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tool_name_construction_and_equality() {
let name = ToolName::new("shell");
assert_eq!(name.as_str(), "shell");
assert_eq!(name, "shell");
assert_eq!(name, "shell".to_owned());
assert_eq!(*"shell", name);
assert_eq!("shell".to_owned(), name);
}
#[test]
fn tool_name_default_is_empty() {
let name = ToolName::default();
assert_eq!(name.as_str(), "");
}
#[test]
fn tool_name_clone_is_cheap() {
let name = ToolName::new("web_scrape");
let name2 = name.clone();
assert_eq!(name, name2);
assert!(Arc::ptr_eq(&name.0, &name2.0));
}
#[test]
fn tool_name_from_impls() {
let from_str: ToolName = ToolName::from("bash");
let from_string: ToolName = ToolName::from("bash".to_owned());
let parsed: ToolName = "bash".parse().unwrap();
assert_eq!(from_str, from_string);
assert_eq!(from_str, parsed);
}
#[test]
fn tool_name_as_hashmap_key() {
use std::collections::HashMap;
let mut map: HashMap<ToolName, u32> = HashMap::new();
map.insert(ToolName::new("shell"), 1);
assert_eq!(map.get("shell"), Some(&1));
}
#[test]
fn tool_name_display() {
let name = ToolName::new("my_tool");
assert_eq!(format!("{name}"), "my_tool");
}
#[test]
fn tool_name_serde_transparent() {
let name = ToolName::new("shell");
let json = serde_json::to_string(&name).unwrap();
assert_eq!(json, r#""shell""#);
let back: ToolName = serde_json::from_str(&json).unwrap();
assert_eq!(back, name);
}
#[test]
fn session_id_new_roundtrip() {
let id = SessionId::new("test-session");
assert_eq!(id.as_str(), "test-session");
assert_eq!(id.to_string(), "test-session");
}
#[test]
fn session_id_generate_is_uuid() {
let id = SessionId::generate();
assert_eq!(id.as_str().len(), 36);
assert!(uuid::Uuid::parse_str(id.as_str()).is_ok());
}
#[test]
fn session_id_default_is_generated() {
let id = SessionId::default();
assert!(!id.as_str().is_empty());
assert_eq!(id.as_str().len(), 36);
}
#[test]
fn session_id_from_uuid() {
let u = uuid::Uuid::new_v4();
let id = SessionId::from(u);
assert_eq!(id.as_str(), u.to_string());
}
#[test]
fn session_id_deref_slicing() {
let id = SessionId::new("abcdefgh");
assert_eq!(&id[..4], "abcd");
}
#[test]
fn session_id_serde_transparent() {
let id = SessionId::new("sess-abc");
let json = serde_json::to_string(&id).unwrap();
assert_eq!(json, r#""sess-abc""#);
let back: SessionId = serde_json::from_str(&json).unwrap();
assert_eq!(back, id);
}
#[test]
fn session_id_from_str_parses() {
let id: SessionId = "my-session".parse().unwrap();
assert_eq!(id.as_str(), "my-session");
}
#[test]
fn session_id_try_new_accepts_valid_uuid() {
let id = SessionId::try_new(uuid::Uuid::new_v4().to_string()).unwrap();
assert_eq!(id.as_str().len(), 36);
}
#[test]
fn session_id_try_new_accepts_plain_string() {
let id = SessionId::try_new("sess-abc123").unwrap();
assert_eq!(id.as_str(), "sess-abc123");
}
#[test]
fn session_id_try_new_rejects_empty() {
assert_eq!(SessionId::try_new(""), Err(SessionIdError::Empty));
}
#[test]
fn session_id_try_new_rejects_path_traversal() {
assert_eq!(
SessionId::try_new("../../etc/passwd"),
Err(SessionIdError::UnsafeCharacters)
);
assert_eq!(
SessionId::try_new("foo/../bar"),
Err(SessionIdError::UnsafeCharacters)
);
}
#[test]
fn session_id_try_new_rejects_forward_slash() {
assert_eq!(
SessionId::try_new("foo/bar"),
Err(SessionIdError::UnsafeCharacters)
);
}
#[test]
fn session_id_try_new_rejects_backslash() {
assert_eq!(
SessionId::try_new("foo\\bar"),
Err(SessionIdError::UnsafeCharacters)
);
}
#[test]
fn session_id_try_new_rejects_nul_byte() {
assert_eq!(
SessionId::try_new("foo\0bar"),
Err(SessionIdError::UnsafeCharacters)
);
}
#[test]
fn provider_name_construction_and_equality() {
let name = ProviderName::new("fast");
assert_eq!(name.as_str(), "fast");
assert_eq!(name, "fast");
assert_eq!(name, "fast".to_owned());
assert_eq!(*"fast", name);
assert_eq!("fast".to_owned(), name);
}
#[test]
fn provider_name_clone_is_cheap() {
let name = ProviderName::new("quality");
let name2 = name.clone();
assert_eq!(name, name2);
assert!(Arc::ptr_eq(&name.0, &name2.0));
}
#[test]
fn provider_name_from_impls() {
let from_str: ProviderName = ProviderName::from("fast");
let from_string: ProviderName = ProviderName::from("fast".to_owned());
let parsed: ProviderName = "fast".parse().unwrap();
assert_eq!(from_str, from_string);
assert_eq!(from_str, parsed);
}
#[test]
fn provider_name_as_hashmap_key() {
use std::collections::HashMap;
let mut map: HashMap<ProviderName, u32> = HashMap::new();
map.insert(ProviderName::new("fast"), 1);
assert_eq!(map.get("fast"), Some(&1));
}
#[test]
fn provider_name_display() {
let name = ProviderName::new("ollama-local");
assert_eq!(format!("{name}"), "ollama-local");
}
#[test]
fn provider_name_serde_transparent() {
let name = ProviderName::new("quality");
let json = serde_json::to_string(&name).unwrap();
assert_eq!(json, r#""quality""#);
let back: ProviderName = serde_json::from_str(&json).unwrap();
assert_eq!(back, name);
}
#[test]
fn skill_name_construction_and_equality() {
let name = SkillName::new("rust-agents");
assert_eq!(name.as_str(), "rust-agents");
assert_eq!(name, "rust-agents");
assert_eq!(name, "rust-agents".to_owned());
assert_eq!(*"rust-agents", name);
assert_eq!("rust-agents".to_owned(), name);
}
#[test]
fn skill_name_default_is_empty() {
let name = SkillName::default();
assert_eq!(name.as_str(), "");
}
#[test]
fn skill_name_clone_is_cheap() {
let name = SkillName::new("readme-generator");
let name2 = name.clone();
assert_eq!(name, name2);
assert!(Arc::ptr_eq(&name.0, &name2.0));
}
#[test]
fn skill_name_from_impls() {
let from_str: SkillName = SkillName::from("rust-agents");
let from_string: SkillName = SkillName::from("rust-agents".to_owned());
let parsed: SkillName = "rust-agents".parse().unwrap();
assert_eq!(from_str, from_string);
assert_eq!(from_str, parsed);
}
#[test]
fn skill_name_as_hashmap_key() {
use std::collections::HashMap;
let mut map: HashMap<SkillName, u32> = HashMap::new();
map.insert(SkillName::new("rust-agents"), 1);
assert_eq!(map.get("rust-agents"), Some(&1));
}
#[test]
fn skill_name_display() {
let name = SkillName::new("readme-generator");
assert_eq!(format!("{name}"), "readme-generator");
}
#[test]
fn skill_name_serde_transparent() {
let name = SkillName::new("rust-agents");
let json = serde_json::to_string(&name).unwrap();
assert_eq!(json, r#""rust-agents""#);
let back: SkillName = serde_json::from_str(&json).unwrap();
assert_eq!(back, name);
}
}