trl-token-reduction 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
plugin/__init__.py ADDED
@@ -0,0 +1,2 @@
1
+ """Claude Code plugin: expose the AST retrieval engine so the agent pulls only
2
+ the relevant code slices instead of grepping + dumping whole files."""
@@ -0,0 +1,16 @@
1
+ {
2
+ "mcpServers": {
3
+ "trl-retrieve": {
4
+ "command": "python",
5
+ "args": [
6
+ "-m",
7
+ "plugin.mcp_server"
8
+ ],
9
+ "cwd": "/ABSOLUTE/PATH/TO/trl-token-reduction",
10
+ "env": {
11
+ "PYTHONPATH": "/ABSOLUTE/PATH/TO/trl-token-reduction",
12
+ "TRL_REPO": "."
13
+ }
14
+ }
15
+ }
16
+ }
@@ -0,0 +1,7 @@
1
+ # trl-retrieve (code retrieval)
2
+
3
+ When you need to understand this codebase, DO NOT grep and read whole files.
4
+ Call the MCP tool **`retrieve_code`** (server `trl-retrieve`) with a natural-language
5
+ question — it returns the exact relevant source slices (function/class/method) at a
6
+ fraction of the tokens. Use **`explain_symbol`** to fetch the full source of a named
7
+ symbol. Only fall back to reading whole files if retrieval returns nothing relevant.
plugin/cli.py ADDED
@@ -0,0 +1,31 @@
1
+ """Bash-callable retrieval for Claude Code (and for humans).
2
+
3
+ python -m plugin.cli "how does auth work?" [--repo DIR] [--budget 1200] [--k 8]
4
+
5
+ Prints only the relevant code slices (file:line + source) -- a drop-in for
6
+ `grep -r` + reading whole files, at a fraction of the tokens."""
7
+ import argparse, os, sys
8
+ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
9
+ from plugin.index_store import get_index
10
+ from trl.retrieval import retrieve
11
+ from trl.util import count_tokens
12
+
13
+
14
+ def main():
15
+ ap = argparse.ArgumentParser()
16
+ ap.add_argument("query")
17
+ ap.add_argument("--repo", default=None) # None -> resolve $TRL_REPO/$CLAUDE_PROJECT_DIR/git-root/cwd
18
+ ap.add_argument("--budget", type=int, default=1200)
19
+ ap.add_argument("--k", type=int, default=8)
20
+ a = ap.parse_args()
21
+ idx = get_index(a.repo) # index_store resolves + caches at <repo>/.trl/index.json
22
+ r = retrieve(idx, a.query, token_budget=a.budget, k=a.k)
23
+ if not r["symbols"]:
24
+ print("(no relevant symbols found)"); return
25
+ hdr = (f"# {len(r['symbols'])} slices, {r['tokens']} tokens "
26
+ f"(indexed {len(idx['files'])} files, {len(idx['symbols'])} symbols)")
27
+ print(hdr); print(r["context"])
28
+
29
+
30
+ if __name__ == "__main__":
31
+ main()
@@ -0,0 +1,15 @@
1
+ # Codex retrieval plugin — merge this into ~/.codex/config.toml (global) or
2
+ # .codex/config.toml (project-scoped; trusted projects only).
3
+ #
4
+ # Same STDIO MCP server as the Claude Code plugin — the protocol is client-neutral,
5
+ # so nothing about the server changes. Edit the two paths if the project moves.
6
+ # Set TRL_REPO to the repo you want Codex to retrieve from ("." = the cwd Codex runs in).
7
+
8
+ [mcp_servers.trl-retrieve]
9
+ command = "python"
10
+ args = ["-m", "plugin.mcp_server"]
11
+ env = { PYTHONPATH = "/ABSOLUTE/PATH/TO/trl-token-reduction", TRL_REPO = "." }
12
+
13
+ # Optional — route Codex's API-key calls through the token-reduction proxy too
14
+ # (applies to API-key mode, not ChatGPT-subscription mode):
15
+ # openai_base_url = "http://localhost:8899/v1"
plugin/index_store.py ADDED
@@ -0,0 +1,71 @@
1
+ """Shared: resolve the user's project at RUNTIME and build-or-load a persistent,
2
+ incremental AST index for it.
3
+
4
+ Why runtime resolution: a Claude Code plugin MCP server's cwd is the PLUGIN root
5
+ (not the user's project), and ${...} interpolation inside a plugin's .mcp.json is
6
+ unreliable. So we never trust the manifest to pass the path -- we resolve it here.
7
+ """
8
+ import os, sys
9
+ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
10
+ from trl.retrieval import build_index, save_index, load_index
11
+
12
+ INDEX_DIR = ".trl"
13
+ INDEX_NAME = "index.json"
14
+
15
+
16
+ def _git_root(start):
17
+ """Nearest ancestor of `start` containing a .git dir, else None."""
18
+ d = os.path.abspath(start)
19
+ while True:
20
+ if os.path.isdir(os.path.join(d, ".git")):
21
+ return d
22
+ parent = os.path.dirname(d)
23
+ if parent == d:
24
+ return None
25
+ d = parent
26
+
27
+
28
+ def _resolve_repo(repo=None):
29
+ """Resolve the target project dir at RUNTIME. Order:
30
+ explicit arg -> $TRL_REPO -> $CLAUDE_PROJECT_DIR -> nearest .git above cwd -> cwd.
31
+ Always returns a path."""
32
+ if repo:
33
+ return os.path.abspath(repo)
34
+ for env in ("TRL_REPO", "CLAUDE_PROJECT_DIR"):
35
+ v = os.environ.get(env)
36
+ if v:
37
+ return os.path.abspath(v)
38
+ g = _git_root(os.getcwd())
39
+ return g if g else os.getcwd()
40
+
41
+
42
+ def _has_explicit_repo(repo=None):
43
+ """True only for a RELIABLE project signal (explicit arg or env). git-root/cwd are
44
+ NOT reliable for a plugin MCP server, whose cwd is the plugin root -- callers that
45
+ run there must gate on this to avoid indexing the wrong tree."""
46
+ return bool(repo or os.environ.get("TRL_REPO") or os.environ.get("CLAUDE_PROJECT_DIR"))
47
+
48
+
49
+ def index_path(repo):
50
+ return os.path.join(repo, INDEX_DIR, INDEX_NAME)
51
+
52
+
53
+ def get_index(repo=None):
54
+ """Up-to-date index for the resolved repo, cached at <repo>/.trl/index.json so
55
+ `/trl-index` and `retrieve_code` agree regardless of cwd. Incremental: only
56
+ changed files are re-parsed."""
57
+ repo = _resolve_repo(repo)
58
+ cache = index_path(repo)
59
+ prev = None
60
+ if os.path.exists(cache):
61
+ try:
62
+ prev = load_index(cache)
63
+ except Exception:
64
+ prev = None
65
+ idx = build_index(repo, prev=prev)
66
+ try:
67
+ os.makedirs(os.path.dirname(cache), exist_ok=True)
68
+ save_index(idx, cache)
69
+ except Exception:
70
+ pass
71
+ return idx
plugin/mcp_server.py ADDED
@@ -0,0 +1,143 @@
1
+ """MCP server exposing the retrieval engine to Claude Code (or any MCP client).
2
+
3
+ The target project is resolved AT RUNTIME (see plugin.index_store._resolve_repo):
4
+ $TRL_REPO -> $CLAUDE_PROJECT_DIR -> nearest .git above cwd -> cwd. We do NOT rely on
5
+ manifest ${...} interpolation. Because an MCP server's cwd is the PLUGIN root, if no
6
+ reliable signal is present the tools return a "run /trl-index or set TRL_REPO" hint
7
+ instead of silently indexing the plugin's own tree. Both tools accept repo= override.
8
+
9
+ Tools:
10
+ retrieve_code(query, repo=?) -> relevant slices instead of whole files
11
+ explain_symbol(name, repo=?) -> exact source of a named function/class/method
12
+ """
13
+ import os, sys
14
+ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
15
+ from plugin.index_store import get_index, _resolve_repo, _has_explicit_repo
16
+ from trl.retrieval import retrieve
17
+ from plugin.retrieval_policy import (RetrievalGuard, DEFAULT_BUDGET, DEFAULT_K,
18
+ ADAPTIVE_CEIL, IDLE_RESET_S)
19
+
20
+ try:
21
+ from mcp.server.fastmcp import FastMCP
22
+ except Exception:
23
+ FastMCP = None
24
+
25
+ _PLUGIN_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
26
+ _UNRESOLVED = ("TRL: couldn't resolve your project -- this MCP server runs from the "
27
+ "plugin directory. Run /trl-index in your project, set TRL_REPO, or "
28
+ "call this tool with repo='/path/to/your/project'.")
29
+
30
+
31
+ def _target(repo=None):
32
+ """Resolved project path, or None if we can't CONFIDENTLY resolve it (no explicit
33
+ arg/env and resolution falls back to the plugin root)."""
34
+ if _has_explicit_repo(repo):
35
+ return _resolve_repo(repo)
36
+ resolved = _resolve_repo(repo)
37
+ if os.path.abspath(resolved) == _PLUGIN_ROOT:
38
+ return None
39
+ return resolved
40
+
41
+
42
+ def _log_savings(kind, query, target, slice_tokens, symbols):
43
+ """Best-effort adoption + savings log (opt-in via env TRL_SAVINGS_LOG). Records the
44
+ slices actually returned vs the whole-file counterfactual (the files those slices came
45
+ from), so cumulative savings sum across sessions. NEVER raises -- must not break retrieval."""
46
+ path = os.environ.get("TRL_SAVINGS_LOG")
47
+ if path and path.strip().lower() == "off":
48
+ return
49
+ if not path:
50
+ # Deterministic default: log into the repo we are actually retrieving from, so
51
+ # tracking never depends on launcher env / cwd (fixes the split/undercounted logs).
52
+ try:
53
+ if target and os.path.isdir(target):
54
+ path = os.path.join(target, ".trl", "savings.jsonl")
55
+ else:
56
+ return
57
+ except Exception:
58
+ return
59
+ try:
60
+ import json, time
61
+ from trl.util import count_tokens
62
+ files = {s.file for s in symbols}
63
+ whole = 0
64
+ for f in files:
65
+ try:
66
+ whole += count_tokens(open(f, encoding="utf-8", errors="ignore").read())
67
+ except Exception:
68
+ pass
69
+ rec = {"ts": time.time(), "tool": kind, "repo": target, "query": (query or "")[:120],
70
+ "slice_tokens": int(slice_tokens), "wholefile_tokens": whole,
71
+ "saved": max(0, whole - int(slice_tokens)), "n_slices": len(symbols),
72
+ "n_files": len(files)}
73
+ os.makedirs(os.path.dirname(os.path.abspath(path)) or ".", exist_ok=True)
74
+ with open(path, "a", encoding="utf-8") as fh:
75
+ fh.write(json.dumps(rec) + "\n")
76
+ except Exception:
77
+ pass
78
+
79
+
80
+ if FastMCP is not None:
81
+ mcp = FastMCP("trl-retrieve")
82
+ _GUARD = RetrievalGuard(idle_reset_s=IDLE_RESET_S) # over-retrieval guard, resets on idle
83
+
84
+ @mcp.tool()
85
+ def retrieve_code(query: str = "", budget: int = DEFAULT_BUDGET, k: int = DEFAULT_K,
86
+ repo: str = "", token_budget: int = 0, limit: int = 0, path: str = "",
87
+ project: str = "") -> str:
88
+ """Return the most relevant code slices for a question, instead of reading whole
89
+ files. Make ONE good, specific query then ANSWER from the result -- do not re-run
90
+ variations of the same question. Args: query (REQUIRED), k (max slices, default 5),
91
+ budget (token budget, default 700), repo (optional project path). Aliases
92
+ token_budget/limit/path/project accepted."""
93
+ query = (query or "").strip()
94
+ if not query:
95
+ return ("retrieve_code needs a 'query' string, e.g. "
96
+ "retrieve_code(query=\"how does the stash update\").")
97
+ # Over-retrieval guard: nudge/stop a runaway loop before it re-bills huge history.
98
+ blocked, prefix = _GUARD.check(query)
99
+ if blocked:
100
+ return prefix
101
+ # Clamp to lean caps so a large model-supplied budget/k can't defeat the fix.
102
+ budget = min(token_budget or budget or DEFAULT_BUDGET, DEFAULT_BUDGET)
103
+ k = min(limit or k or DEFAULT_K, DEFAULT_K)
104
+ target = _target(repo or path or project or None)
105
+ if target is None:
106
+ return _UNRESOLVED
107
+ r = retrieve(get_index(target), query, token_budget=budget, k=k, adaptive_ceil=ADAPTIVE_CEIL)
108
+ _log_savings("retrieve_code", query, target, r["tokens"], r["symbols"])
109
+ return prefix + (r["context"] or "(no relevant symbols found)")
110
+
111
+ @mcp.tool()
112
+ def explain_symbol(name: str = "", repo: str = "", path: str = "", project: str = "") -> str:
113
+ """Return the exact source of a function/class/method by name. Use sparingly, after
114
+ retrieve_code, for one specific symbol. Args: name (REQUIRED), repo (optional project
115
+ path). Aliases path/project accepted."""
116
+ name = (name or "").strip()
117
+ if not name:
118
+ return "explain_symbol needs a 'name', e.g. explain_symbol(name=\"StashService\")."
119
+ blocked, prefix = _GUARD.check("symbol:" + name)
120
+ if blocked:
121
+ return prefix
122
+ target = _target(repo or path or project or None)
123
+ if target is None:
124
+ return _UNRESOLVED
125
+ idx = get_index(target)
126
+ hits = [s for s in idx["symbols"] if s.name == name]
127
+ if not hits:
128
+ return f"(no symbol named {name})"
129
+ picked = hits[:3]
130
+ from trl.util import count_tokens
131
+ _log_savings("explain_symbol", name, target, sum(count_tokens(s.source) for s in picked), picked)
132
+ return prefix + "\n\n".join(f"# {s.file}:{s.start_line}-{s.end_line} ({s.kind})\n{s.source}"
133
+ for s in picked)
134
+
135
+ def main():
136
+ mcp.run()
137
+ else:
138
+ def main():
139
+ print("install the MCP SDK: pip install mcp", file=sys.stderr); sys.exit(1)
140
+
141
+
142
+ if __name__ == "__main__":
143
+ main()
@@ -0,0 +1,109 @@
1
+ """Single source of truth for the retrieval POLICY, shared by the shipping MCP
2
+ (plugin/mcp_server.py) and the paired billed-token harness
3
+ (validate/paired_billed_bench.py) so the benchmark measures exactly what ships.
4
+
5
+ Background: the trl-retrieve MCP is quality-neutral and cheaper per retrieval, BUT
6
+ had an over-retrieval failure mode -- an agent could loop retrieve_code 25+ times and
7
+ the re-billed slice history ballooned to 280-380k input tokens, flipping the total
8
+ cost above a plain grep+read agent. This module fixes that with two levers, applied
9
+ identically in the tool and in the harness:
10
+
11
+ Lever 1 -- leaner slices: smaller k + token budget + a lower adaptive ceiling, so
12
+ each result is small and per-call history growth is bounded. (retrieve() defaults to
13
+ an adaptive ceiling of 6000 tokens/call; we cap the retrieval path far below that.)
14
+
15
+ Lever 2 -- a per-session repeat-call guard: after a few calls, nudge the agent to
16
+ ANSWER instead of retrieving more; after a hard cap, stop returning slices entirely
17
+ so the loop cannot run away. Exact-duplicate queries are short-circuited immediately.
18
+
19
+ Positioning guardrail (do NOT regress): this makes the tool cheaper AND quality-neutral
20
+ vs a shell agent post-fix. It does NOT restore any blanket "93% billed savings" claim;
21
+ 93% was content reduction vs whole-file reads, not billed savings in an agent loop.
22
+ """
23
+
24
+ # --- Lever 1: leaner slices (was budget=1200, k=8, adaptive ceiling=6000) ---
25
+ DEFAULT_BUDGET = 700 # token budget per retrieve_code call
26
+ DEFAULT_K = 5 # max slices per call
27
+ ADAPTIVE_CEIL = 1500 # hard cap on per-call returned context (overrides retrieve()'s 6000)
28
+
29
+ # --- Lever 2: repeat-call guard, per session / per agent loop ---
30
+ SOFT_NUDGE_AFTER = 3 # calls 4.. get a "you likely have enough, answer now" prefix
31
+ HARD_STOP_AFTER = 6 # calls 7.. return only a stop-and-answer nudge (no new slices)
32
+ IDLE_RESET_S = 90 # (shipping MCP only) auto-reset the counter after this much idle,
33
+ # so a NEW user turn starts fresh while an intra-turn runaway
34
+ # loop (back-to-back calls) still hits the cap.
35
+
36
+ RETRIEVE_DESC = (
37
+ "Return the most relevant code slices for a question, instead of reading whole "
38
+ "files. Make ONE good, specific query and then ANSWER from the result -- do NOT "
39
+ "re-run variations of the same question; you rarely need more than a couple of "
40
+ "retrievals. Args: query (the question, REQUIRED string), k (max slices, default 5), "
41
+ "budget (token budget, default 700).")
42
+ EXPLAIN_DESC = (
43
+ "Return the exact source of a function/class/method/table by name. Use sparingly, "
44
+ "after retrieve_code, when you need one specific named symbol in full. Args: name (REQUIRED).")
45
+
46
+ _SOFT_MSG = ("NOTE: you have already retrieved code {n} times this turn -- you very likely "
47
+ "have enough context now. Prefer to ANSWER the question. Only retrieve again if "
48
+ "a specific, named thing is still missing.\n\n")
49
+ _HARD_MSG = ("You have retrieved code {n} times this turn -- that is enough context. Stop "
50
+ "retrieving and ANSWER the user's question now from what you already have. "
51
+ "(Further retrieval is paused for this turn to prevent runaway token cost.)")
52
+ _DUP_MSG = ("You already ran this exact query this turn and the result has not changed. "
53
+ "Answer now from the context you already retrieved instead of re-retrieving.")
54
+
55
+
56
+ def _norm(query):
57
+ return " ".join((query or "").lower().split())
58
+
59
+
60
+ class RetrievalGuard:
61
+ """Per-session counter shared by both retrieval tools. One instance lives for the
62
+ life of ONE agent loop: the MCP server process (shipping), or one run_session in the
63
+ harness (which builds a fresh guard per session).
64
+
65
+ Usage: call check(query) BEFORE retrieving.
66
+ returns (blocked, prefix)
67
+ blocked=True -> return `prefix` as the WHOLE tool result; do NOT retrieve
68
+ (no slices are added to history -- this is what caps the loop).
69
+ blocked=False -> prepend `prefix` (possibly "") to the real slices.
70
+
71
+ idle_reset_s: if set (shipping MCP), the counter auto-resets when more than that many
72
+ seconds elapse between calls, so distinct user turns don't share one budget while a
73
+ back-to-back runaway loop still trips the cap. Left None in the harness for determinism.
74
+ """
75
+
76
+ def __init__(self, idle_reset_s=None, clock=None):
77
+ self.idle_reset_s = idle_reset_s
78
+ self._clock = clock
79
+ self._last = None
80
+ self.n = 0
81
+ self._seen = set()
82
+
83
+ def _now(self):
84
+ if self._clock is not None:
85
+ return self._clock()
86
+ import time
87
+ return time.monotonic()
88
+
89
+ def reset(self):
90
+ self.n = 0
91
+ self._seen = set()
92
+
93
+ def check(self, query):
94
+ if self.idle_reset_s is not None:
95
+ now = self._now()
96
+ if self._last is not None and (now - self._last) > self.idle_reset_s:
97
+ self.reset()
98
+ self._last = now
99
+ norm = _norm(query)
100
+ if norm and norm in self._seen:
101
+ return True, _DUP_MSG # exact repeat: don't count, don't retrieve
102
+ if self.n >= HARD_STOP_AFTER:
103
+ return True, _HARD_MSG.format(n=self.n) # cap reached: stop, no slices
104
+ self.n += 1
105
+ if norm:
106
+ self._seen.add(norm)
107
+ if self.n > SOFT_NUDGE_AFTER:
108
+ return False, _SOFT_MSG.format(n=self.n) # nudge, but still retrieve
109
+ return False, ""
proxy/__init__.py ADDED
@@ -0,0 +1,5 @@
1
+ """v0 product skin: an OpenAI-compatible drop-in proxy. Point your client's
2
+ base_url at this server and it applies the token-reduction levers to every
3
+ request before forwarding upstream -- zero code change for the caller."""
4
+ from .transform import transform_chat_request
5
+ __all__ = ["transform_chat_request", "transform_anthropic_request"]
@@ -0,0 +1,49 @@
1
+ """Local /compress endpoint for the composer-compression browser extension.
2
+
3
+ No upstream call. Takes the user's pasted bulk context, runs the SAME
4
+ compression + fact-guard (or the text retriever) the rest of the engine uses,
5
+ and returns the shorter text plus token counts and the numeric facts the guard
6
+ guarantees survived -- so the extension can show a trustworthy pre-send preview.
7
+
8
+ Isolated + pure so it unit-tests offline (like transform.py)."""
9
+ import re
10
+
11
+ from trl.message import Message, TOOL_RESULT
12
+ from trl.compress import compress_request
13
+ from trl.util import count_tokens
14
+
15
+ _NUM = re.compile(r"-?\d[\d,]*")
16
+
17
+
18
+ def handle_compress(req: dict, engine) -> dict:
19
+ text = (req.get("text") or "").strip()
20
+ if not text:
21
+ return {"error": "no text"}
22
+ mode = req.get("mode", "compress")
23
+ question = (req.get("question") or "").strip()
24
+ budget = int(req.get("budget", 1200))
25
+ before = count_tokens(text)
26
+
27
+ if mode == "retrieve":
28
+ # Doc-slice mode: keep only the passages relevant to `question`.
29
+ from trl.retrieval import build_text_index, retrieve_text
30
+ idx = build_text_index({"pasted": text})
31
+ r = retrieve_text(idx, question or text, token_budget=budget, rerank=False)
32
+ out = r.get("context", "") or text
33
+ else:
34
+ # Compression mode: summarize the tail, then the deterministic fact-guard
35
+ # re-injects any dropped number. Short text (<200 chars) passes through.
36
+ msg = Message("tool", TOOL_RESULT, text)
37
+ new_msgs, _ = compress_request([msg], engine.mode, engine.local)
38
+ out = "\n".join(m.content for m in new_msgs).strip() or text
39
+
40
+ after = count_tokens(out)
41
+ preserved = sorted({n.replace(",", "") for n in _NUM.findall(text)}, key=len)
42
+ return {
43
+ "compressed": out,
44
+ "tokens_before": before,
45
+ "tokens_after": after,
46
+ "saved_pct": round(100 * (1 - after / before), 1) if before else 0.0,
47
+ "preserved_facts": preserved,
48
+ "mode": mode,
49
+ }
proxy/server.py ADDED
@@ -0,0 +1,179 @@
1
+ """Drop-in OpenAI-compatible proxy. Point your client's base_url here; every
2
+ /v1/chat/completions request gets the token-reduction levers applied, then is
3
+ forwarded upstream with YOUR Authorization header (the proxy never stores keys).
4
+
5
+ Run: python -m proxy.server # listens on :8899, forwards to OpenAI
6
+ Then: client base_url = http://localhost:8899/v1 (key unchanged)
7
+
8
+ Stdlib only (http.server + urllib) so it has zero extra deps. Adds response
9
+ header X-TRL-Tokens-Saved so you can see the reduction per request."""
10
+ from __future__ import annotations
11
+
12
+ import json
13
+ import os
14
+ import sys
15
+ import urllib.request
16
+ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
17
+
18
+ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
19
+ from trl import Engine
20
+ from trl.util import load_config
21
+ from proxy.transform import transform_chat_request, transform_anthropic_request
22
+ from proxy.compress_endpoint import handle_compress
23
+
24
+ _DEFAULT_CFG = {
25
+ "arms": {"treatment": {"native_prompt_cache": True, "compress_history": True,
26
+ "compress_tool_outputs": True, "compression_mode": "safe"}},
27
+ "local_model": {"provider": "none"},
28
+ "retrieval": {"enabled": True, "token_budget": 800, "k": 8, "expand_call_graph": True},
29
+ }
30
+
31
+
32
+ def _load_cfg():
33
+ # Prefer an explicit path, then ./config.yaml, then a repo-relative copy;
34
+ # fall back to a sane built-in default so the pip-installed CLI runs anywhere.
35
+ import os as _os
36
+ for cand in (_os.environ.get("TRL_CONFIG"), "config.yaml",
37
+ _os.path.join(_os.path.dirname(__file__), "..", "config.yaml")):
38
+ if cand and _os.path.exists(cand):
39
+ try:
40
+ return load_config(cand)
41
+ except Exception:
42
+ pass
43
+ return _DEFAULT_CFG
44
+
45
+
46
+ _CFG = _load_cfg()
47
+ _ENGINE = Engine(_CFG)
48
+ _UPSTREAM = os.environ.get("TRL_UPSTREAM", "https://api.openai.com")
49
+
50
+
51
+ _MAX_BODY = 10 * 1024 * 1024 # 10 MB request-body cap (reject before read)
52
+ _ALLOWED_ORIGINS = {"https://claude.ai"}
53
+ _CORS_BASE = {
54
+ "Access-Control-Allow-Methods": "POST, OPTIONS",
55
+ "Access-Control-Allow-Headers": "Content-Type",
56
+ "Access-Control-Allow-Private-Network": "true",
57
+ }
58
+
59
+
60
+ class Handler(BaseHTTPRequestHandler):
61
+ def _cors(self):
62
+ # Scope CORS: the browser extension's origin is chrome-extension://<id>
63
+ # (random per install, so allow the scheme) plus claude.ai. Arbitrary
64
+ # websites get NO Access-Control-Allow-Origin, so a browser refuses to let
65
+ # them read responses from the local proxy. Non-browser SDK clients ignore
66
+ # CORS entirely, so this doesn't affect the base_url-swap proxy usage.
67
+ h = dict(_CORS_BASE)
68
+ origin = self.headers.get("Origin", "")
69
+ if origin.startswith("chrome-extension://") or origin in _ALLOWED_ORIGINS:
70
+ h["Access-Control-Allow-Origin"] = origin
71
+ h["Vary"] = "Origin"
72
+ return h
73
+
74
+ def do_OPTIONS(self):
75
+ self._send(204, b"", self._cors())
76
+
77
+ def do_GET(self):
78
+ # Lightweight liveness probe so clients (e.g. the browser extension) can
79
+ # tell "engine running" from "engine down" and show a useful message.
80
+ if self.path.rstrip("/").endswith("/health"):
81
+ return self._send(200, b'{"status":"ok","service":"trl-proxy"}',
82
+ {"Content-Type": "application/json", **self._cors()})
83
+ return self._send(404, b'{"error":"not found"}',
84
+ {"Content-Type": "application/json", **self._cors()})
85
+
86
+ def _send(self, code, body: bytes, headers=None):
87
+ self.send_response(code)
88
+ for k, v in (headers or {}).items():
89
+ self.send_header(k, v)
90
+ self.end_headers()
91
+ self.wfile.write(body)
92
+
93
+ def do_POST(self):
94
+ try:
95
+ n = int(self.headers.get("Content-Length") or 0)
96
+ except (TypeError, ValueError):
97
+ return self._send(411, b'{"error":"missing or invalid Content-Length"}',
98
+ {"Content-Type": "application/json", **self._cors()})
99
+ if n > _MAX_BODY: # reject oversized bodies BEFORE read
100
+ return self._send(413, b'{"error":"request body too large"}',
101
+ {"Content-Type": "application/json", **self._cors()})
102
+ raw = self.rfile.read(n)
103
+ path = self.path.rstrip("/")
104
+ if path.endswith("/compress"):
105
+ try:
106
+ result = handle_compress(json.loads(raw), _ENGINE)
107
+ except Exception as e:
108
+ return self._send(400, json.dumps({"error": str(e)}).encode(), self._cors())
109
+ return self._send(200, json.dumps(result).encode(),
110
+ {"Content-Type": "application/json", **self._cors()})
111
+ is_anthropic = path.endswith("/messages")
112
+ is_openai = path.endswith("/chat/completions")
113
+ if not (is_anthropic or is_openai):
114
+ return self._send(404, b'{"error":"only /v1/chat/completions or /v1/messages"}',
115
+ {"Content-Type": "application/json"})
116
+ try:
117
+ req = json.loads(raw)
118
+ if is_anthropic:
119
+ new_req, meta = transform_anthropic_request(req, _ENGINE)
120
+ else:
121
+ new_req, meta = transform_chat_request(req, _ENGINE)
122
+ except Exception as e:
123
+ return self._send(400, json.dumps({"error": str(e)}).encode())
124
+
125
+ fwd_headers = {"Content-Type": "application/json"}
126
+ if self.headers.get("Authorization"):
127
+ fwd_headers["Authorization"] = self.headers.get("Authorization")
128
+ for h in ("x-api-key", "anthropic-version", "anthropic-beta", "OpenAI-Organization"):
129
+ if self.headers.get(h):
130
+ fwd_headers[h] = self.headers.get(h)
131
+ up = urllib.request.Request(_UPSTREAM + self.path,
132
+ data=json.dumps(new_req).encode(), headers=fwd_headers)
133
+ saved_hdrs = {"X-TRL-Tokens-Saved": str(meta["tokens_saved"]),
134
+ "X-TRL-Tokens-Before": str(meta["tokens_before"]),
135
+ "X-TRL-Tokens-After": str(meta["tokens_after"])}
136
+ streaming = False
137
+ try:
138
+ with urllib.request.urlopen(up, timeout=300) as r:
139
+ if new_req.get("stream"):
140
+ # stream the upstream SSE response through, unchanged
141
+ streaming = True
142
+ self.send_response(r.status)
143
+ self.send_header("Content-Type",
144
+ r.headers.get("Content-Type", "text/event-stream"))
145
+ for k, v in saved_hdrs.items():
146
+ self.send_header(k, v)
147
+ self.end_headers()
148
+ while True:
149
+ chunk = r.read(2048)
150
+ if not chunk:
151
+ break
152
+ self.wfile.write(chunk); self.wfile.flush()
153
+ else:
154
+ body = r.read()
155
+ self._send(r.status, body, {"Content-Type": "application/json", **saved_hdrs})
156
+ except urllib.error.HTTPError as e:
157
+ self._send(e.code, e.read())
158
+ except Exception as e:
159
+ if streaming:
160
+ # headers already sent mid-stream; a second status line would
161
+ # corrupt the response -- just drop the connection.
162
+ self.close_connection = True
163
+ return
164
+ self._send(502, json.dumps({"error": f"upstream: {e}"}).encode())
165
+
166
+ def log_message(self, *a): # quiet
167
+ pass
168
+
169
+
170
+ def main(port=None):
171
+ if port is None: # honor TRL_PORT for the pip-installed `trl-proxy` script too
172
+ port = int(os.environ.get("TRL_PORT", "8899"))
173
+ print(f"TRL proxy on http://localhost:{port} -> {_UPSTREAM}")
174
+ print("point your client base_url at http://localhost:%d/v1 (key unchanged)" % port)
175
+ ThreadingHTTPServer(("127.0.0.1", port), Handler).serve_forever()
176
+
177
+
178
+ if __name__ == "__main__":
179
+ main()