Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ jobs:
strategy:
fail-fast: false
matrix:
python-version: ["3.13"]
python-version: ["3.10", "3.11", "3.12", "3.13"]

steps:
- uses: actions/checkout@v4
Expand Down
263 changes: 148 additions & 115 deletions README.md

Large diffs are not rendered by default.

Binary file added assets/context-kernel.drawio.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file added assets/context_kernel.gif
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
64 changes: 59 additions & 5 deletions context_kernel/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
"""
from __future__ import annotations

import re
import sys
import threading
import time
Expand All @@ -23,6 +24,55 @@
from .memory.storage import StorageEngine
from .pruners.shell_pruner import ShellPruner

_USD_PER_MILLION_INPUT_TOKENS = 3.0

_ANSI_OSC = re.compile(r"\x1b\][^\x07\x1b]*(?:\x07|\x1b\\)")
_ANSI_CSI = re.compile(r"\x1b\[[0-9;?]*[ -/]*[@-~]")
_ANSI_OTHER = re.compile(r"\x1b[@-Z\\-_]")
_CTRL_CHARS = re.compile(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]")


def _sanitize(text: str) -> str:
"""Strip escape sequences and control chars from stored output before echo."""
text = _ANSI_OSC.sub("", text)
text = _ANSI_CSI.sub("", text)
text = _ANSI_OTHER.sub("", text)
return _CTRL_CHARS.sub("", text)


def _print_session_summary(
storage: StorageEngine,
session_id: str,
stats: OrchestratorStats,
) -> None:
"""Print a one-glance summary of what ACK saved this session (to stderr)."""
db_stats = storage.stats(session_id)
chunks = db_stats["total_entries"]
pruned = db_stats["pruned_entries"]
raw_pruned_tokens = db_stats["tokens_saved"]
saved = stats.tokens_saved

if chunks == 0:
return

elapsed = max(1, int(time.monotonic() - stats.session_start))
pct = (saved / raw_pruned_tokens * 100) if raw_pruned_tokens else 0.0
cost = saved / 1_000_000 * _USD_PER_MILLION_INPUT_TOKENS
rate = f"{_USD_PER_MILLION_INPUT_TOKENS:g}"

lines = [
click.style("[ACK] Session summary", fg="cyan", bold=True),
f" Chunks intercepted : {chunks:>8,}",
f" Chunks pruned : {pruned:>8,}",
f" Tokens saved : {saved:>8,} (~{pct:.0f}% of pruned output)",
f" Est. cost saved : ${cost:>7.2f} (at ${rate}/M input tokens)",
f" Elapsed : {elapsed:>7}s",
]
try:
click.echo("\n" + "\n".join(lines), err=True)
except (BrokenPipeError, OSError):
pass


class StatsPanel(Static):
stats: reactive[OrchestratorStats] = reactive(OrchestratorStats())
Expand Down Expand Up @@ -209,10 +259,14 @@ def cmd_run(
)

if tui:
exit_code = AckDashboard(orchestrator=orch).run()
sys.exit(exit_code or 0)
exit_code = AckDashboard(orchestrator=orch).run() or 0
else:
sys.exit(orch.run())
exit_code = orch.run()
if not no_annotate:
_print_session_summary(storage, session.session_id, orch.stats)

storage.close()
sys.exit(exit_code)


@main.command(name="search")
Expand Down Expand Up @@ -259,7 +313,7 @@ def cmd_search(
pruned = " [pruned]" if row["was_pruned"] else ""
click.echo(click.style(f"[{ts}] session={sid_short}… type={etype}{pruned}", fg="cyan"))

content: str = row["compressed_summary"] or row["raw_content"]
content = _sanitize(row["compressed_summary"] or row["raw_content"])
preview = content[:300].strip()
if len(content) > 300:
preview += "\n …"
Expand Down Expand Up @@ -308,7 +362,7 @@ def cmd_sessions(limit: int, db: Path | None) -> None:
click.echo("─" * 90)
for row in rows:
sid = row["session_id"]
cmd = row["agent_command"][:40]
cmd = _sanitize(row["agent_command"])[:40]
ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(row["started_at"]))
ended = " ✓" if row["ended_at"] else " …"
click.echo(f"{sid} {ts} {cmd}{ended}")
Expand Down
134 changes: 96 additions & 38 deletions context_kernel/core/orchestrator.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,13 @@
import fcntl
import os
import pty
import re
import select
import signal
import struct
import sys
import termios
import threading
import time
import tty
from collections.abc import Callable
Expand All @@ -26,10 +28,10 @@

_DIM = "\033[2m"
_CYAN = "\033[36m"
_BOLD = "\033[1m"
_RESET = "\033[0m"

_SELECT_TIMEOUT = 0.08
_ANSI_ESC = re.compile(r"\x1b\[[0-9;]*[mGKHFJA-Z]")
_TRACEBACK_MARKER = "Traceback (most recent call last):"


@dataclass
Expand All @@ -38,6 +40,7 @@ class OrchestratorConfig:
buffer_flush_timeout: float = 0.15
read_chunk_bytes: int = 8192
annotate_injections: bool = True
max_buffer_bytes: int = 262144


@dataclass
Expand Down Expand Up @@ -86,7 +89,7 @@ def __init__(
self.command = command
self.session_id = session_id
self.storage = storage
self.pruners:list[BasePruner] = pruners or []
self.pruners: list[BasePruner] = pruners or []
self.config = config or OrchestratorConfig()
self.stats_callback = stats_callback
self.text_callback: Callable[[str], None] | None = None
Expand All @@ -102,9 +105,6 @@ def __init__(
def stats(self) -> OrchestratorStats:
return self._stats

def add_pruner(self, pruner: BasePruner) -> None:
self.pruners.append(pruner)

def run(self) -> int:
"""Spawn the agent, block until it exits, and return its exit code.

