gemini_genai_rs/tunings/
mod.rs

1//! Tunings API — create, list, get, cancel tuning jobs.
2//!
3//! Feature-gated behind `tunings`.
4
5use serde::{Deserialize, Serialize};
6
7use crate::client::Client;
8use crate::client::http::HttpError;
9use crate::transport::auth::ServiceEndpoint;
10
11/// State of a tuning job.
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
13#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
14pub enum TuningJobState {
15    /// State not set by the server.
16    StateUnspecified,
17    /// Tuning job is being created.
18    Creating,
19    /// Tuned model is ready for use.
20    Active,
21    /// Terminated with an error.
22    Failed,
23}
24
25/// Supervised tuning specification.
26#[derive(Debug, Clone, Serialize, Deserialize)]
27#[serde(rename_all = "camelCase")]
28pub struct SupervisedTuningSpec {
29    /// Training dataset configuration.
30    pub training_dataset_uri: Option<String>,
31    /// Validation dataset URI.
32    #[serde(default)]
33    pub validation_dataset_uri: Option<String>,
34    /// Hyperparameters.
35    #[serde(default)]
36    pub hyper_parameters: Option<TuningHyperParameters>,
37}
38
39/// Tuning hyperparameters.
40#[derive(Debug, Clone, Serialize, Deserialize)]
41#[serde(rename_all = "camelCase")]
42pub struct TuningHyperParameters {
43    /// Number of training epochs.
44    #[serde(default)]
45    pub epoch_count: Option<u32>,
46    /// Batch size.
47    #[serde(default)]
48    pub batch_size: Option<u32>,
49    /// Learning rate.
50    #[serde(default)]
51    pub learning_rate: Option<f64>,
52}
53
54/// A tuning job resource.
55#[derive(Debug, Clone, Serialize, Deserialize)]
56#[serde(rename_all = "camelCase")]
57pub struct TuningJob {
58    /// Resource name.
59    #[serde(default)]
60    pub name: String,
61    /// Base model being tuned.
62    #[serde(default)]
63    pub base_model: Option<String>,
64    /// Tuned model name (output).
65    #[serde(default)]
66    pub tuned_model: Option<String>,
67    /// Display name.
68    #[serde(default)]
69    pub display_name: Option<String>,
70    /// State of the tuning job.
71    #[serde(default)]
72    pub state: Option<TuningJobState>,
73    /// Supervised tuning spec.
74    #[serde(default)]
75    pub supervised_tuning_spec: Option<SupervisedTuningSpec>,
76    /// Creation time (RFC3339).
77    #[serde(default)]
78    pub create_time: Option<String>,
79    /// Update time (RFC3339).
80    #[serde(default)]
81    pub update_time: Option<String>,
82    /// Error details if state is Failed.
83    #[serde(default)]
84    pub error: Option<serde_json::Value>,
85}
86
87/// Configuration for creating a tuning job.
88#[derive(Debug, Clone)]
89pub struct CreateTuningJobConfig {
90    /// Base model to tune.
91    pub base_model: String,
92    /// Display name.
93    pub display_name: Option<String>,
94    /// Supervised tuning spec.
95    pub supervised_tuning_spec: SupervisedTuningSpec,
96}
97
98/// Response from listTuningJobs.
99#[derive(Debug, Clone, Serialize, Deserialize)]
100#[serde(rename_all = "camelCase")]
101pub struct ListTuningJobsResponse {
102    /// List of tuning jobs.
103    #[serde(default)]
104    pub tuning_jobs: Vec<TuningJob>,
105    /// Pagination token for the next page.
106    #[serde(default)]
107    pub next_page_token: Option<String>,
108}
109
110/// Errors from the Tunings API.
111#[derive(Debug, thiserror::Error)]
112pub enum TuningsError {
113    #[error(transparent)]
114    /// Transport-level HTTP failure.
115    Http(#[from] HttpError),
116    #[error("Failed to parse response: {0}")]
117    /// Response body failed to parse.
118    Parse(#[from] serde_json::Error),
119    #[error("Auth error: {0}")]
120    /// Authentication/authorization failure.
121    Auth(#[from] crate::session::AuthError),
122}
123
124impl Client {
125    /// List tuning jobs.
126    pub async fn list_tuning_jobs(&self) -> Result<ListTuningJobsResponse, TuningsError> {
127        let url = self.rest_url(ServiceEndpoint::TuningJobs);
128        let headers = self.auth_headers().await?;
129        let json = self.http_client().get_json(&url, headers).await?;
130        if json.is_null() {
131            return Ok(ListTuningJobsResponse {
132                tuning_jobs: vec![],
133                next_page_token: None,
134            });
135        }
136        Ok(serde_json::from_value(json)?)
137    }
138
139    /// Get a tuning job by name.
140    pub async fn get_tuning_job(&self, name: &str) -> Result<TuningJob, TuningsError> {
141        let base_url = self.rest_url(ServiceEndpoint::TuningJobs);
142        let url = format!("{base_url}/{name}");
143        let headers = self.auth_headers().await?;
144        let json = self.http_client().get_json(&url, headers).await?;
145        Ok(serde_json::from_value(json)?)
146    }
147
148    /// Create a new tuning job.
149    pub async fn create_tuning_job(
150        &self,
151        config: CreateTuningJobConfig,
152    ) -> Result<TuningJob, TuningsError> {
153        let url = self.rest_url(ServiceEndpoint::TuningJobs);
154        let headers = self.auth_headers().await?;
155
156        let mut body = serde_json::json!({
157            "baseModel": config.base_model,
158            "supervisedTuningSpec": config.supervised_tuning_spec,
159        });
160
161        if let Some(name) = config.display_name {
162            body["displayName"] = serde_json::Value::String(name);
163        }
164
165        let json = self.http_client().post_json(&url, headers, &body).await?;
166        Ok(serde_json::from_value(json)?)
167    }
168
169    /// Cancel a tuning job by name.
170    pub async fn cancel_tuning_job(&self, name: &str) -> Result<(), TuningsError> {
171        let base_url = self.rest_url(ServiceEndpoint::TuningJobs);
172        let url = format!("{base_url}/{name}:cancel");
173        let headers = self.auth_headers().await?;
174        self.http_client()
175            .post_json(&url, headers, &serde_json::json!({}))
176            .await?;
177        Ok(())
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184
185    #[test]
186    fn parse_tuning_job() {
187        let json = serde_json::json!({
188            "name": "tuningJobs/123",
189            "baseModel": "models/gemini-1.5-flash",
190            "tunedModel": "tunedModels/my-model",
191            "displayName": "My Tuning",
192            "state": "ACTIVE",
193            "createTime": "2026-03-01T00:00:00Z"
194        });
195        let job: TuningJob = serde_json::from_value(json).unwrap();
196        assert_eq!(job.name, "tuningJobs/123");
197        assert_eq!(job.state, Some(TuningJobState::Active));
198        assert_eq!(job.tuned_model, Some("tunedModels/my-model".to_string()));
199    }
200
201    #[test]
202    fn parse_list_tuning_jobs_response() {
203        let json = serde_json::json!({
204            "tuningJobs": [
205                {"name": "tuningJobs/1", "state": "CREATING"},
206                {"name": "tuningJobs/2", "state": "ACTIVE"}
207            ]
208        });
209        let resp: ListTuningJobsResponse = serde_json::from_value(json).unwrap();
210        assert_eq!(resp.tuning_jobs.len(), 2);
211        assert_eq!(resp.tuning_jobs[0].state, Some(TuningJobState::Creating));
212    }
213
214    #[test]
215    fn tuning_job_state_serialization() {
216        assert_eq!(
217            serde_json::to_value(TuningJobState::Active).unwrap(),
218            "ACTIVE"
219        );
220        assert_eq!(
221            serde_json::to_value(TuningJobState::Creating).unwrap(),
222            "CREATING"
223        );
224        assert_eq!(
225            serde_json::to_value(TuningJobState::Failed).unwrap(),
226            "FAILED"
227        );
228    }
229
230    #[test]
231    fn supervised_tuning_spec_serialization() {
232        let spec = SupervisedTuningSpec {
233            training_dataset_uri: Some("gs://bucket/train.jsonl".to_string()),
234            validation_dataset_uri: None,
235            hyper_parameters: Some(TuningHyperParameters {
236                epoch_count: Some(5),
237                batch_size: Some(32),
238                learning_rate: Some(0.001),
239            }),
240        };
241        let json = serde_json::to_value(&spec).unwrap();
242        assert_eq!(json["trainingDatasetUri"], "gs://bucket/train.jsonl");
243        assert_eq!(json["hyperParameters"]["epochCount"], 5);
244    }
245
246    #[test]
247    fn empty_list_response() {
248        let json = serde_json::json!({"tuningJobs": []});
249        let resp: ListTuningJobsResponse = serde_json::from_value(json).unwrap();
250        assert!(resp.tuning_jobs.is_empty());
251    }
252}