1use futures::StreamExt;
7use serde::Deserialize;
8
9use crate::driver::{Exchange, Model, Opened, Opening, Transport};
10use crate::error::ProviderError;
11use crate::providers::gemini::cached_content::CachedContents;
12use crate::providers::gemini::caching::{
13 CacheBook, Create, Lease, Parsed, cache_body, digests, is_user_text, parse, short_digest,
14 stripped,
15};
16use crate::providers::gemini::completion::{GenerateContent, ThoughtReplay};
17use crate::providers::internal::wire::classify_untyped_line;
18use crate::wire::{Body, Encoded, Framing, Mode, WireEvent, WireFrame};
19
20impl CacheBook {
21 pub async fn prove<T>(&self, caches: &Model<CachedContents, T>) -> usize
24 where
25 T: Transport<CachedContents>,
26 {
27 let mut survived = 0;
28 for lease in self.leases() {
29 let alive = match caches.get(&lease.name).await {
30 Ok(resource) => resource
31 .display_name
32 .as_deref()
33 .is_some_and(|name| name.ends_with(short_digest(&lease.digest))),
34 Err(_) => false,
35 };
36 if alive {
37 survived += 1;
38 } else {
39 self.lost(&lease);
40 }
41 }
42 survived
43 }
44
45 pub async fn close<T>(&self, caches: &Model<CachedContents, T>)
48 where
49 T: Transport<CachedContents>,
50 {
51 for lease in self.leases() {
52 let result = caches.delete(&lease.name).await;
53 let gone = matches!(result, Ok(()) | Err(ProviderError::CacheExpired { .. }));
54 if !gone {
55 tracing::warn!(target: "gemini.cache", name = %lease.name, "cache delete failed");
56 }
57 self.retired(&lease);
58 }
59 }
60}
61
62#[derive(Clone)]
65pub struct Caching<T> {
66 inner: T,
67 config: crate::providers::gemini::GeminiConfig,
68 book: CacheBook,
69 thought_replay: ThoughtReplay,
70}
71
72impl<T> std::fmt::Debug for Caching<T> {
73 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
74 f.debug_struct("Caching")
75 .field("book", &self.book)
76 .field("thought_replay", &self.thought_replay)
77 .finish_non_exhaustive()
78 }
79}
80
81impl<T> Model<GenerateContent, T> {
82 pub fn caching(self, book: &CacheBook) -> Model<GenerateContent, Caching<T>> {
85 let caching = Caching {
86 inner: self.transport,
87 config: self.wire.provider.clone(),
88 book: book.clone(),
89 thought_replay: self.wire.thought_replay,
90 };
91 Model::new(self.wire, caching)
92 }
93}
94
95fn model_of(path: &str) -> Option<String> {
96 let rest = path.split("/models/").nth(1)?;
97 let (model, verb) = rest.split_once(':')?;
98 let verb = verb.split('?').next()?;
99 matches!(verb, "generateContent" | "streamGenerateContent").then(|| model.to_owned())
100}
101
102fn with_body(payload: &Encoded, bytes: Vec<u8>) -> Result<Encoded, ProviderError> {
104 let mut builder = http::Request::builder()
105 .method(payload.request.method().clone())
106 .uri(payload.request.uri().clone());
107 for (name, value) in payload.request.headers() {
108 builder = builder.header(name, value);
109 }
110 let request = builder
111 .body(Body::Bytes(bytes))
112 .map_err(|error| ProviderError::request(error.to_string()))?;
113 Ok(Encoded {
114 request,
115 framing: payload.framing,
116 request_id_header: payload.request_id_header,
117 relaxed_content_type: payload.relaxed_content_type,
118 route: payload.route,
119 project: payload.project,
120 analysis_only: payload.analysis_only,
121 })
122}
123
124struct Reply {
126 status: Option<u16>,
127 body: String,
128 error: Option<String>,
129}
130
131impl<T> Caching<T>
132where
133 T: Transport<GenerateContent>,
134{
135 async fn resource(&self, method: http::Method, path: &str, body: Option<Vec<u8>>) -> Reply {
137 let request = http::Request::builder()
138 .method(method)
139 .uri(self.config.uri(path))
140 .header("Content-Type", "application/json")
141 .body(Body::Bytes(body.unwrap_or_default()));
142 let request = match request {
143 Ok(request) => request,
144 Err(error) => {
145 return Reply {
146 status: None,
147 body: String::new(),
148 error: Some(error.to_string()),
149 };
150 }
151 };
152 let exchange = Exchange {
153 mode: Mode::Unary,
154 observation: None,
155 };
156 let opened = match self
157 .inner
158 .send(Encoded::new(request, Framing::Whole), exchange)
159 .await
160 {
161 Ok(opened) => opened,
162 Err(error) => {
163 return Reply {
164 status: error.provider_response_status().map(|s| s.as_u16()),
165 body: String::new(),
166 error: Some(error.to_string()),
167 };
168 }
169 };
170 let mut frames = opened.frames;
171 let mut body = String::new();
172 while let Some(frame) = frames.next().await {
173 match frame {
174 Ok(frame) => body.push_str(&frame.as_str()),
175 Err(error) => {
176 return Reply {
177 status: error
178 .provider_response_status()
179 .map(|status| status.as_u16()),
180 body: error
181 .provider_response_body()
182 .unwrap_or_default()
183 .to_owned(),
184 error: Some(error.to_string()),
185 };
186 }
187 }
188 }
189 Reply {
190 status: None,
191 body,
192 error: None,
193 }
194 }
195
196 async fn delete(&self, lease: &Lease) {
197 let reply = self
198 .resource(
199 http::Method::DELETE,
200 &format!("/v1beta/{}", lease.name),
201 None,
202 )
203 .await;
204 if reply.error.is_some() && reply.status != Some(403) && reply.status != Some(404) {
205 let status = reply.status.unwrap_or_default();
206 tracing::warn!(target: "gemini.cache", name = %lease.name, status, "cache delete failed");
207 }
208 self.book.retired(lease);
209 }
210
211 async fn extend(&self, lease: &Lease, ttl_secs: u64) {
212 let body = format!("{{\"ttl\":\"{ttl_secs}s\"}}").into_bytes();
213 let reply = self
214 .resource(
215 http::Method::PATCH,
216 &format!("/v1beta/{}?updateMask=ttl", lease.name),
217 Some(body),
218 )
219 .await;
220 if reply.error.is_none() {
221 self.book.extended(lease, ttl_secs);
222 }
223 }
224
225 async fn create(
228 &self,
229 model: &str,
230 parsed: &Parsed,
231 create: &Create,
232 line: &str,
233 coverable: u64,
234 ) -> Option<Lease> {
235 let _single = self.book.create.lock().await;
236 if let Some(existing) = self.book.book().leases.get(&create.digest).cloned() {
237 return Some(existing);
238 }
239 let display_name = format!(
240 "{}{}",
241 self.book.display_prefix,
242 short_digest(&create.digest)
243 );
244 let body = cache_body(model, parsed, create.covers, &display_name, create.ttl_secs)?;
245 let reply = self
246 .resource(http::Method::POST, "/v1beta/cachedContents", Some(body))
247 .await;
248 if reply.error.is_some() {
249 #[derive(Deserialize)]
250 struct Envelope {
251 error: ErrorBody,
252 }
253 #[derive(Deserialize)]
254 struct ErrorBody {
255 message: String,
256 }
257 let message = known(classify_untyped_line::<Envelope>(reply.body.as_bytes()))
258 .map(|envelope| envelope.error.message)
259 .or(reply.error)
260 .unwrap_or_default();
261 self.book
262 .create_failed(line, coverable, reply.status, message);
263 return None;
264 }
265 #[derive(Deserialize)]
266 #[serde(rename_all = "camelCase")]
267 struct Resource {
268 name: String,
269 #[serde(default)]
270 usage_metadata: Option<Usage>,
271 }
272 #[derive(Deserialize)]
273 #[serde(rename_all = "camelCase")]
274 struct Usage {
275 #[serde(default)]
276 total_token_count: u64,
277 }
278 let resource = known(classify_untyped_line::<Resource>(reply.body.as_bytes()))?;
279 let tokens = resource
280 .usage_metadata
281 .map_or(0, |usage| usage.total_token_count);
282 Some(
283 self.book
284 .created(line, create, resource.name, tokens, model),
285 )
286 }
287}
288
289fn known<T>(event: WireEvent<T>) -> Option<T> {
291 match event {
292 WireEvent::Known(value) => Some(value),
293 _ => None,
294 }
295}
296
297fn final_cached(frame: &WireFrame) -> Option<u64> {
300 #[derive(Deserialize)]
301 #[serde(rename_all = "camelCase")]
302 struct Usage {
303 #[serde(default)]
304 cached_content_token_count: u64,
305 }
306 #[derive(Deserialize)]
307 #[serde(rename_all = "camelCase")]
308 struct Candidate {
309 finish_reason: Option<String>,
310 }
311 #[derive(Deserialize)]
312 #[serde(rename_all = "camelCase")]
313 struct Reply {
314 usage_metadata: Option<Usage>,
315 #[serde(default)]
316 candidates: Vec<Candidate>,
317 }
318 let reply = known(classify_untyped_line::<Reply>(frame.as_str().as_bytes()))?;
319 let finished = reply
320 .candidates
321 .iter()
322 .any(|candidate| candidate.finish_reason.is_some());
323 finished.then_some(reply.usage_metadata?.cached_content_token_count)
324}
325
326impl<T> Transport<GenerateContent> for Caching<T>
327where
328 T: Transport<GenerateContent>,
329{
330 fn send(&self, payload: Encoded, exchange: Exchange) -> Opening<WireFrame> {
331 let this = self.clone();
332 Opening::new(async move {
333 let Exchange { mode, observation } = exchange;
334 let path = payload.request.uri().path().to_owned();
335 let bytes = match payload.request.body() {
336 Body::Bytes(bytes) => Some(bytes.clone()),
337 Body::Multipart(_) => None,
338 };
339 let parsed = bytes.as_deref().and_then(parse);
340 let (Some(bytes), Some(model), Some(parsed)) = (bytes, model_of(&path), parsed) else {
341 return this
342 .inner
343 .send(payload, Exchange { mode, observation })
344 .await;
345 };
346 if parsed.has_cached_content {
347 return this
348 .inner
349 .send(payload, Exchange { mode, observation })
350 .await;
351 }
352
353 let d = digests(&model, &parsed);
354 let roll_allowed = match this.thought_replay {
355 ThoughtReplay::All => true,
356 ThoughtReplay::CurrentTurn => {
357 parsed.contents.last().is_some_and(|c| is_user_text(c))
358 }
359 };
360 let plan = this.book.plan(&model, &d, &parsed, roll_allowed);
361 for idle in &plan.retire {
362 this.delete(idle).await;
363 }
364 let mut read = plan.read.clone();
365 let mut line = plan.line.clone();
366 if let Some(create) = &plan.create
367 && let Some(lease) = this
368 .create(&model, &parsed, create, &plan.line, plan.coverable)
369 .await
370 {
371 if let Some(replaced) = &create.replaces
372 && replaced.name != lease.name
373 {
374 this.delete(replaced).await;
375 }
376 if lease.covers > 0 {
377 line = lease.digest.clone();
378 }
379 read = Some(lease);
380 }
381 if let Some((lease, ttl_secs)) = &plan.extend
382 && read.as_ref().is_some_and(|read| read.name == lease.name)
383 {
384 this.extend(lease, *ttl_secs).await;
385 }
386
387 let sent = match &read {
388 Some(lease) => match stripped(&parsed, lease) {
389 Some(body) => with_body(&payload, body)?,
390 None => {
391 read = None;
392 with_body(&payload, bytes.clone())?
393 }
394 },
395 None => with_body(&payload, bytes.clone())?,
396 };
397 let mut opened: Opened<WireFrame> = this
398 .inner
399 .send(
400 sent,
401 Exchange {
402 mode,
403 observation: observation.clone(),
404 },
405 )
406 .await?;
407 if let Some(lease) = read.clone() {
408 let first = opened.frames.next().await;
409 let forbidden = matches!(
410 &first,
411 Some(Err(error)) if error.provider_response_status() == Some(http::StatusCode::FORBIDDEN)
412 );
413 if forbidden {
414 this.book.lost(&lease);
416 read = None;
417 opened = this
418 .inner
419 .send(with_body(&payload, bytes)?, Exchange { mode, observation })
420 .await?;
421 } else {
422 this.book.touched(&lease);
423 let rest =
424 std::mem::replace(&mut opened.frames, Box::pin(futures::stream::empty()));
425 opened.frames = Box::pin(futures::stream::iter(first).chain(rest));
426 }
427 }
428
429 let book = this.book.clone();
430 let read_name = read.map(|lease| lease.name);
431 let coverable = plan.coverable;
432 Ok(opened.map_frames(move |frames| {
433 frames.inspect(move |frame| {
434 if let Ok(frame) = frame
435 && let Some(cached) = final_cached(frame)
436 {
437 book.observe(&line, read_name.as_deref(), coverable, cached);
438 }
439 })
440 }))
441 })
442 }
443}