Expand All @@ -114,7 +114,7 @@ def run(self) -> int:
self._child_pid = child_pid

self._enter_raw_mode()
self._install_sigwinch_handler()
self._install_signal_handlers()

try:
return self._io_loop(child_pid)
Expand Down Expand Up @@ -153,27 +153,34 @@ def _spawn_in_pty(self) -> tuple[int, int]:
os.close(slave_fd)
os.close(master_fd)

os.execvp(self.command[0], self.command)
try:
os.execvp(self.command[0], self.command)
except OSError as exc:
os.write(2, f"ack: cannot run {self.command[0]!r}: {exc.strerror}\n".encode())
os._exit(127)

os.close(slave_fd)
return master_fd, child_pid

def _io_loop(self, child_pid: int) -> int:
"""Pump I/O between the user and the child until the child exits.

While nothing is buffered the loop blocks in select(); it only polls on
the flush timeout while it is still holding data to emit, so it stays at
~0% CPU when idle. stdin is dropped from the watch set once it reaches
EOF so a closed/piped input never spins the loop.
"""
assert self._master_fd is not None
master_fd = self._master_fd
stdin_fd = sys.stdin.fileno()
stdout_fd = sys.stdout.fileno()
exit_code = 0
watched = [master_fd, stdin_fd]

while True:
timeout = self.config.buffer_flush_timeout if self._buffer else None
try:
rlist, _, _ = select.select(
[master_fd, stdin_fd],
[],
[],
_SELECT_TIMEOUT,
)
rlist, _, _ = select.select(watched, [], [], timeout)
except InterruptedError:
continue
except (ValueError, OSError):
Expand Down Expand Up @@ -201,13 +208,13 @@ def _io_loop(self, child_pid: int) -> int:
os.write(master_fd, keys)
except OSError:
pass
else:
watched = [master_fd]

elapsed_since_data = time.monotonic() - self._last_data_monotonic
if (
self._buffer
and elapsed_since_data >= self.config.buffer_flush_timeout
):
self._flush_buffer(stdout_fd, force=True)
if self._buffer:
elapsed_since_data = time.monotonic() - self._last_data_monotonic
if elapsed_since_data >= self.config.buffer_flush_timeout:
self._flush_buffer(stdout_fd, force=True)

if self._buffer:
self._flush_buffer(stdout_fd, force=True)
Expand All @@ -229,9 +236,10 @@ def _accumulate(self, chunk: bytes) -> None:
def _flush_buffer(self, stdout_fd: int, *, force: bool = False) -> None:
"""Emit the buffered output, pruning it first if it qualifies.

force controls only *whether to emit now* (silence timeout / EOF),
never *whether to prune*: output below pruning_threshold_lines and
interactive prompts always pass through verbatim.
force controls only whether to emit now (silence timeout / EOF), never
whether to prune: output below pruning_threshold_lines and interactive
prompts always pass through verbatim. An unfinished traceback is held
(up to max_buffer_bytes) so it prunes as one unit across PTY reads.
"""
if not self._buffer:
return
Expand All @@ -251,6 +259,14 @@ def _flush_buffer(self, stdout_fd: int, *, force: bool = False) -> None:
self._buffer.append(raw)
return

if (
not force
and len(raw) < self.config.max_buffer_bytes
and self._is_incomplete_traceback(text)
):
self._buffer.append(raw)
return

if below_threshold or self._text_is_prompt(text):
self._emit(stdout_fd, raw)
self._persist(text, pruned=False)
Expand All @@ -260,19 +276,18 @@ def _flush_buffer(self, stdout_fd: int, *, force: bool = False) -> None:

