#!/usr/bin/env python3
"""OptMem: a permanent, append-only memory for AI agents.

  memo wake [part]         print who you are
  memo note "..."          record a memory
  memo sleep [id "..."]    compress
  memo recall <regex>      search the raw log
  memo forget <id>         drop a wrong summary so it is compressed again
  memo import <file>       bulk-append historical memories (bootstrap only)

Everything lives in $MEMORY_DIR as two append-only text files. See README.md.
"""

import datetime
import fcntl
import os
import re
import sys

sys.path.insert(0, os.path.dirname(os.path.realpath(__file__)))
from blocks import complete, cover  # noqa: E402

ENTRY_CHARS = 280
WAKE_LINES = 320
RAW_MAX = 16  # blocks up to this many memories compress from the raw log

# Every agent harness silently truncates a command that prints too much --
# Codex at 10 KiB or 256 lines, Claude Code at 30k chars, pi at 50 KB -- and
# each drops a different part. So the memory is handed over in parts that fit
# the strictest of them. These are transport limits, not memory limits.
PART_CHARS = 8000
PART_LINES = 200


# Records are FIXED WIDTH, so a memory or a block is found by seeking to its
# offset -- no scanning, no index file to keep in sync. Position IS identity:
# memory i lives at i*LOG_REC of LOG.txt, and block [k*s,(k+1)*s) lives at
# k*TREE_REC of TREE/<s>. Padding costs ~2x on disk and buys O(1) everywhere.
LOG_REC = 320
TREE_REC = 288


# ---------------------------------------------------------------- store

def store():
    d = os.environ.get("MEMORY_DIR")
    if not d:
        die("MEMORY_DIR is not set. It must name this machine's memory "
            "directory, e.g. export MEMORY_DIR=~/memory")
    d = os.path.expanduser(d)
    os.makedirs(os.path.join(d, "TREE"), exist_ok=True)
    p = os.path.join(d, "LOG.txt")
    if not os.path.exists(p):
        open(p, "a").close()
    return d


def config(d):
    global ENTRY_CHARS, WAKE_LINES, PART_CHARS, PART_LINES
    p = os.path.join(d, "config")
    if not os.path.exists(p):
        with open(p, "w") as f:
            f.write("ENTRY_CHARS=%d\nWAKE_LINES=%d\nPART_CHARS=%d\nPART_LINES=%d\n"
                    % (ENTRY_CHARS, WAKE_LINES, PART_CHARS, PART_LINES))
        return
    for line in open(p):
        line = line.split("#")[0].strip()
        if "=" not in line:
            continue
        k, v = (s.strip() for s in line.split("=", 1))
        if k == "ENTRY_CHARS":
            ENTRY_CHARS = int(v)
        elif k == "WAKE_LINES":
            WAKE_LINES = int(v)
        elif k == "PART_CHARS":
            PART_CHARS = int(v)
        elif k == "PART_LINES":
            PART_LINES = int(v)


def log_path(d):
    return os.path.join(d, "LOG.txt")


def tree_path(d, size):
    return os.path.join(d, "TREE", str(size))


def count(path, rec):
    try:
        return os.path.getsize(path) // rec
    except OSError:
        return 0


def log_len(d):
    return count(log_path(d), LOG_REC)


def log_get(d, i):
    """(id, date, text) of memory i, in one seek."""
    with open(log_path(d), "rb") as f:
        f.seek(i * LOG_REC)
        line = f.read(LOG_REC).decode().rstrip()
    head, _, rest = line.partition(" ")
    date, _, text = rest.partition(" ")
    return int(head[1:]), date, text


def log_slice(d, lo, hi):
    """Memories [lo,hi) in one read."""
    with open(log_path(d), "rb") as f:
        f.seek(lo * LOG_REC)
        buf = f.read((hi - lo) * LOG_REC).decode()
    out = []
    for i in range(hi - lo):
        line = buf[i * LOG_REC:(i + 1) * LOG_REC].rstrip()
        head, _, rest = line.partition(" ")
        date, _, text = rest.partition(" ")
        out.append((int(head[1:]), date, text))
    return out


