Skip to content

Context Window Planner — Rust source

Paste your system prompt, docs, and history — see how they fill any model's context window, with overflow warnings and output headroom.

This is the Rust implementation — the same logic the interactive tool runs, in a shareable, citable form.

//! Context Window Planner — plan labeled prompt sections against a model's
//! context window.
//!
//! Language: Rust (edition 2021, standard library only)
//! Source:   CosmoDev polyglot showcase port of the Context Window Planner
//!           tool, ported from src/lib/contextPlanner.ts (the canonical
//!           TypeScript implementation).
//! Live at:  https://dev.cosmolabs.org/tools/context-window-planner
//! License:  display source — part of CosmoDev's polyglot tool pages.
//!
//! Design goals:
//!   - Pure + deterministic; never panics (public API returns plain values).
//!   - Functionally equivalent to the TS reference: same inputs -> same outputs.
//!   - Self-contained: std only (no crates.io dependencies — no `serde_json`,
//!     no `regex`).
//!
//! Port notes: the TS lib delegates to two siblings — `estimateTokens` from
//! src/lib/tokenEstimator.ts and `fitsWindow` from src/lib/ai/models.ts (which
//! defaults to the bundled pricing snapshot, src/data/ai-models.json). A
//! dependency-free port cannot load that file, so the estimator is inlined
//! below in the exact form the planner uses it (`estimateTokens(text).tokens`,
//! auto content type — the full heuristic lives in the token-estimator port),
//! window math is inlined from `fitsWindow` and `models` is an explicit
//! parameter, never re-derived.
//!
//! Faithfulness notes (the places Rust's std silently differs from JS):
//!   - Length: TS's `String.length` counts UTF-16 code units (an astral-plane
//!     character — emoji, rare CJK ext-B ideographs — counts as 2). Rust
//!     `str::chars()` counts Unicode scalar values, so line arithmetic goes
//!     through `s.encode_utf16().count()` to count the same unit.
//!   - JSON: std has no JSON parser and `serde_json` is off-limits, so this
//!     port ships a small strict recursive-descent validator (`is_valid_json`)
//!     implementing exactly the grammar `JSON.parse` accepts. Whole-text JSON
//!     detection therefore behaves identically, not approximately.
//!   - Rounding: `f64::round` rounds halfway cases away from zero, which for
//!     the non-negative numbers used here is exactly JS `Math.round`.

/// One labeled block of the prompt (system / docs / history / ...).
/// Mirrors the TS `PlanSection` interface.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PlanSection {
    /// Section label, e.g. "system" or "docs".
    pub label: String,
    /// The section's raw text.
    pub text: String,
}

/// Convenience constructor mirroring the TS object literal `{ label, text }`.
pub fn sec(label: &str, text: &str) -> PlanSection {
    PlanSection { label: label.to_string(), text: text.to_string() }
}

/// The subset of the TS `AiModel` record the planner reads. Production code
/// passes the full snapshot entry; only these fields influence the plan.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Model {
    /// Model id, e.g. "beta-pro".
    pub id: String,
    /// Total context window in tokens.
    pub context_window: usize,
    /// The model's output cap (informational).
    pub max_output: usize,
}

/// Sample table for standalone use (mirrors the shared test fixtures).
/// Production code passes the model snapshot instead. Rust's `const` can't
/// allocate `String`s, so the table is a constructor function — the moral
/// twin of the JS/Python/PHP `SAMPLE_MODELS` constant.
pub fn sample_models() -> Vec<Model> {
    vec![
        Model { id: "alpha-mini".to_string(), context_window: 200_000, max_output: 10_000 },
        Model { id: "beta-pro".to_string(), context_window: 1_000_000, max_output: 10_000 },
        Model { id: "gamma-open".to_string(), context_window: 100_000, max_output: 10_000 },
    ]
}

