rig_core/providers/together/
extension.rs1use 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#[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#[non_exhaustive]
32#[derive(Clone, Debug, Default, PartialEq, Serialize)]
33pub struct TogetherOptions {
34 #[serde(rename = "*")]
36 pub shared: TogetherShared,
37}
38
39#[non_exhaustive]
41#[derive(Clone, Debug, Default, PartialEq, Serialize)]
42pub struct TogetherShared {
43 #[serde(skip_serializing_if = "Map::is_empty")]
45 pub chat_template_kwargs: Map<String, Value>,
46 #[serde(skip_serializing_if = "Option::is_none")]
48 pub top_k: Option<i32>,
49 #[serde(skip_serializing_if = "Option::is_none")]
51 pub min_p: Option<f64>,
52 #[serde(skip_serializing_if = "Option::is_none")]
54 pub repetition_penalty: Option<f64>,
55 #[serde(skip_serializing_if = "Option::is_none")]
57 pub safety_model: Option<String>,
58}
59
60impl TogetherOptions {
61 pub fn new() -> Self {
63 Self::default()
64 }
65
66 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 pub fn top_k(mut self, k: i32) -> Self {
74 self.shared.top_k = Some(k);
75 self
76 }
77
78 pub fn min_p(mut self, p: f64) -> Self {
80 self.shared.min_p = Some(p);
81 self
82 }
83
84 pub fn repetition_penalty(mut self, penalty: f64) -> Self {
86 self.shared.repetition_penalty = Some(penalty);
87 self
88 }
89
90 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#[non_exhaustive]
103#[derive(Clone, Debug, Default, PartialEq)]
104pub struct TogetherExtras {
105 pub warnings: Option<Vec<Value>>,
107 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;