use serde::Serialize;
use serde_json::{Map, Value};
use crate::completion::provider_options::reply_field;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct TogetherExt;
impl ProviderExtension for TogetherExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = TogetherOptions;
type Extras = TogetherExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct TogetherOptions {
#[serde(rename = "*")]
pub shared: TogetherShared,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct TogetherShared {
#[serde(skip_serializing_if = "Map::is_empty")]
pub chat_template_kwargs: Map<String, Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub repetition_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub safety_model: Option<String>,
}
impl TogetherOptions {
pub fn new() -> Self {
Self::default()
}
pub fn chat_template_kwarg(mut self, key: impl Into<String>, value: Value) -> Self {
self.shared.chat_template_kwargs.insert(key.into(), value);
self
}
pub fn top_k(mut self, k: i32) -> Self {
self.shared.top_k = Some(k);
self
}
pub fn min_p(mut self, p: f64) -> Self {
self.shared.min_p = Some(p);
self
}
pub fn repetition_penalty(mut self, penalty: f64) -> Self {
self.shared.repetition_penalty = Some(penalty);
self
}
pub fn safety_model(mut self, model: impl Into<String>) -> Self {
self.shared.safety_model = Some(model.into());
self
}
}
impl ExtensionOptions for TogetherOptions {
type Ext = TogetherExt;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct TogetherExtras {
pub warnings: Option<Vec<Value>>,
pub reasoning: Option<String>,
}
impl ReplyExtras for TogetherExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
warnings: reply_field(raw, "/warnings")?,
reasoning: reply_field(raw, "/choices/0/message/reasoning")?,
})
}
}
#[cfg(test)]
mod tests;