/// Result of `plan_window`. Field-for-field twin of the TS `WindowPlan`
/// interface.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct WindowPlan {
    /// The model id planned against.
    pub id: String,
    /// Sum of per-section token estimates.
    pub input_tokens: usize,
    /// The model's context window.
    pub context_window: usize,
    /// Context tokens left after the request; underflow-safe signed view of
    /// `context_window - input_tokens` (negative on overflow).
    pub free: isize,
    /// Raw fit: free >= 0.
    pub fits: bool,
    /// Room for the output reserve: free >= reserve.
    pub output_reserve_ok: bool,
    /// The model's output cap (informational).
    pub max_output: usize,
}

/// Content classification of a single line. The planner only needs each
/// type's chars-per-token rate (mirrors `CHARS_PER_TOKEN` in
/// src/lib/tokenEstimator.ts: prose 4, code 3.5, json 3, cjk 1.5).
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ContentType {
    Prose,
    Code,
    Json,
    Cjk,
}

impl ContentType {
    fn chars_per_token(self) -> f64 {
        match self {
            ContentType::Prose => 4.0,
            ContentType::Code => 3.5,
            ContentType::Json => 3.0,
            ContentType::Cjk => 1.5,
        }
    }
}

/// Length of `s` in UTF-16 code units — the unit TS's `String.length`
/// counts. BMP code points are one unit, astral-plane ones two.
fn utf16_len(s: &str) -> usize {
    s.encode_utf16().count()
}

/// Reports whether `s` contains a CJK ideograph (U+4E00–U+9FFF), kana
/// (U+3040–U+30FF), or a Hangul syllable (U+AC00–U+D7AF). Mirrors `CJK_RE`
/// in the TS lib.
fn has_cjk(s: &str) -> bool {
    s.chars().any(|c| {
        ('\u{4E00}'..='\u{9FFF}').contains(&c)
            || ('\u{3040}'..='\u{30FF}').contains(&c)
            || ('\u{AC00}'..='\u{D7AF}').contains(&c)
    })
}

/// Reports whether `c` is one of the code-flavored symbols counted by
/// `CODE_SYMBOL_RE` (`{}();=<>[]#`).
fn is_code_symbol(c: char) -> bool {
    matches!(c, '{' | '}' | '(' | ')' | ';' | '=' | '<' | '>' | '[' | ']' | '#')
}

/// Splits `text` on LF or CRLF, mirroring `text.split(/\r?\n/)`: strip the
/// optional CR that belongs to the newline, then split on LF. A lone CR is
/// NOT a line break.
fn split_lines(text: &str) -> Vec<&str> {
    text.split('\n')
        .map(|line| line.strip_suffix('\r').unwrap_or(line))
        .collect()
}

/// Classifies a single line by its shape. Order: json, cjk, code, prose.
/// Inlined from `detectLineType()` in src/lib/tokenEstimator.ts.
fn detect_line_type(line: &str) -> ContentType {
    let trimmed = line.trim();
    // JSON-ish: opens like a JSON fragment AND carries a separator.
    let starts_jsonish = trimmed.starts_with('{')
        || trimmed.starts_with('}')
        || trimmed.starts_with('[')
        || trimmed.starts_with('"');
    if starts_jsonish && (line.contains(':') || line.contains(',')) {
        return ContentType::Json;
    }
    // CJK ideographs / kana / Hangul pack roughly one token per 1.5 chars.
    if has_cjk(line) {
        return ContentType::Cjk;
    }
    // Code: symbol-dense, or a statement terminator / block opener at EOL.
    let length = utf16_len(line);
    let symbols = line.chars().filter(|c| is_code_symbol(*c)).count();
    let density = if length > 0 { symbols as f64 / length as f64 } else { 0.0 };
    if density > 0.08 || trimmed.ends_with(';') || trimmed.ends_with('{') || trimmed.ends_with('}')
    {
        return ContentType::Code;
    }
    ContentType::Prose
}

