Skip to main content

rig_core/client/
gemini_caching.rs

1//! The [`Caching`] transport and the cache book's I/O: the part of Gemini's
2//! automatic caching that sends requests. What to read, create and retire is
3//! decided by [`CacheBook`] in [`crate::providers::gemini::caching`], from
4//! bytes alone.
5
6use 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    /// Keep each lease only if Google still has a cache under its name whose
22    /// display name ends with the lease's digest. Returns how many survived.
23    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    /// Delete every cache the book still holds. A cache already gone (403)
46    /// counts as deleted.
47    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/// A transport that caches GenerateContent requests through `inner`, as
63/// its [`CacheBook`] decides. Build it with [`Model::caching`].
64#[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    /// This model, reading and creating explicit caches as `book` decides.
83    /// See [`crate::providers::gemini::caching`].
84    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
102/// `payload` with `bytes` as its body.
103fn 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
124/// A `cachedContents` reply: its status (when it failed) and its body.
125struct 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    /// Send one `cachedContents` request through the inner transport.
136    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    /// Create the cache `create` describes, unless one already exists for
226    /// its digest. Returns the lease to read.
227    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
289/// The decoded value of a classified payload, when it decoded.
290fn known<T>(event: WireEvent<T>) -> Option<T> {
291    match event {
292        WireEvent::Known(value) => Some(value),
293        _ => None,
294    }
295}
296
297/// The final `cachedContentTokenCount` a reply frame reports, when the frame
298/// ends the reply.
299fn 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                    // Deleted elsewhere or expired early: forget it and send inline, once.
415                    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}