Skip to main content

rig_core/providers/together/
extension.rs

1//! Together AI's typed request options and reply extras
2//! (<https://docs.together.ai/reference/chat-completions-1>).
3//!
4//! ```
5//! use rig_core::completion::CompletionRequest;
6//! use rig_core::providers::together::extension::{TogetherOptions};
7//!
8//! let options = TogetherOptions::new().top_k(40).repetition_penalty(1.1);
9//! let request = CompletionRequest::new("hi").provider_option(options);
10//! # let _ = request;
11//! ```
12
13use serde::Serialize;
14use serde_json::{Map, Value};
15
16use crate::completion::provider_options::reply_field;
17use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
18use crate::message::Api;
19
20/// Together AI's extension marker.
21#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
22pub struct TogetherExt;
23
24impl ProviderExtension for TogetherExt {
25    const PROVIDER: &'static str = super::PROVIDER_NAME;
26    type Options = TogetherOptions;
27    type Extras = TogetherExtras;
28}
29
30/// Together AI's request options.
31#[non_exhaustive]
32#[derive(Clone, Debug, Default, PartialEq, Serialize)]
33pub struct TogetherOptions {
34    /// The fields every route takes.
35    #[serde(rename = "*")]
36    pub shared: TogetherShared,
37}
38
39/// The fields Together AI takes.
40#[non_exhaustive]
41#[derive(Clone, Debug, Default, PartialEq, Serialize)]
42pub struct TogetherShared {
43    /// Arguments to the model's chat template, such as `enable_thinking`.
44    #[serde(skip_serializing_if = "Map::is_empty")]
45    pub chat_template_kwargs: Map<String, Value>,
46    /// Sample from the `k` most likely tokens.
47    #[serde(skip_serializing_if = "Option::is_none")]
48    pub top_k: Option<i32>,
49    /// The minimum probability of a token, relative to the most likely one.
50    #[serde(skip_serializing_if = "Option::is_none")]
51    pub min_p: Option<f64>,
52    /// Penalizes tokens already in the prompt and the output.
53    #[serde(skip_serializing_if = "Option::is_none")]
54    pub repetition_penalty: Option<f64>,
55    /// The moderation model that screens the request.
56    #[serde(skip_serializing_if = "Option::is_none")]
57    pub safety_model: Option<String>,
58}
59
60impl TogetherOptions {
61    /// No option set.
62    pub fn new() -> Self {
63        Self::default()
64    }
65
66    /// Pass `key`, `value` to the model's chat template.
67    pub fn chat_template_kwarg(mut self, key: impl Into<String>, value: Value) -> Self {
68        self.shared.chat_template_kwargs.insert(key.into(), value);
69        self
70    }
71
72    /// Sample from the `k` most likely tokens.
73    pub fn top_k(mut self, k: i32) -> Self {
74        self.shared.top_k = Some(k);
75        self
76    }
77
78    /// Set the minimum relative token probability.
79    pub fn min_p(mut self, p: f64) -> Self {
80        self.shared.min_p = Some(p);
81        self
82    }
83
84    /// Set the repetition penalty.
85    pub fn repetition_penalty(mut self, penalty: f64) -> Self {
86        self.shared.repetition_penalty = Some(penalty);
87        self
88    }
89
90    /// Screen the request with moderation model `model`.
91    pub fn safety_model(mut self, model: impl Into<String>) -> Self {
92        self.shared.safety_model = Some(model.into());
93        self
94    }
95}
96
97impl ExtensionOptions for TogetherOptions {
98    type Ext = TogetherExt;
99}
100
101/// Together AI's reply fields. Each is `None` when the reply lacks it.
102#[non_exhaustive]
103#[derive(Clone, Debug, Default, PartialEq)]
104pub struct TogetherExtras {
105    /// Warnings about the request.
106    pub warnings: Option<Vec<Value>>,
107    /// The first choice's reasoning.
108    pub reasoning: Option<String>,
109}
110
111impl ReplyExtras for TogetherExtras {
112    fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
113        Ok(Self {
114            warnings: reply_field(raw, "/warnings")?,
115            reasoning: reply_field(raw, "/choices/0/message/reasoning")?,
116        })
117    }
118}
119
120#[cfg(test)]
121mod tests;