/// A strict JSON syntax validator — the exact grammar `JSON.parse` accepts,
/// walked with a byte cursor. (Multibyte UTF-8 inside strings never contains
/// an ASCII byte, so byte-level scanning is safe.)
struct JsonParser<'a> {
    bytes: &'a [u8],
    pos: usize,
}

impl<'a> JsonParser<'a> {
    fn new(text: &'a str) -> Self {
        JsonParser { bytes: text.as_bytes(), pos: 0 }
    }

    fn skip_ws(&mut self) {
        while let Some(&b) = self.bytes.get(self.pos) {
            if b == b' ' || b == b'\t' || b == b'\n' || b == b'\r' {
                self.pos += 1;
            } else {
                break;
            }
        }
    }

    fn peek(&self) -> Option<u8> {
        self.bytes.get(self.pos).copied()
    }

    fn eat(&mut self, b: u8) -> bool {
        if self.peek() == Some(b) {
            self.pos += 1;
            true
        } else {
            false
        }
    }

    fn literal(&mut self, lit: &[u8]) -> bool {
        if self.bytes[self.pos..].starts_with(lit) {
            self.pos += lit.len();
            true
        } else {
            false
        }
    }

    /// value := ws* (object | array | string | number | 'true' | 'false' | 'null') ws*
    fn value(&mut self) -> bool {
        self.skip_ws();
        match self.peek() {
            Some(b'{') => self.object(),
            Some(b'[') => self.array(),
            Some(b'"') => self.string(),
            Some(b'-') | Some(b'0'..=b'9') => self.number(),
            Some(b't') => self.literal(b"true"),
            Some(b'f') => self.literal(b"false"),
            Some(b'n') => self.literal(b"null"),
            _ => false,
        }
    }

    /// object := '{' ws* (string ws* ':' value (ws* ',' ws* string ws* ':' value)*)? ws* '}'
    fn object(&mut self) -> bool {
        if !self.eat(b'{') {
            return false;
        }
        self.skip_ws();
        if self.eat(b'}') {
            return true;
        }
        loop {
            if !self.string() {
                return false;
            }
            self.skip_ws();
            if !self.eat(b':') {
                return false;
            }
            if !self.value() {
                return false;
            }
            self.skip_ws();
            if self.eat(b',') {
                self.skip_ws();
            } else {
                return self.eat(b'}');
            }
        }
    }

    /// array := '[' ws* (value (ws* ',' ws* value)*)? ws* ']'
    fn array(&mut self) -> bool {
        if !self.eat(b'[') {
            return false;
        }
        self.skip_ws();
        if self.eat(b']') {
            return true;
        }
        loop {
            if !self.value() {
                return false;
            }
            self.skip_ws();
            if self.eat(b',') {
                self.skip_ws();
            } else {
                return self.eat(b']');
            }
        }
    }

    /// string := '"' (escape | any byte >= 0x20)* '"'
    /// escape := '\' ("\"" | "/" | "\" | 'b' | 'f' | 'n' | 'r' | 't' | 'u' hex hex hex hex)
    fn string(&mut self) -> bool {
        if !self.eat(b'"') {
            return false;
        }
        while let Some(b) = self.peek() {
            match b {
                b'"' => {
                    self.pos += 1;
                    return true;
                }
                b'\\' => {
                    self.pos += 1;
                    let Some(esc) = self.peek() else { return false };
                    self.pos += 1;
                    match esc {
                        b'"' | b'/' | b'\\' | b'b' | b'f' | b'n' | b'r' | b't' => {}
                        b'u' => {
                            for _ in 0..4 {
                                match self.peek() {
                                    Some(h @ (b'0'..=b'9' | b'a'..=b'f' | b'A'..=b'F')) => {
                                        self.pos += 1;
                                        let _ = h;
                                    }
                                    _ => return false,
                                }
                            }
                        }
                        _ => return false,
                    }
                }
                // Raw control characters are not allowed inside strings.
                0x00..=0x1F => return false,
                _ => self.pos += 1,
            }
        }
        false // unterminated string
    }

