claude - move api types to their own module, simplify message buffer computations
This commit is contained in:
@@ -0,0 +1,120 @@
|
|||||||
|
//! Anthropic Claude API types.
|
||||||
|
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
use serde_json::Value;
|
||||||
|
use signal_gateway_assistant::Tool;
|
||||||
|
|
||||||
|
/// Request body for the Claude Messages API.
|
||||||
|
#[derive(Serialize)]
|
||||||
|
pub(crate) struct MessagesRequest<'a> {
|
||||||
|
pub model: &'a str,
|
||||||
|
pub max_tokens: u32,
|
||||||
|
pub system: &'a [SystemContent],
|
||||||
|
pub messages: &'a [MessageContent],
|
||||||
|
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||||
|
pub tools: Vec<Tool>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A content block in the system prompt array.
|
||||||
|
#[derive(Clone, Default, Serialize)]
|
||||||
|
pub(crate) struct SystemContent {
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
content_type: &'static str,
|
||||||
|
text: String,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
cache_control: Option<CacheControl>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl SystemContent {
|
||||||
|
pub fn text(text: impl Into<String>) -> Self {
|
||||||
|
Self {
|
||||||
|
content_type: "text",
|
||||||
|
text: text.into(),
|
||||||
|
cache_control: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn set_cached(&mut self) {
|
||||||
|
self.cache_control = Some(CacheControl::ephemeral());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Cache control directive for prompt caching.
|
||||||
|
#[derive(Clone, Serialize)]
|
||||||
|
struct CacheControl {
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
cache_type: &'static str,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl CacheControl {
|
||||||
|
fn ephemeral() -> Self {
|
||||||
|
Self {
|
||||||
|
cache_type: "ephemeral",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A message in the conversation (can have multiple content blocks).
|
||||||
|
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||||
|
pub(crate) struct MessageContent {
|
||||||
|
pub role: Box<str>,
|
||||||
|
pub content: Vec<ContentBlock>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl MessageContent {
|
||||||
|
pub fn assistant(blocks: Vec<ContentBlock>) -> Self {
|
||||||
|
Self {
|
||||||
|
role: "assistant".into(),
|
||||||
|
content: blocks,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn tool_result(tool_use_id: String, content: String, is_error: bool) -> Self {
|
||||||
|
Self {
|
||||||
|
role: "user".into(),
|
||||||
|
content: vec![ContentBlock::ToolResult {
|
||||||
|
tool_use_id: tool_use_id.into(),
|
||||||
|
content: content.into(),
|
||||||
|
is_error: if is_error { Some(true) } else { None },
|
||||||
|
}],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A content block in the request/response.
|
||||||
|
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||||
|
#[serde(tag = "type", rename_all = "snake_case")]
|
||||||
|
pub(crate) enum ContentBlock {
|
||||||
|
Text {
|
||||||
|
text: Box<str>,
|
||||||
|
},
|
||||||
|
ToolUse {
|
||||||
|
id: Box<str>,
|
||||||
|
name: Box<str>,
|
||||||
|
input: Value,
|
||||||
|
},
|
||||||
|
ToolResult {
|
||||||
|
tool_use_id: Box<str>,
|
||||||
|
content: Box<str>,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
is_error: Option<bool>,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Response from the Claude Messages API.
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
pub(crate) struct MessagesResponse {
|
||||||
|
pub content: Vec<ContentBlock>,
|
||||||
|
pub stop_reason: Box<str>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Error response from the Claude API.
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
pub(crate) struct ErrorResponse {
|
||||||
|
pub error: ApiErrorDetail,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
pub(crate) struct ApiErrorDetail {
|
||||||
|
pub message: Box<str>,
|
||||||
|
}
|
||||||
@@ -1,12 +1,14 @@
|
|||||||
//! Claude API implementation of the Assistant trait.
|
//! Claude API implementation of the Assistant trait.
|
||||||
|
|
||||||
|
mod api;
|
||||||
mod message_buffer;
|
mod message_buffer;
|
||||||
|
|
||||||
|
use api::{
|
||||||
|
ContentBlock, ErrorResponse, MessageContent, MessagesRequest, MessagesResponse, SystemContent,
|
||||||
|
};
|
||||||
use conf::Conf;
|
use conf::Conf;
|
||||||
use message_buffer::MessageBuffer;
|
use message_buffer::MessageBuffer;
|
||||||
use serde::{Deserialize, Serialize};
|
use signal_gateway_assistant::{Assistant, AssistantResponse, ChatMessage, ToolExecutor};
|
||||||
use serde_json::Value;
|
|
||||||
use signal_gateway_assistant::{Assistant, AssistantResponse, ChatMessage, Tool, ToolExecutor};
|
|
||||||
use std::{path::PathBuf, sync::Weak, time::Duration};
|
use std::{path::PathBuf, sync::Weak, time::Duration};
|
||||||
use tokio_util::sync::CancellationToken;
|
use tokio_util::sync::CancellationToken;
|
||||||
use tracing::{error, info, warn};
|
use tracing::{error, info, warn};
|
||||||
@@ -485,119 +487,3 @@ fn message_to_content(msg: ChatMessage) -> MessageContent {
|
|||||||
content: vec![ContentBlock::Text { text: text.into() }],
|
content: vec![ContentBlock::Text { text: text.into() }],
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// ---- API Types ----
|
|
||||||
|
|
||||||
/// Request body for the Claude Messages API.
|
|
||||||
#[derive(Serialize)]
|
|
||||||
struct MessagesRequest<'a> {
|
|
||||||
model: &'a str,
|
|
||||||
max_tokens: u32,
|
|
||||||
system: &'a [SystemContent],
|
|
||||||
messages: &'a [MessageContent],
|
|
||||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
|
||||||
tools: Vec<Tool>,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A content block in the system prompt array.
|
|
||||||
#[derive(Clone, Default, Serialize)]
|
|
||||||
struct SystemContent {
|
|
||||||
#[serde(rename = "type")]
|
|
||||||
content_type: &'static str,
|
|
||||||
text: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
cache_control: Option<CacheControl>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl SystemContent {
|
|
||||||
fn text(text: impl Into<String>) -> Self {
|
|
||||||
Self {
|
|
||||||
content_type: "text",
|
|
||||||
text: text.into(),
|
|
||||||
cache_control: None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn set_cached(&mut self) {
|
|
||||||
self.cache_control = Some(CacheControl::ephemeral());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Cache control directive for prompt caching.
|
|
||||||
#[derive(Clone, Serialize)]
|
|
||||||
struct CacheControl {
|
|
||||||
#[serde(rename = "type")]
|
|
||||||
cache_type: &'static str,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl CacheControl {
|
|
||||||
fn ephemeral() -> Self {
|
|
||||||
Self {
|
|
||||||
cache_type: "ephemeral",
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A message in the conversation (can have multiple content blocks).
|
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
|
||||||
struct MessageContent {
|
|
||||||
role: Box<str>,
|
|
||||||
content: Vec<ContentBlock>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl MessageContent {
|
|
||||||
fn assistant(blocks: Vec<ContentBlock>) -> Self {
|
|
||||||
Self {
|
|
||||||
role: "assistant".into(),
|
|
||||||
content: blocks,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn tool_result(tool_use_id: String, content: String, is_error: bool) -> Self {
|
|
||||||
Self {
|
|
||||||
role: "user".into(),
|
|
||||||
content: vec![ContentBlock::ToolResult {
|
|
||||||
tool_use_id: tool_use_id.into(),
|
|
||||||
content: content.into(),
|
|
||||||
is_error: if is_error { Some(true) } else { None },
|
|
||||||
}],
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A content block in the request/response.
|
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
|
||||||
#[serde(tag = "type", rename_all = "snake_case")]
|
|
||||||
enum ContentBlock {
|
|
||||||
Text {
|
|
||||||
text: Box<str>,
|
|
||||||
},
|
|
||||||
ToolUse {
|
|
||||||
id: Box<str>,
|
|
||||||
name: Box<str>,
|
|
||||||
input: Value,
|
|
||||||
},
|
|
||||||
ToolResult {
|
|
||||||
tool_use_id: Box<str>,
|
|
||||||
content: Box<str>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
is_error: Option<bool>,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Response from the Claude Messages API.
|
|
||||||
#[derive(Debug, Deserialize)]
|
|
||||||
struct MessagesResponse {
|
|
||||||
content: Vec<ContentBlock>,
|
|
||||||
stop_reason: Box<str>,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Error response from the Claude API.
|
|
||||||
#[derive(Deserialize)]
|
|
||||||
struct ErrorResponse {
|
|
||||||
error: ApiErrorDetail,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Deserialize)]
|
|
||||||
struct ApiErrorDetail {
|
|
||||||
message: Box<str>,
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,15 +1,13 @@
|
|||||||
//! Message buffer with cached character count.
|
//! Message buffer with cached character count.
|
||||||
|
|
||||||
|
use crate::api::{ContentBlock, MessageContent};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use std::collections::VecDeque;
|
use std::collections::VecDeque;
|
||||||
use std::fmt;
|
|
||||||
|
|
||||||
use crate::{ContentBlock, MessageContent};
|
|
||||||
|
|
||||||
/// A buffer of messages with cached total character count.
|
/// A buffer of messages with cached total character count.
|
||||||
///
|
#[derive(Clone, Debug, Default)]
|
||||||
/// Uses `VecDeque` for efficient front removal during compaction.
|
|
||||||
pub struct MessageBuffer {
|
pub struct MessageBuffer {
|
||||||
|
/// `VecDeque` to easily remove oldest messages if needed
|
||||||
messages: VecDeque<MessageContent>,
|
messages: VecDeque<MessageContent>,
|
||||||
/// Cached total character count of all messages.
|
/// Cached total character count of all messages.
|
||||||
total_chars: usize,
|
total_chars: usize,
|
||||||
@@ -19,12 +17,12 @@ impl MessageBuffer {
|
|||||||
/// Create a new empty message buffer.
|
/// Create a new empty message buffer.
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
messages: VecDeque::new(),
|
messages: VecDeque::with_capacity(256),
|
||||||
total_chars: 0,
|
total_chars: 0,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Push a message to the back of the buffer.
|
/// Push a message to the back.
|
||||||
pub fn push(&mut self, msg: MessageContent) {
|
pub fn push(&mut self, msg: MessageContent) {
|
||||||
self.total_chars += message_chars(&msg);
|
self.total_chars += message_chars(&msg);
|
||||||
self.messages.push_back(msg);
|
self.messages.push_back(msg);
|
||||||
@@ -58,7 +56,7 @@ impl MessageBuffer {
|
|||||||
self.total_chars
|
self.total_chars
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Returns a reference to the last message, if any.
|
/// Get the last message, if any.
|
||||||
pub fn last(&self) -> Option<&MessageContent> {
|
pub fn last(&self) -> Option<&MessageContent> {
|
||||||
self.messages.back()
|
self.messages.back()
|
||||||
}
|
}
|
||||||
@@ -71,22 +69,6 @@ impl MessageBuffer {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Default for MessageBuffer {
|
|
||||||
fn default() -> Self {
|
|
||||||
Self::new()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl fmt::Debug for MessageBuffer {
|
|
||||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
|
||||||
f.debug_struct("MessageBuffer")
|
|
||||||
.field("len", &self.messages.len())
|
|
||||||
.field("total_chars", &self.total_chars)
|
|
||||||
.field("messages", &self.messages)
|
|
||||||
.finish()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Calculate the character count for a single message.
|
/// Calculate the character count for a single message.
|
||||||
fn message_chars(msg: &MessageContent) -> usize {
|
fn message_chars(msg: &MessageContent) -> usize {
|
||||||
msg.content
|
msg.content
|
||||||
@@ -99,23 +81,25 @@ fn message_chars(msg: &MessageContent) -> usize {
|
|||||||
.sum()
|
.sum()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Estimate the serialized JSON size of a value.
|
// Estimate the serialized JSON size of a value.
|
||||||
|
// https://github.com/serde-rs/json/issues/784#issuecomment-877688512
|
||||||
fn estimate_json_size(value: &Value) -> usize {
|
fn estimate_json_size(value: &Value) -> usize {
|
||||||
match value {
|
use serde::Serialize;
|
||||||
Value::Null => 4,
|
use std::io::{Result, Write};
|
||||||
Value::Bool(true) => 4,
|
|
||||||
Value::Bool(false) => 5,
|
struct ByteCount(usize);
|
||||||
Value::Number(n) => n.to_string().len(),
|
|
||||||
Value::String(s) => s.len() + 2,
|
impl Write for ByteCount {
|
||||||
Value::Array(arr) => {
|
fn write(&mut self, buf: &[u8]) -> Result<usize> {
|
||||||
2 + arr.iter().map(estimate_json_size).sum::<usize>() + arr.len().saturating_sub(1)
|
self.0 += buf.len();
|
||||||
|
Ok(buf.len())
|
||||||
}
|
}
|
||||||
Value::Object(obj) => {
|
fn flush(&mut self) -> Result<()> {
|
||||||
2 + obj
|
Ok(())
|
||||||
.iter()
|
|
||||||
.map(|(k, v)| k.len() + 3 + estimate_json_size(v))
|
|
||||||
.sum::<usize>()
|
|
||||||
+ obj.len().saturating_sub(1)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let mut ser = serde_json::Serializer::new(ByteCount(0));
|
||||||
|
value.serialize(&mut ser).unwrap();
|
||||||
|
ser.into_inner().0
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -25,10 +25,11 @@ pub enum Input {
|
|||||||
/// Background worker that processes assistant requests serially.
|
/// Background worker that processes assistant requests serially.
|
||||||
pub struct AssistantWorker {
|
pub struct AssistantWorker {
|
||||||
assistant: Box<dyn Assistant>,
|
assistant: Box<dyn Assistant>,
|
||||||
|
// Serialize input
|
||||||
input_rx: mpsc::Receiver<Input>,
|
input_rx: mpsc::Receiver<Input>,
|
||||||
/// Used when the user wants to cancel the current requests, but not shutdown the service
|
// Used when the user wants to cancel the current requests, but not shutdown the service
|
||||||
stop_rx: mpsc::Receiver<()>,
|
stop_rx: mpsc::Receiver<()>,
|
||||||
/// Used when the user wants to shut down the service
|
// Used when the user wants to shut down the service
|
||||||
cancellation_token: CancellationToken,
|
cancellation_token: CancellationToken,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user