def tree_get(d, lo, hi):
    """The summary of block [lo,hi), in one seek. None if not built yet."""
    size = hi - lo
    with open(tree_path(d, size), "rb") as f:
        f.seek((lo // size) * TREE_REC)
        rec = f.read(TREE_REC)
    return rec.decode().rstrip() or None


def pad(text, rec):
    b = text.encode()
    if len(b) > rec - 1:
        die("REJECTED: %d bytes, over the %d-byte record." % (len(b), rec - 1))
    return b + b" " * (rec - 1 - len(b)) + b"\n"


def locked(d):
    lock = open(os.path.join(d, ".lock"), "w")
    fcntl.flock(lock, fcntl.LOCK_EX)
    return lock


def log_append(d, entries):
    """Append memories. The only way LOG.txt ever changes."""
    lock = locked(d)
    try:
        with open(log_path(d), "ab") as f:
            for e in entries:
                f.write(pad(e, LOG_REC))
            f.flush()
            os.fsync(f.fileno())
    finally:
        lock.close()


def tree_put(d, lo, hi, text):
    """Write block [lo,hi). Blocks are built in order, so this only ever
    appends one record to one level file."""
    size = hi - lo
    lock = locked(d)
    try:
        p = tree_path(d, size)
        if count(p, TREE_REC) != lo // size:
            return False
        with open(p, "ab") as f:
            f.write(pad(text, TREE_REC))
            f.flush()
            os.fsync(f.fileno())
        return True
    finally:
        lock.close()


def tree_drop(d, lo, hi):
    """Forget block [lo,hi) and every block built from it, by truncating each
    level back to that point. Later blocks at those levels go too and are
    rebuilt; the log is never touched, so nothing is lost."""
    gone, size = [], hi - lo
    lock = locked(d)
    try:
        while size <= log_len(d):
            p, k = tree_path(d, size), lo // size
            n = count(p, TREE_REC)
            if n > k:
                gone += [(i * size, (i + 1) * size) for i in range(k, n)]
                with open(p, "r+b") as f:
                    f.truncate(k * TREE_REC)
            size *= 2
        return gone
    finally:
        lock.close()


def die(msg):
    print(msg, file=sys.stderr)
    sys.exit(1)


def check(text):
    text = text.strip()
    if not text:
        die("REJECTED: empty.")
    if "\n" in text or "\r" in text:
        die("REJECTED: %d lines. A memory is exactly one line." %
            (text.count("\n") + 1))
    if len(text) > ENTRY_CHARS:
        die("REJECTED: %d chars, %d over the %d limit. Compress it further."
            % (len(text), len(text) - ENTRY_CHARS, ENTRY_CHARS))
    return text


# ---------------------------------------------------------------- naps

def pending(d, T, limit=None):
    """Blocks that can be built and have not been, smallest first. Each level
    file holds a dense prefix, so its length says exactly how far that level
    got: this costs one stat per level, never a scan."""
    todo, size = [], 2
    while size <= T:
        have = count(tree_path(d, size), TREE_REC)
        for k in range(have, T // size):
            todo.append((k * size, (k + 1) * size))
            if limit and len(todo) >= limit:
                return todo
        size *= 2
    return todo


def pending_count(d, T):
    n, size = 0, 2
    while size <= T:
        n += T // size - count(tree_path(d, size), TREE_REC)
        size *= 2
    return n


def nap_prompt(d, lo, hi, left):
    if hi - lo <= RAW_MAX:
        body = "\n".join("  #%d %s %s" % e for e in log_slice(d, lo, hi))
        what = "these %d memories" % (hi - lo)
    else:
        mid = (lo + hi) // 2
        body = "\n".join("  " + tree_get(d, a, b) for a, b in
                          ((lo, mid), (mid, hi)))
        what = "these two summaries"
    return (
        "You are dreaming. Compress {what} into ONE line of at most {n} "
        "characters.\nKeep every name, number, date and decision. Drop wording, "
        "not facts.\nInvent nothing. Write it as a memory, not as a description "
        "of memories.\n\n{body}\n\nThen run exactly:\n"
        "  memo sleep {lo}-{hi} \"<your line>\"\n\n"
        "{left} nap(s) left after this one."
    ).format(what=what, n=ENTRY_CHARS, body=body, lo=lo, hi=hi, left=left)


def next_nap(d, T):
    todo = pending(d, T, limit=1)
    if not todo:
        return None
    lo, hi = todo[0]
    return nap_prompt(d, lo, hi, pending_count(d, T) - 1)


# ---------------------------------------------------------------- commands

def paginate(lines):
    """Split the document into parts that survive any harness's output cap."""
    parts, cur, size = [], [], 0
    for line in lines:
        if cur and (len(cur) >= PART_LINES or size + len(line) > PART_CHARS):
            parts.append(cur)
            cur, size = [], 0
        cur.append(line)
        size += len(line) + 1
    if cur:
        parts.append(cur)
    return parts


def cmd_wake(d, args):
    T = log_len(d)
    nap = next_nap(d, T)
    if nap:
        print("YOU CANNOT WAKE UP YET: %d compression(s) are pending, and a "
              "memory\nwith work left in it is not yet the truth. Sleep first "
              "-- it is quick.\n" % pending_count(d, T))
        print(nap)
        sys.exit(1)
    if not T:
        print("You have no memories yet. This is your first moment.\n"
              "Record what matters with: memo note \"...\"")
        return
    lines = []
    for lo, hi in cover(T, WAKE_LINES):
        if hi - lo == 1:
            lines.append("#%d %s %s" % log_get(d, lo))
        else:
            lines.append("#%d-%d %s" % (lo, hi - 1, tree_get(d, lo, hi)))
    parts = paginate(lines)
    k = 1
    if args:
        if len(args) != 1 or not args[0].isdigit():
            die("usage: memo wake [part]")
        k = int(args[0])
        if not 1 <= k <= len(parts):
            die("There is no part %d. Your memory has %d." % (k, len(parts)))
    if len(parts) > 1:
        print("--- your memory, part %d of %d (oldest first) ---" % (k, len(parts)))
    print("\n".join(parts[k - 1]))
    if k < len(parts):
        print("--- END OF PART %d. You are still half awake: you have NOT seen "
              "your most recent\n    memories yet, and acting now would be "
              "acting blind. Run `memo wake %d`\n    RIGHT NOW, and keep going "
              "until a part tells you that was the last one. ---" % (k, k + 1))
    elif len(parts) > 1:
        print("--- END OF PART %d, the last one. That is your whole memory. "
              "You are awake. ---" % k)


def cmd_note(d, args):
    if len(args) != 1:
        die("usage: memo note \"<one line, at most %d chars>\"" % ENTRY_CHARS)
    text = check(args[0])
    T = log_len(d)
    log_append(d, ["#%d %s %s" % (T, datetime.date.today().isoformat(), text)])
    print("ok, memory #%d." % T)
    nap = next_nap(d, T + 1)
    if nap:
        print("\n" + nap)


def cmd_sleep(d, args):
    T = log_len(d)
    if args:
        if len(args) != 2:
            die("usage: memo sleep <lo>-<hi> \"<one line>\"")
        m = re.fullmatch(r"(\d+)-(\d+)", args[0])
        if not m:
            die("REJECTED: '%s' is not a block id. Copy it from the prompt."
                % args[0])
        lo, hi = int(m.group(1)), int(m.group(2))
        todo = pending(d, T, limit=1)
        if not todo:
            die("REJECTED: nothing is pending. You are already awake.")
        if (lo, hi) != todo[0]:
            die("REJECTED: %d-%d is not the block to compress. Blocks are built "
                "in order,\nand yours is %d-%d. Run `memo sleep` to see it."
                % (lo, hi, todo[0][0], todo[0][1]))
        if not tree_put(d, lo, hi, check(args[1])):
            print("Already dreamt; another session got there first. Skipping.")
        else:
            print("ok, %d-%d remembered." % (lo, hi))
    nap = next_nap(d, T)
    if not nap:
        print("You woke up. Nothing left to compress.")
        return
    print("\n" + nap)


def cmd_forget(d, args):
    """A summary can be wrong -- mistyped, or a bad compression. Drop it and
    everything built on top of it; the next sleep computes them again. The log
    is untouched, so nothing is ever actually lost."""
    if len(args) != 1:
        die("usage: memo forget <lo>-<hi>")
    m = re.fullmatch(r"(\d+)-(\d+)", args[0])
    if not m:
        die("REJECTED: '%s' is not a block id." % args[0])
    lo, hi = int(m.group(1)), int(m.group(2))
    size = hi - lo
    if size < 2 or size & (size - 1) or lo % size:
        die("REJECTED: %d-%d is not a block. A block covers an aligned "
            "power-of-two range." % (lo, hi))
    gone = tree_drop(d, lo, hi)
    if not gone:
        die("There is no summary at %d-%d to forget." % (lo, hi))
    print("forgot %d summaries (%d-%d and everything built from it). They will "
          "be compressed again on your next sleep."
          % (len(gone), gone[0][0], gone[0][1]))


def cmd_recall(d, args):
    if len(args) != 1:
        die("usage: memo recall <regex>")
    try:
        pat = re.compile(args[0], re.I)
    except re.error as e:
        die("bad regex: %s" % e)
    hits = [e for e in log_slice(d, 0, log_len(d)) if pat.search(e[2])]
    if not hits:
        print("Nothing in your memory matches that.")
        return
    for e in hits:
        print("#%d %s %s" % e)
    print("\n%d memories matched." % len(hits))


def cmd_import(d, args):
    """Bulk-append historical memories: 'YYYY-MM-DD <text>' per line.
    For bootstrapping an identity from older records. Used once."""
    if len(args) != 1:
        die("usage: memo import <file>   # lines of 'YYYY-MM-DD <text>'")
    T = log_len(d)
    last = log_get(d, T - 1)[1] if T else "0000-00-00"
    out = []
    for i, line in enumerate(open(args[0]), 1):
        line = line.rstrip("\n")
        if not line.strip():
            continue
        date, _, text = line.partition(" ")
        if not re.fullmatch(r"\d{4}-\d{2}-\d{2}", date):
            die("line %d: expected 'YYYY-MM-DD <text>', got: %s" % (i, line))
        if date < last:
            die("line %d: date %s is older than the previous memory (%s). "
                "Memories must be in order." % (i, date, last))
        text = text.strip()
        if not text or len(text) > ENTRY_CHARS:
            die("line %d: %d chars (limit %d)." % (i, len(text), ENTRY_CHARS))
        out.append("#%d %s %s" % (T + len(out), date, text))
        last = date
    log_append(d, out)
    print("imported %d memories (#%d..#%d)." % (len(out), T, T + len(out) - 1))
    n = pending_count(d, log_len(d))
    if n:
        print("%d compressions are now pending. Run `memo sleep` until it "
              "says you woke up." % n)


COMMANDS = {"wake": cmd_wake, "note": cmd_note, "sleep": cmd_sleep,
            "recall": cmd_recall, "forget": cmd_forget, "import": cmd_import}


def main():
    if len(sys.argv) < 2 or sys.argv[1] not in COMMANDS:
        print(__doc__.strip())
        sys.exit(0 if len(sys.argv) < 2 else 1)
    d = store()
    config(d)
    COMMANDS[sys.argv[1]](d, sys.argv[2:])


if __name__ == "__main__":
    main()