    /// number := '-'? int frac? exp?
    /// int := '0' | [1-9][0-9]*  (no leading zeros, like JSON.parse)
    /// frac := '.' [0-9]+ ; exp := [eE] [+-]? [0-9]+
    fn number(&mut self) -> bool {
        self.eat(b'-');
        match self.peek() {
            Some(b'0') => self.pos += 1,
            Some(b'1'..=b'9') => {
                while matches!(self.peek(), Some(b'0'..=b'9')) {
                    self.pos += 1;
                }
            }
            _ => return false,
        }
        if self.peek() == Some(b'.') {
            self.pos += 1;
            let mut digits = 0;
            while matches!(self.peek(), Some(b'0'..=b'9')) {
                self.pos += 1;
                digits += 1;
            }
            if digits == 0 {
                return false;
            }
        }
        if matches!(self.peek(), Some(b'e' | b'E')) {
            self.pos += 1;
            if matches!(self.peek(), Some(b'+' | b'-')) {
                self.pos += 1;
            }
            let mut digits = 0;
            while matches!(self.peek(), Some(b'0'..=b'9')) {
                self.pos += 1;
                digits += 1;
            }
            if digits == 0 {
                return false;
            }
        }
        true
    }
}

/// Whole-text JSON gate: a document that parses as JSON is json all the way
/// down. Mirrors `isValidJson()` (`JSON.parse` in a try/catch);
/// empty/whitespace text is not.
fn is_valid_json(text: &str) -> bool {
    if text.trim().is_empty() {
        return false;
    }
    let mut p = JsonParser::new(text);
    p.value()
        && {
            p.skip_ws();
            p.pos == p.bytes.len() // reject trailing garbage
        }
}

/// Token count of `text` under auto content detection — exactly the slice of
/// `estimateTokens()` the planner consumes (`.tokens`): per non-empty line,
/// `max(1, round(utf16_len / chars_per_token))`. Framing tokens are the
/// caller's job.
fn estimate_tokens(text: &str) -> usize {
    // AUTO + whole-text JSON: json's 3 chars/token rate applies to every
    // line, not just the reported content type.
    let whole_text_json = is_valid_json(text);
    let mut tokens = 0usize;
    for line in split_lines(text) {
        if line.trim().is_empty() {
            continue;
        }
        let t = if whole_text_json { ContentType::Json } else { detect_line_type(line) };
        let line_tokens = (utf16_len(line) as f64 / t.chars_per_token()).round().max(1.0) as usize;
        tokens += line_tokens;
    }
    tokens
}

/// Sum of per-section token estimates (framing tokens are the caller's job).
/// Mirrors `inputTokenTotal()` in the TS lib.
pub fn input_token_total(sections: &[PlanSection]) -> usize {
    sections.iter().map(|s| estimate_tokens(&s.text)).sum()
}

/// Plan one section set against one model's context window. Returns `None`
/// for an unknown model id (window math is `fitsWindow`'s, never re-derived).
/// Mirrors `planWindow()` in the TS lib.
pub fn plan_window(
    sections: &[PlanSection],
    model_id: &str,
    output_reserve: usize,
    models: &[Model],
) -> Option<WindowPlan> {
    let input_tokens = input_token_total(sections);
    // Fit check inlined from fitsWindow() in src/lib/ai/models.ts.
    let m = models.iter().find(|m| m.id == model_id)?;
    let free = m.context_window as isize - input_tokens as isize;
    Some(WindowPlan {
        id: model_id.to_string(),
        input_tokens,
        context_window: m.context_window,
        free,
        fits: free >= 0,
        output_reserve_ok: free >= output_reserve as isize,
        max_output: m.max_output,
    })
}

