Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,10 @@ path = "examples/openrouter.rs"
name = "structured-generation"
path = "examples/structured_generation.rs"

[[example]]
name = "text-generation"
path = "examples/text_generation.rs"

[[example]]
name = "tracing"
path = "examples/tracing.rs"
48 changes: 47 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
## Quick Start

```rust
use rsai::{llm, Message, ChatRole, ApiKey, Provider, completion_schema};
use rsai::{llm, Message, ChatRole, ApiKey, Provider, TextResponse, completion_schema};

#[completion_schema]
struct Analysis {
Expand All @@ -20,6 +20,24 @@ let analysis = llm::with(Provider::OpenAI)
}])
.complete::<Analysis>()
.await?;

let reply = llm::with(Provider::OpenAI)
.api_key(ApiKey::Default)?
.model("gpt-4o-mini")
.messages(vec![
Message {
role: ChatRole::System,
content: "You are friendly and concise.".to_string(),
},
Message {
role: ChatRole::User,
content: "Share a fun fact about Rust.".to_string(),
},
])
.complete::<TextResponse>()
.await?;

println!("{}", reply.text);
```

## Structured Generation
Expand Down Expand Up @@ -113,6 +131,34 @@ struct CustomType {

> **Note**: The library automatically handles provider-specific requirements. For example, OpenAI requires root schemas to be objects, so non-object types like enums are transparently wrapped and unwrapped.

## Text Generation

To obtain plain-text output without defining a schema, target `TextResponse`:

```rust
use rsai::{llm, ApiKey, Provider, Message, ChatRole, TextResponse};

let fact = llm::with(Provider::OpenAI)
.api_key(ApiKey::Default)?
.model("gpt-4o-mini")
.messages(vec![
Message {
role: ChatRole::System,
content: "You explain things clearly but briefly.".to_string(),
},
Message {
role: ChatRole::User,
content: "What makes Rust's borrow checker special?".to_string(),
},
])
.complete::<TextResponse>()
.await?;

println!("{}", fact.text);
```

See `examples/text_generation.rs` for a runnable version.

## Known Issues

- ..
Expand Down
27 changes: 27 additions & 0 deletions examples/text_generation.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
use dotenv::dotenv;
use rsai::{ApiKey, ChatRole, Message, Provider, TextResponse, llm};

#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
dotenv().ok();

let response = llm::with(Provider::OpenAI)
.api_key(ApiKey::Default)?
.model("gpt-4o-mini")
.messages(vec![
Message {
role: ChatRole::System,
content: "You are a concise, upbeat assistant.".to_string(),
},
Message {
role: ChatRole::User,
content: "Share a fun fact about Rust programming.".to_string(),
},
])
.complete::<TextResponse>()
.await?;

println!("Assistant:\n{}", response.text);

Ok(())
}
6 changes: 3 additions & 3 deletions src/core.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,11 @@ pub use builder::{ApiKey, LlmBuilder, llm};

pub use error::LlmError;
pub use tool_guard::{ToolCallingConfig, ToolCallingGuard};
pub use traits::{LlmProvider, ToolFunction};
pub use traits::{CompletionTarget, LlmProvider, ToolFunction};

pub use types::StructuredRequest;
pub use types::{
BoxFuture, ChatRole, ConversationMessage, GenerationConfig, LanguageModelUsage, Message,
ResponseMetadata, StructuredResponse, Tool, ToolCall, ToolCallResult, ToolChoice, ToolConfig,
ToolRegistry, ToolSet,
ResponseMetadata, StructuredResponse, TextResponse, Tool, ToolCall, ToolCallResult, ToolChoice,
ToolConfig, ToolRegistry, ToolSet,
};
74 changes: 48 additions & 26 deletions src/core/builder.rs
Original file line number Diff line number Diff line change
@@ -1,7 +1,5 @@
use std::{env, marker::PhantomData};

use serde::de::Deserialize;

use tracing::{debug, instrument};

