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 +2 -0
- plugin/claude-code/.mcp.json +16 -0
- plugin/claude-code/SKILL.md +7 -0
- plugin/cli.py +31 -0
- plugin/codex/config.toml +15 -0
- plugin/index_store.py +71 -0
- plugin/mcp_server.py +143 -0
- plugin/retrieval_policy.py +109 -0
- proxy/__init__.py +5 -0
- proxy/compress_endpoint.py +49 -0
- proxy/server.py +179 -0
- proxy/test_live.py +59 -0
- proxy/transform.py +220 -0
- trl/__init__.py +12 -0
- trl/cache.py +37 -0
- trl/cascade.py +39 -0
- trl/compress.py +144 -0
- trl/engine.py +82 -0
- trl/local_model.py +168 -0
- trl/message.py +33 -0
- trl/retrieval/__init__.py +10 -0
- trl/retrieval/ast_index.py +353 -0
- trl/retrieval/embed.py +22 -0
- trl/retrieval/llm_rerank.py +39 -0
- trl/retrieval/pdf.py +33 -0
- trl/retrieval/retrieve.py +313 -0
- trl/retrieval/text_index.py +113 -0
- trl/util.py +87 -0
- trl_token_reduction-0.1.0.dist-info/METADATA +391 -0
- trl_token_reduction-0.1.0.dist-info/RECORD +35 -0
- trl_token_reduction-0.1.0.dist-info/WHEEL +5 -0
- trl_token_reduction-0.1.0.dist-info/entry_points.txt +4 -0
- trl_token_reduction-0.1.0.dist-info/licenses/LICENSE +201 -0
- trl_token_reduction-0.1.0.dist-info/licenses/NOTICE +4 -0
- trl_token_reduction-0.1.0.dist-info/top_level.txt +3 -0
plugin/__init__.py
ADDED
|
@@ -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()
|
plugin/codex/config.toml
ADDED
|
@@ -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()
|