/// Plan against several models; unknown ids are dropped from the result.
/// Mirrors `planAll()` in the TS lib.
pub fn plan_all(
    sections: &[PlanSection],
    model_ids: &[&str],
    output_reserve: usize,
    models: &[Model],
) -> Vec<WindowPlan> {
    model_ids
        .iter()
        .filter_map(|id| plan_window(sections, id, output_reserve, models))
        .collect()
}

// ---------- tests (showcase-only; the canonical suite lives in src/lib) ----------
#[cfg(test)]
mod tests {
    use super::*;

    // A 1600-char single line of 'a' is pure prose: 1600 / 4 = 400 tokens.
    fn two_sections() -> Vec<PlanSection> {
        let line_a = "a".repeat(1600);
        vec![sec("sys", &line_a), sec("docs", &line_a)]
    }

    #[test]
    fn input_totals() {
        assert_eq!(input_token_total(&two_sections()), 800);
        assert_eq!(input_token_total(&[]), 0);
        assert_eq!(input_token_total(&[sec("sys", "")]), 0);
    }

    #[test]
    fn plans_two_400_token_sections_against_beta_pro() {
        let p = plan_window(&two_sections(), "beta-pro", 0, &sample_models()).unwrap();
        assert_eq!(p.id, "beta-pro");
        assert_eq!(p.input_tokens, 800);
        assert_eq!(p.context_window, 1_000_000);
        assert_eq!(p.free, 999_200);
        assert!(p.fits);
        assert!(p.output_reserve_ok);
        assert_eq!(p.max_output, 10_000);
    }

    #[test]
    fn reserve_larger_than_free_leaves_raw_fit_true() {
        let p = plan_window(&two_sections(), "beta-pro", 1_000_000, &sample_models()).unwrap();
        assert!(p.fits);
        assert!(!p.output_reserve_ok);
    }

    #[test]
    fn reserve_exactly_equal_to_free_is_ok() {
        let p = plan_window(&two_sections(), "beta-pro", 999_200, &sample_models()).unwrap();
        assert!(p.output_reserve_ok);
    }

    #[test]
    fn smaller_window_leaves_199_200_free() {
        let p = plan_window(&two_sections(), "alpha-mini", 0, &sample_models()).unwrap();
        assert_eq!(p.context_window, 200_000);
        assert_eq!(p.free, 199_200);
        assert!(p.fits);
    }

    #[test]
    fn unknown_model_id_returns_none() {
        assert!(plan_window(&two_sections(), "ghost", 0, &sample_models()).is_none());
    }

    #[test]
    fn no_sections_full_window_free() {
        let p = plan_window(&[], "beta-pro", 0, &sample_models()).unwrap();
        assert_eq!(p.input_tokens, 0);
        assert_eq!(p.free, 1_000_000);
        assert!(p.fits);
    }

    #[test]
    fn overflow_fits_false_reserve_false() {
        let sections = [sec("big", &"z".repeat(4_400_000))];
        let p = plan_window(&sections, "beta-pro", 0, &sample_models()).unwrap();
        assert_eq!(p.input_tokens, 1_100_000);
        assert_eq!(p.free, -100_000);
        assert!(!p.fits);
        assert!(!p.output_reserve_ok);
    }

    #[test]
    fn plan_all_drops_unknown_ids_and_keeps_order() {
        let plans = plan_all(&two_sections(), &["beta-pro", "alpha-mini", "ghost"], 0, &sample_models());
        assert_eq!(plans.len(), 2);
        assert_eq!(plans[0].id, "beta-pro");
        assert_eq!(plans[1].id, "alpha-mini");
        assert_eq!(plans[1].free, 199_200);
        assert!(plan_all(&two_sections(), &[], 0, &sample_models()).is_empty());
    }
}

Also available in 13 other languages

Every CosmoDev tool ships its pure logic in TypeScript (web) and Go (CLI), with authored implementations in a dozen-plus languages — the same contract, ported. Compare all languages side by side →