use std::fmt::Display;
use arxiv;
use arxiv::{Arxiv, ArxivQuery};
use sapiens::tools::{Describe, ProtoToolDescribe, ProtoToolInvoke, ToolDescription, ToolUseError};
use sapiens_derive::{Describe, ProtoToolDescribe, ProtoToolInvoke};
use serde::{Deserialize, Serialize};
#[derive(Debug, ProtoToolInvoke, ProtoToolDescribe)]
#[tool(name = "Arxiv", input = "ArxivToolInput", output = "ArxivToolOutput")]
pub struct ArxivTool {}
#[derive(Debug, Deserialize, Serialize, Default, Clone)]
pub enum SortOrder {
#[serde(rename = "ascending")]
Ascending,
#[serde(rename = "descending")]
#[default]
Descending,
}
impl Display for SortOrder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SortOrder::Ascending => write!(f, "ascending"),
SortOrder::Descending => write!(f, "descending"),
}
}
}
#[derive(Debug, Deserialize, Serialize, Default, Clone)]
pub enum SortBy {
#[serde(rename = "relevance")]
#[default]
Relevance,
#[serde(rename = "lastUpdatedDate")]
LastUpdatedDate,
#[serde(rename = "submittedDate")]
SubmittedDate,
}
impl Display for SortBy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SortBy::Relevance => write!(f, "relevance"),
SortBy::LastUpdatedDate => write!(f, "lastUpdatedDate"),
SortBy::SubmittedDate => write!(f, "submittedDate"),
}
}
}
#[derive(Debug, Deserialize, Serialize, Describe)]
pub struct ArxivToolInput {
pub search_query: String,
pub id_list: Option<String>,
pub start: Option<i32>,
pub max_results: Option<i32>,
pub sort_by: Option<SortBy>,
pub sort_order: Option<SortOrder>,
pub show_pdf_url: Option<bool>,
pub show_authors: Option<bool>,
pub show_comments: Option<bool>,
pub show_summary: Option<bool>,
}
impl From<&ArxivToolInput> for ArxivQuery {
fn from(input: &ArxivToolInput) -> Self {
ArxivQuery {
base_url: "https://export.arxiv.org/api/query?".to_string(),
search_query: input.search_query.clone(),
id_list: input.id_list.clone().unwrap_or_default(),
start: input.start,
max_results: input.max_results,
sort_by: input.sort_by.clone().unwrap_or_default().to_string(),
sort_order: input.sort_order.clone().unwrap_or_default().to_string(),
}
}
}
#[derive(Debug, Deserialize, Serialize, Describe)]
pub struct ArxivToolOutput {
result: Vec<ArxivResult>,
}
#[derive(Debug, Deserialize, Serialize, Describe)]
pub struct ArxivResult {
pub id: String,
pub updated: String,
pub published: String,
pub title: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub summary: Option<String>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub authors: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pdf_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub comment: Option<String>,
}
impl From<Arxiv> for ArxivResult {
fn from(arxiv: Arxiv) -> Self {
ArxivResult {
id: arxiv.id,
updated: arxiv.updated,
published: arxiv.published,
title: arxiv.title,
summary: Some(arxiv.summary),
authors: arxiv.authors,
pdf_url: Some(arxiv.pdf_url),
comment: arxiv.comment,
}
}
}
impl ArxivTool {
pub async fn new() -> ArxivTool {
ArxivTool {}
}
#[tracing::instrument(skip(self))]
async fn invoke_typed(&self, input: &ArxivToolInput) -> Result<ArxivToolOutput, ToolUseError> {
let query = ArxivQuery::from(input);
if query.max_results.unwrap_or(0) > 100 {
return Err(ToolUseError::InvocationFailed(
"max_results cannot be greater than 100".to_string(),
));
}
let result = arxiv::fetch_arxivs(query)
.await
.map_err(|e| ToolUseError::InvocationFailed(e.to_string()))?;
let vec = result
.into_iter()
.map(|x| x.into())
.map(|mut x: ArxivResult| {
if !(input.show_pdf_url.unwrap_or(false)) {
x.pdf_url = None;
}
if !(input.show_comments.unwrap_or(false)) {
x.comment = None;
}
if !(input.show_summary.unwrap_or(false)) {
x.summary = None;
}
if !(input.show_authors.unwrap_or(false)) {
x.authors = vec![];
}
x
})
.collect();
Ok(ArxivToolOutput { result: vec })
}
}
#[cfg(test)]
mod tests {
use indoc::indoc;
use insta::assert_yaml_snapshot;
use super::*;
#[tokio::test]
async fn test_arxiv() {
let tool = ArxivTool::new().await;
let input = ArxivToolInput {
search_query: "cat:cs.AI".to_string(),
id_list: None,
start: None,
max_results: None,
sort_by: Some(SortBy::Relevance),
sort_order: Some(SortOrder::Ascending),
show_authors: None,
show_comments: None,
show_summary: Some(false),
show_pdf_url: Some(false),
};
let output = tool.invoke_typed(&input).await.unwrap();
assert!(!output.result.is_empty())
}
#[tokio::test]
async fn test_arxiv_from_yaml() {
let tool = ArxivTool::new().await;
let input = indoc! {"
search_query: cat:cs.AI
show_authors: true
"};
let input: ArxivToolInput = serde_yaml::from_str(input).unwrap();
assert_yaml_snapshot!(input);
let output = tool.invoke_typed(&input).await.unwrap();
assert!(!output.result.is_empty());
assert!(!output.result[0].authors.is_empty());
}
#[tokio::test]
async fn test_arxiv_from_yaml_2() {
let tool = ArxivTool::new().await;
let input = indoc! {"
search_query: cat:cs.DB
max_results: 4
show_authors: true
show_pdf_url: true
"};
let input: ArxivToolInput = serde_yaml::from_str(input).unwrap();
assert_yaml_snapshot!(input);
let output = tool.invoke_typed(&input).await.unwrap();
assert_eq!(output.result.len(), 4);
assert!(!output.result[0].authors.is_empty());
let yaml = serde_yaml::to_value(&output).unwrap();
assert_yaml_snapshot!(yaml);
}
}