1use std::time::Duration;
20
21use serde::{Deserialize, Serialize};
22
23use serde_json::{Value, json};
24
25use crate::error::EncodeError;
26use crate::error::ProviderError;
27use crate::json_utils::Lenient;
28use crate::operation::Whole;
29use crate::providers::internal::{
30 wire::{classify_or, classify_untyped_line},
31 with_query_pairs,
32};
33use crate::wire::{
34 Body, Call, Decoder, Descriptor, Encoded, Flow, Framing, Free, Mode, Operation, Out, Wire,
35 WireEvent, WireFrame,
36};
37
38const CACHED_CONTENTS_PATH: &str = "/v1beta/cachedContents";
40
41const MAX_PAGE_SIZE: usize = 1000;
43
44pub(crate) fn on_handle(error: ProviderError, name: &str) -> ProviderError {
48 match error {
49 ProviderError::ProviderResponse(response)
50 if matches!(
51 response.status,
52 Some(http::StatusCode::FORBIDDEN | http::StatusCode::NOT_FOUND)
53 ) =>
54 {
55 ProviderError::CacheExpired {
56 name: name.to_owned(),
57 response,
58 }
59 }
60 other => other,
61 }
62}
63
64#[derive(Clone, Debug, PartialEq, Eq)]
66pub enum CacheExpiry {
67 Ttl(Duration),
70 ExpireTime(String),
72}
73
74impl CacheExpiry {
75 pub fn ttl(ttl: Duration) -> Self {
76 Self::Ttl(ttl)
77 }
78
79 pub fn expire_time(timestamp: impl Into<String>) -> Self {
80 Self::ExpireTime(timestamp.into())
81 }
82
83 fn ttl_string(ttl: Duration) -> String {
85 format!("{}.{:09}s", ttl.as_secs(), ttl.subsec_nanos())
86 }
87}
88
89#[derive(Debug, Default, Serialize)]
93#[serde(rename_all = "camelCase")]
94pub struct NewCachedContent {
95 model: String,
98 #[serde(skip_serializing_if = "Vec::is_empty")]
99 contents: Vec<Value>,
100 #[serde(skip_serializing_if = "Option::is_none")]
101 system_instruction: Option<Value>,
102 #[serde(skip_serializing_if = "Option::is_none")]
103 tools: Option<Vec<Value>>,
104 #[serde(skip_serializing_if = "Option::is_none")]
105 tool_config: Option<Value>,
106 #[serde(skip_serializing_if = "Option::is_none")]
107 display_name: Option<String>,
108 #[serde(skip_serializing_if = "Option::is_none")]
109 ttl: Option<String>,
110 #[serde(skip_serializing_if = "Option::is_none")]
111 expire_time: Option<String>,
112}
113
114impl NewCachedContent {
115 pub fn new(model: impl AsRef<str>) -> Self {
121 Self {
122 model: qualify_model(model.as_ref()),
123 ..Default::default()
124 }
125 }
126
127 pub fn content(mut self, text: impl Into<String>) -> Self {
129 let text = text.into();
130 self.contents
131 .push(json!({ "parts": [{ "text": text, "thought": false }], "role": "user" }));
132 self
133 }
134
135 pub fn content_block(mut self, content: Value) -> Self {
137 self.contents.push(content);
138 self
139 }
140
141 pub fn system_instruction(mut self, text: impl Into<String>) -> Self {
142 let text = text.into();
143 self.system_instruction =
144 Some(json!({ "parts": [{ "text": text, "thought": false }], "role": "model" }));
145 self
146 }
147
148 pub fn tools(mut self, tools: Vec<Value>) -> Self {
153 self.tools = Some(tools);
154 self
155 }
156
157 pub fn tool_config(mut self, tool_config: Value) -> Self {
162 self.tool_config = Some(tool_config);
163 self
164 }
165
166 pub fn display_name(mut self, name: impl Into<String>) -> Self {
167 self.display_name = Some(name.into());
168 self
169 }
170
171 pub fn expiry(mut self, expiry: CacheExpiry) -> Self {
174 match expiry {
175 CacheExpiry::Ttl(ttl) => {
176 self.ttl = Some(CacheExpiry::ttl_string(ttl));
177 self.expire_time = None;
178 }
179 CacheExpiry::ExpireTime(at) => {
180 self.expire_time = Some(at);
181 self.ttl = None;
182 }
183 }
184 self
185 }
186
187 fn validate(&self) -> Result<(), EncodeError> {
188 if self.contents.is_empty() && self.system_instruction.is_none() {
189 return Err(EncodeError::request(
190 "a cached content needs contents or a system instruction; an empty cache would \
191 bill for storage and cache nothing",
192 ));
193 }
194 Ok(())
195 }
196}
197
198#[derive(Clone, Debug, Default, Deserialize, Serialize)]
200#[serde(rename_all = "camelCase")]
201pub struct CachedContentUsage {
202 #[serde(default)]
205 pub total_token_count: u64,
206}
207
208#[derive(Clone, Debug, Deserialize, Serialize)]
210#[serde(rename_all = "camelCase")]
211pub struct CachedContent {
212 pub name: String,
215 #[serde(default)]
217 pub model: String,
218 #[serde(default, skip_serializing_if = "Option::is_none")]
219 pub display_name: Option<String>,
220 #[serde(default, skip_serializing_if = "Option::is_none")]
221 pub create_time: Option<String>,
222 #[serde(default, skip_serializing_if = "Option::is_none")]
223 pub update_time: Option<String>,
224 #[serde(default, skip_serializing_if = "Option::is_none")]
227 pub expire_time: Option<String>,
228 #[serde(default, skip_serializing_if = "Option::is_none")]
229 pub usage_metadata: Option<CachedContentUsage>,
230}
231
232#[derive(Debug)]
234pub enum CachedContentRequest {
235 Create(NewCachedContent),
237 Get(String),
239 List {
242 page_token: Option<String>,
244 },
245 UpdateExpiry { name: String, expiry: CacheExpiry },
248 Delete(String),
250}
251
252#[derive(Clone, Debug, Deserialize)]
254#[serde(rename_all = "camelCase")]
255pub struct CachedContentPage {
256 pub cached_contents: Vec<CachedContent>,
259 #[serde(default)]
263 pub next_page_token: Option<String>,
264}
265
266#[derive(Clone, Debug, Default)]
269pub enum CachedContentReply {
270 Resource(CachedContent),
272 Page(CachedContentPage),
274 #[default]
277 Acknowledged,
278}
279
280impl CachedContentReply {
281 pub fn resource(self) -> Result<CachedContent, ProviderError> {
284 match self {
285 Self::Resource(resource) => Ok(resource),
286 other => Err(other.mismatch("one cached content")),
287 }
288 }
289
290 pub fn next_page_token(&self) -> Option<String> {
293 match self {
294 Self::Page(page) => page
295 .next_page_token
296 .clone()
297 .filter(|token| !token.is_empty()),
298 _ => None,
299 }
300 }
301
302 pub fn entries(self) -> Result<Vec<CachedContent>, ProviderError> {
305 match self {
306 Self::Page(page) => Ok(page.cached_contents),
307 Self::Acknowledged => Ok(Vec::new()),
308 other => Err(other.mismatch("a listing page")),
309 }
310 }
311
312 fn mismatch(&self, wanted: &str) -> ProviderError {
314 let carried = match self {
315 Self::Resource(_) => "one cached content",
316 Self::Page(_) => "a listing page",
317 Self::Acknowledged => "nothing to read, only a success status",
318 };
319 ProviderError::Response(format!("the reply carried {carried}, not {wanted}"))
320 }
321}
322
323#[derive(Debug, Clone, Copy, PartialEq, Eq)]
326pub struct ContextCache;
327
328impl Operation for ContextCache {
329 type Request = CachedContentRequest;
330 type Event = std::convert::Infallible;
331 type End = CachedContentReply;
332 type Response = CachedContentReply;
333 type Fold = Whole<Self>;
334 type Emit = Free;
335
336 fn fold(_request: &Self::Request, _call: &mut Call<'_>) -> Self::Fold {
337 Whole::new()
338 }
339}
340
341#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
347pub struct CachedContents {
348 pub provider: super::GeminiConfig,
350 pub page_size: usize,
352}
353
354impl CachedContents {
355 pub fn new(provider: super::GeminiConfig) -> Self {
357 Self {
358 provider,
359 page_size: MAX_PAGE_SIZE,
360 }
361 }
362
363 pub fn with_page_size(mut self, page_size: usize) -> Self {
365 self.page_size = page_size;
366 self
367 }
368
369 fn list_request(&self, page_token: Option<&str>) -> Result<http::Request<Body>, http::Error> {
372 let page_size = self.page_size.to_string();
373 let mut pairs = vec![("pageSize", page_size.as_str())];
374 if let Some(token) = page_token {
375 pairs.push(("pageToken", token));
376 }
377 let path = with_query_pairs(CACHED_CONTENTS_PATH, &pairs);
378 http::Request::get(self.provider.uri(&path)).body(Body::empty())
379 }
380}
381
382pub fn with_cached_content(
389 body: &mut serde_json::Map<String, Value>,
390 name: &str,
391) -> Result<(), EncodeError> {
392 let fail = |message: String| Err(EncodeError::request(message));
393 if !name.starts_with("cachedContents/") {
394 return fail(format!(
395 "gemini cached content handle should look like `cachedContents/<id>`, got `{name}`"
396 ));
397 }
398 if let Some(existing) = body
399 .get("cachedContent")
400 .and_then(Value::as_str)
401 .filter(|existing| *existing != name)
402 {
403 return fail(format!(
404 "a Gemini request set cached content twice, to `{existing}` and `{name}`: set it \
405 one way or the other"
406 ));
407 }
408 let conflicts: Vec<&str> = [
409 (
410 &super::completion::SYSTEM_INSTRUCTION[..],
411 "a system instruction (preamble)",
412 ),
413 (&["tools"][..], "tools"),
414 (&super::completion::TOOL_CONFIG[..], "a tool choice"),
415 ]
416 .into_iter()
417 .filter_map(|(spellings, what)| super::completion::present(body, spellings).map(|_| what))
418 .collect();
419 if conflicts.is_empty() {
420 body.insert("cachedContent".to_owned(), Value::String(name.to_owned()));
421 return Ok(());
422 }
423 let tools = body
426 .get("tools")
427 .and_then(Value::as_array)
428 .map_or(&[][..], Vec::as_slice);
429 let declares_functions = tools.iter().any(|tool| {
430 !tool.arr("functionDeclarations").is_empty()
431 || !tool.arr("function_declarations").is_empty()
432 });
433 let caveat = match declares_functions {
434 true => {
435 " Function declarations in a cache are declarations only: an `Agent` dispatches \
436 only tools it advertised, so a cached function tool runs only when you drive \
437 `GenerateContent` yourself. Hosted tools such as `codeExecution` are fine to cache."
438 }
439 false => "",
440 };
441 fail(format!(
442 "a Gemini request using cached content `{name}` also set {}. The cached content owns \
443 the system instruction, tools and tool choice of every request that uses it: move \
444 them into the cache, or drop the cache handle.{caveat}",
445 conflicts.join(" and ")
446 ))
447}
448
449impl super::GeminiConfig {
450 pub(crate) fn cached_contents(&self) -> CachedContents {
452 CachedContents::new(self.clone())
453 }
454}
455
456impl Wire for CachedContents {
457 type Op = ContextCache;
458 type Payload = crate::wire::Encoded;
459 type Frame = crate::wire::WireFrame;
460 type Decoder<'id> = CachedContentsDecoder;
461 type Reassembler = crate::wire::document::Unreassembled;
462
463 fn describe(&self) -> Descriptor<'_> {
464 Descriptor::new(super::PROVIDER_NAME)
465 }
466
467 fn encode(&self, request: CachedContentRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
471 let request = match request {
472 CachedContentRequest::Create(new) => {
473 new.validate()?;
474 http::Request::post(self.provider.uri(CACHED_CONTENTS_PATH))
475 .body(Body::Bytes(serde_json::to_vec(&new)?))?
476 }
477 CachedContentRequest::Get(name) => {
478 http::Request::get(self.provider.uri(&resource_path(&name)?)).body(Body::empty())?
479 }
480 CachedContentRequest::List { page_token } => {
481 self.list_request(page_token.as_deref())?
482 }
483 CachedContentRequest::UpdateExpiry { name, expiry } => {
484 let (patch, mask) = expiry_patch(expiry)?;
485 let path = format!("{}?updateMask={mask}", resource_path(&name)?);
487 http::Request::patch(self.provider.uri(&path)).body(Body::Bytes(patch))?
488 }
489 CachedContentRequest::Delete(name) => {
490 http::Request::delete(self.provider.uri(&resource_path(&name)?))
491 .body(Body::empty())?
492 }
493 };
494 Ok(Encoded::new(request, Framing::Whole))
495 }
496
497 fn decoder<'id>(&self) -> Self::Decoder<'id> {
498 CachedContentsDecoder
499 }
500}
501
502pub struct CachedContentsDecoder;
504
505impl<'id> Decoder<'id, ContextCache> for CachedContentsDecoder {
506 type Event = CachedContentReply;
507
508 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
511 let body = frame.as_str();
512 classify_or(&body, as_page, |data| {
513 classify_or(data, as_acknowledgement, as_resource)
514 })
515 }
516
517 fn decode(
518 &mut self,
519 reply: Self::Event,
520 out: Out<'id, ContextCache>,
521 ) -> Result<Flow, ProviderError> {
522 Ok(out.end(reply))
523 }
524
525 fn eof(&mut self, out: Out<'id, ContextCache>) -> Result<Flow, ProviderError> {
528 Ok(out.end(CachedContentReply::Acknowledged))
529 }
530}
531
532fn as_page(data: &str) -> WireEvent<CachedContentReply> {
534 classify_untyped_line::<CachedContentPage>(data.as_bytes()).map(CachedContentReply::Page)
535}
536
537fn as_acknowledgement(data: &str) -> WireEvent<CachedContentReply> {
539 #[derive(Deserialize)]
540 #[serde(deny_unknown_fields)]
541 struct Acknowledgement {}
542
543 classify_untyped_line::<Acknowledgement>(data.as_bytes())
544 .map(|_| CachedContentReply::Acknowledged)
545}
546
547fn as_resource(data: &str) -> WireEvent<CachedContentReply> {
549 classify_untyped_line::<CachedContent>(data.as_bytes()).map(CachedContentReply::Resource)
550}
551
552fn expiry_patch(expiry: CacheExpiry) -> Result<(Vec<u8>, &'static str), EncodeError> {
554 let (field, value) = match expiry {
555 CacheExpiry::Ttl(ttl) => ("ttl", CacheExpiry::ttl_string(ttl)),
556 CacheExpiry::ExpireTime(at) => ("expireTime", at),
557 };
558 let patch = serde_json::Map::from_iter([(field.to_owned(), serde_json::Value::String(value))]);
559 Ok((serde_json::to_vec(&patch)?, field))
560}
561
562fn qualify_model(model: &str) -> String {
564 if model.starts_with("models/") {
565 model.to_owned()
566 } else {
567 format!("models/{model}")
568 }
569}
570
571fn resource_path(name: &str) -> Result<String, EncodeError> {
575 let id = name.strip_prefix("cachedContents/").unwrap_or(name);
576 let is_id_char = |ch: char| ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_');
577 if id.is_empty() || !id.chars().all(is_id_char) {
578 return Err(EncodeError::request(format!(
579 "`{name}` is not a cached content handle; expected `cachedContents/<id>` or a bare \
580 `<id>` of letters, digits, `-` and `_`. The id is spliced into the request path, \
581 where a `?`, `#` or `/` silently retargets the call at a different resource — \
582 and this is the path that deletes"
583 )));
584 }
585 Ok(format!("{CACHED_CONTENTS_PATH}/{id}"))
586}
587
588#[cfg(test)]
589mod tests;
590
591#[cfg(test)]
592mod exhaustive_validation_tests;
593
594#[cfg(test)]
595mod status_triage_tests;