use crate::{
Expand All @@ -13,8 +11,8 @@ use super::{
error::LlmError,
traits::LlmProvider,
types::{
ConversationMessage, GenerationConfig, Message, StructuredRequest, StructuredResponse,
ToolChoice, ToolConfig, ToolRegistry,
ConversationMessage, GenerationConfig, Message, StructuredRequest, ToolChoice, ToolConfig,
ToolRegistry,
},
};

Expand Down Expand Up @@ -206,14 +204,14 @@ impl<State: private::Completable> LlmBuilder<State> {
self
}

/// Execute the LLM request and return structured output of type T.
/// The type T must implement Deserialize and JsonSchema for structured output generation as well as
/// be annotated with `#[schemars(deny_unknown_fields)]`.
/// Use the `completion_schema` attribute macro to easily define structured output types.
/// Execute the LLM request and return an output defined by `T`.
///
/// The target type `T` must implement [`CompletionTarget`]. Structured schemas can be created
/// with the `#[completion_schema]` macro, while plain text responses can use [`TextResponse`].
///
/// # Example
/// ```no_run
/// # use rsai::{completion_schema, llm, Message, ChatRole, ApiKey, Provider};
/// # use rsai::{completion_schema, llm, Message, ChatRole, ApiKey, Provider, TextResponse};
/// # #[tokio::main]
/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
/// #[completion_schema]
Expand All @@ -231,11 +229,21 @@ impl<State: private::Completable> LlmBuilder<State> {
/// }])
/// .complete::<Analysis>()
/// .await?;
///
/// let text = llm::with(Provider::OpenAI)
/// .api_key(ApiKey::Default)?
/// .model("gpt-4o-mini")
/// .messages(vec![Message {
/// role: ChatRole::User,
/// content: "Say hello".to_string(),
/// }])
/// .complete::<TextResponse>()
/// .await?;
/// # Ok(())
/// # }
/// ```
#[instrument(
name = "generate_structured",
name = "generate_completion",
skip(self),
fields(
model = ?self.fields.model,
Expand All @@ -244,23 +252,33 @@ impl<State: private::Completable> LlmBuilder<State> {
),
err
)]
pub async fn complete<T>(self) -> Result<StructuredResponse<T>, LlmError>
pub async fn complete<T>(self) -> Result<T::Output, LlmError>
where
T: for<'a> Deserialize<'a> + Send + schemars::JsonSchema,
T: super::traits::CompletionTarget + Send,
{
debug!("Starting structured generation request");
debug!("Starting generation request");
let (messages, provider, model) = self.fields.validate()?;
let model_string = model.to_string();
let messages = messages.to_vec();
let format = T::format()?;

if !T::supports_tools() && self.fields.tool_registry.is_some() {
return Err(LlmError::Builder(
"Tools are only supported with structured completion targets".to_string(),
));
}

// Deferred error handling for tool registry errors in case of a poisoned lock.
let tool_schemas = self
.fields
.tool_registry
.as_ref()
.map(|registry| registry.get_schemas())
.transpose()?
.map(|tools| tools.into_boxed_slice());
let tool_schemas = if T::supports_tools() {
self.fields
.tool_registry
.as_ref()
.map(|registry| registry.get_schemas())
.transpose()?
.map(|tools| tools.into_boxed_slice())
} else {
None
};

match provider {
Provider::OpenAI => {
Expand All @@ -272,8 +290,8 @@ impl<State: private::Completable> LlmBuilder<State> {
let req = StructuredRequest {
model: model_string,
messages: conversation_messages,
tool_config: Some(ToolConfig {
tools: tool_schemas.clone(),
tool_config: tool_schemas.map(|tools| ToolConfig {
tools: Some(tools),
tool_choice: self.fields.tool_choice.clone(),
parallel_tool_calls: self.fields.parallel_tool_calls,
}),
Expand All @@ -285,7 +303,11 @@ impl<State: private::Completable> LlmBuilder<State> {
};
let client = openai::create_openai_client_from_builder(&self)?;
client
.generate_structured(req, self.fields.tool_registry.as_ref())
.generate_completion::<T>(
req,
format.clone(),
self.fields.tool_registry.as_ref(),
)
.await
}
Provider::OpenRouter => {
Expand All @@ -297,8 +319,8 @@ impl<State: private::Completable> LlmBuilder<State> {
let req = StructuredRequest {
model: model_string,
messages: conversation_messages,
tool_config: Some(ToolConfig {
tools: tool_schemas.clone(),
tool_config: tool_schemas.map(|tools| ToolConfig {
tools: Some(tools),
tool_choice: self.fields.tool_choice.clone(),
parallel_tool_calls: self.fields.parallel_tool_calls,
}),
Expand All @@ -310,7 +332,7 @@ impl<State: private::Completable> LlmBuilder<State> {
};
let client = openrouter::create_openrouter_client_from_builder(&self)?;
client
.generate_structured(req, self.fields.tool_registry.as_ref())
.generate_completion::<T>(req, format, self.fields.tool_registry.as_ref())
.await
}
}
Expand Down
26 changes: 22 additions & 4 deletions src/core/traits.rs
Original file line number Diff line number Diff line change
@@ -1,19 +1,25 @@
use async_trait::async_trait;

use crate::{
Provider,
responses::{request::Format, response::Response},
};

use super::{
error::LlmError,
types::{BoxFuture, StructuredRequest, StructuredResponse, Tool, ToolRegistry},
types::{BoxFuture, StructuredRequest, Tool, ToolRegistry},
};

#[async_trait]
pub trait LlmProvider {
async fn generate_structured<T>(
async fn generate_completion<T>(
&self,
request: StructuredRequest,
format: Format,
tool_registry: Option<&ToolRegistry>,
) -> Result<StructuredResponse<T>, LlmError>
) -> Result<T::Output, LlmError>
where
T: serde::de::DeserializeOwned + Send + schemars::JsonSchema;
T: CompletionTarget + Send;
}

pub trait ToolFunction: Send + Sync {
Expand All @@ -23,3 +29,15 @@ pub trait ToolFunction: Send + Sync {
params: serde_json::Value,
) -> BoxFuture<'a, Result<serde_json::Value, LlmError>>;
}

pub trait CompletionTarget: Sized + Send {
type Output;

fn format() -> Result<Format, LlmError>;

fn parse_response(res: Response, provider: Provider) -> Result<Self::Output, LlmError>;

fn supports_tools() -> bool {
true
}
}
43 changes: 42 additions & 1 deletion src/core/types.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
use crate::core::{LlmError, traits::ToolFunction};
use crate::core::{LlmError, traits::CompletionTarget, traits::ToolFunction};
use crate::provider::Provider;
use crate::responses::{self, request::Format, response::Response};
use schemars::JsonSchema;
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::collections::HashMap;
use std::future::Future;
Expand Down Expand Up @@ -97,6 +100,13 @@ pub struct StructuredResponse<T> {
pub metadata: ResponseMetadata,
}

#[derive(Debug, Clone, PartialEq)]
pub struct TextResponse {
pub text: String,
pub usage: LanguageModelUsage,
pub metadata: ResponseMetadata,
}

#[derive(Debug, Clone, PartialEq)]
pub struct LanguageModelUsage {
pub prompt_tokens: i32,
Expand Down Expand Up @@ -255,6 +265,37 @@ impl ToolSet {
}
}

impl<T> CompletionTarget for T
where
T: DeserializeOwned + JsonSchema + Send,
{
type Output = StructuredResponse<T>;

fn format() -> Result<Format, LlmError> {
responses::create_format_for_type::<T>()
}

fn parse_response(res: Response, provider: Provider) -> Result<Self::Output, LlmError> {
responses::create_core_structured_response(res, provider)
}
}

impl CompletionTarget for TextResponse {
type Output = TextResponse;

fn format() -> Result<Format, LlmError> {
Ok(responses::create_text_format())
}

fn parse_response(res: Response, provider: Provider) -> Result<Self::Output, LlmError> {
responses::create_core_text_response(res, provider)
}

fn supports_tools() -> bool {
false
}
}

#[cfg(test)]
mod tests {
use async_trait::async_trait;
Expand Down
Loading