1use serde::{Deserialize, Serialize};
6
7use crate::client::Client;
8use crate::client::http::HttpError;
9use crate::transport::auth::ServiceEndpoint;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
13#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
14pub enum TuningJobState {
15 StateUnspecified,
17 Creating,
19 Active,
21 Failed,
23}
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
27#[serde(rename_all = "camelCase")]
28pub struct SupervisedTuningSpec {
29 pub training_dataset_uri: Option<String>,
31 #[serde(default)]
33 pub validation_dataset_uri: Option<String>,
34 #[serde(default)]
36 pub hyper_parameters: Option<TuningHyperParameters>,
37}
38
39#[derive(Debug, Clone, Serialize, Deserialize)]
41#[serde(rename_all = "camelCase")]
42pub struct TuningHyperParameters {
43 #[serde(default)]
45 pub epoch_count: Option<u32>,
46 #[serde(default)]
48 pub batch_size: Option<u32>,
49 #[serde(default)]
51 pub learning_rate: Option<f64>,
52}
53
54#[derive(Debug, Clone, Serialize, Deserialize)]
56#[serde(rename_all = "camelCase")]
57pub struct TuningJob {
58 #[serde(default)]
60 pub name: String,
61 #[serde(default)]
63 pub base_model: Option<String>,
64 #[serde(default)]
66 pub tuned_model: Option<String>,
67 #[serde(default)]
69 pub display_name: Option<String>,
70 #[serde(default)]
72 pub state: Option<TuningJobState>,
73 #[serde(default)]
75 pub supervised_tuning_spec: Option<SupervisedTuningSpec>,
76 #[serde(default)]
78 pub create_time: Option<String>,
79 #[serde(default)]
81 pub update_time: Option<String>,
82 #[serde(default)]
84 pub error: Option<serde_json::Value>,
85}
86
87#[derive(Debug, Clone)]
89pub struct CreateTuningJobConfig {
90 pub base_model: String,
92 pub display_name: Option<String>,
94 pub supervised_tuning_spec: SupervisedTuningSpec,
96}
97
98#[derive(Debug, Clone, Serialize, Deserialize)]
100#[serde(rename_all = "camelCase")]
101pub struct ListTuningJobsResponse {
102 #[serde(default)]
104 pub tuning_jobs: Vec<TuningJob>,
105 #[serde(default)]
107 pub next_page_token: Option<String>,
108}
109
110#[derive(Debug, thiserror::Error)]
112pub enum TuningsError {
113 #[error(transparent)]
114 Http(#[from] HttpError),
116 #[error("Failed to parse response: {0}")]
117 Parse(#[from] serde_json::Error),
119 #[error("Auth error: {0}")]
120 Auth(#[from] crate::session::AuthError),
122}
123
124impl Client {
125 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 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 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 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}