summary: str | None = None
for pruner in self.pruners:
if pruner.matches(text):
result = pruner.compress(text)
if result is not None:
summary = result
saved = max(0, (len(text) - len(summary)) // 4)
self._stats.tokens_saved += saved
self._stats.total_pruner_hits += 1
break
result = pruner.compress(text)
if result is not None:
summary = result
saved = max(0, (len(text) - len(summary)) // 4)
self._stats.tokens_saved += saved
self._stats.total_pruner_hits += 1
break

if summary is not None:
self._persist(text, pruned=True, summary=summary)
injection = self._format_injection(summary, line_count)
injected = injection.encode("utf-8")
injected = self._terminal_newlines(injection).encode("utf-8")
self._stats.total_bytes_injected += len(injected)
self._emit(stdout_fd, injected)
_display = injection
Expand All @@ -287,6 +302,17 @@ def _flush_buffer(self, stdout_fd: int, *, force: bool = False) -> None:
if self.stats_callback is not None:
self.stats_callback(self._stats)

def _terminal_newlines(self, text: str) -> str:
"""Convert text ACK generates itself to CRLF while the terminal is raw.

Raw mode disables the terminal's NL->CRLF output mapping, so a bare
``\\n`` would leave the cursor in the same column (the "staircase"
effect). Child passthrough already carries CRLF from its own PTY.
"""
if self._saved_tty is None:
return text
return text.replace("\r\n", "\n").replace("\n", "\r\n")

def _emit(self, fd: int, data: bytes) -> None:
offset = 0
while offset < len(data):
Expand Down Expand Up @@ -321,6 +347,24 @@ def _format_injection(self, summary: str, original_lines: int) -> str:
)
return banner + body

def _is_incomplete_traceback(self, text: str) -> bool:
"""True if text holds a traceback still streaming its frames.

A finished traceback ends in a non-indented exception line; while frames
are still arriving the last non-blank line is an indented frame line, or
the header itself.
"""
if _TRACEBACK_MARKER not in text:
return False
clean = _ANSI_ESC.sub("", text)
nonblank = [ln for ln in clean.splitlines() if ln.strip()]
if not nonblank:
return False
last = nonblank[-1]
if _TRACEBACK_MARKER in last:
return True
return last[:1].isspace()

def _tail_is_prompt(self, chunk: bytes) -> bool:
tail = chunk[-200:]
return any(p in tail for p in _PROMPT_BYTES)
Expand All @@ -345,13 +389,27 @@ def _restore_terminal(self) -> None:
termios.tcsetattr(sys.stdin.fileno(), termios.TCSAFLUSH, self._saved_tty)
self._saved_tty = None

def _install_sigwinch_handler(self) -> None:
def _handler(signum: int, frame: object) -> None: # noqa: ARG001
if self._master_fd is not None and sys.stdout.isatty():
rows, cols = self._get_terminal_size()
self._set_winsize(self._master_fd, rows, cols)
def _install_signal_handlers(self) -> None:
if threading.current_thread() is not threading.main_thread():
return
signal.signal(signal.SIGWINCH, self._on_sigwinch)
signal.signal(signal.SIGTERM, self._on_terminate)
signal.signal(signal.SIGHUP, self._on_terminate)

def _on_sigwinch(self, signum: int, frame: object) -> None: # noqa: ARG002
if self._master_fd is not None and sys.stdout.isatty():
rows, cols = self._get_terminal_size()
self._set_winsize(self._master_fd, rows, cols)

signal.signal(signal.SIGWINCH, _handler)
def _on_terminate(self, signum: int, frame: object) -> None: # noqa: ARG002
"""Restore the terminal and forward the signal to the child before exit."""
self._restore_terminal()
if self._child_pid is not None:
try:
os.killpg(self._child_pid, signum)
except OSError:
pass
os._exit(128 + signum)

@staticmethod
def _get_terminal_size() -> tuple[int, int]:
Expand Down
18 changes: 0 additions & 18 deletions context_kernel/memory/pager.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,24 +93,6 @@ def map_file(self, path: Path) -> FileSymbolMap:
self._cache[cache_key] = fmap
return fmap

def page_symbol(self, path: Path, symbol_name: str) -> str | None:
"""Return the source text of one named symbol, or None if not found."""
fmap = self.map_file(path)
entry = next((s for s in fmap.symbols if s.name == symbol_name), None)
if entry is None:
return None

source_lines = path.read_text(encoding="utf-8", errors="replace").splitlines()
return "\n".join(source_lines[entry.start_line - 1 : entry.end_line])

def toc(self, path: Path) -> str:
return self.map_file(path).to_toc()

def invalidate(self, path: Path) -> None:
stale = [k for k in self._cache if k.startswith(str(path.resolve()))]
for k in stale:
del self._cache[k]

def _parse_with_tree_sitter(
self,
path: Path,
Expand Down
Loading
Loading