#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
patch_streaming.py - Stellt alle Band-Textgeneratoren auf Streaming um
und setzt MAX_TOKENS einheitlich auf 64000.
- Sucht rekursiv text_generator*.py in band*/ Unterordnern.
- Pro Datei: Backup, MAX_TOKENS -> 64000, messages.create(...) -> stream-Block.
- Bricht pro Datei ab, wenn das create-Muster nicht eindeutig erkennbar ist.
"""
import os, re, glob, time, shutil, sys

BASE = "."
TARGET = 64000

# alle Textgenerator-Dateien in band*/ finden
dateien = sorted(glob.glob(os.path.join(BASE, "band*", "text_generator*.py")))
if not dateien:
    print("KEINE Textgenerator-Dateien gefunden!"); sys.exit(1)

print("Gefundene Dateien:")
for d in dateien:
    print("  ", d)
print()

# Regex fuer den create-Aufruf:
#   resp = client.messages.create(model=MODEL, max_tokens=MAX_TOKENS,
#       system=..., messages=[...])
# Wir ersetzen "client.messages.create(" -> Streaming-Konstrukt.
# Strategie: finde 'XXX = client.messages.create(' ... ')' (balancierte Klammern)
# und wandle in:
#   with client.messages.stream(<args>) as _stream:
#       XXX = _stream.get_final_message()

def finde_create_bloecke(code):
    """Liefert Liste (start, end, varname, args_str) fuer jeden
    '<var> = client.messages.create( ... )'-Aufruf."""
    treffer = []
    for m in re.finditer(r'(\w+)\s*=\s*client\.messages\.create\(', code):
        var = m.group(1)
        start = m.start()
        # ab der Klammer balancieren
        i = m.end() - 1  # Position der '('
        tiefe = 0
        j = i
        while j < len(code):
            if code[j] == '(':
                tiefe += 1
            elif code[j] == ')':
                tiefe -= 1
                if tiefe == 0:
                    break
            j += 1
        args = code[i+1:j]
        treffer.append((start, j+1, var, args))
    return treffer

gesamt_ok = 0
for pfad in dateien:
    with open(pfad, "r", encoding="utf-8") as f:
        code = f.read()

    orig = code
    aenderungen = []

    # 1) MAX_TOKENS setzen (egal welcher Wert)
    neu_code, n = re.subn(r'MAX_TOKENS\s*=\s*\d+', 'MAX_TOKENS = %d' % TARGET, code)
    if n >= 1:
        code = neu_code
        aenderungen.append("MAX_TOKENS->%d (%dx)" % (TARGET, n))

    # 2) create -> stream
    bloecke = finde_create_bloecke(code)
    if not bloecke:
        print("UEBERSPRUNGEN (kein create):", pfad)
        # trotzdem MAX_TOKENS speichern, falls geaendert
        if code != orig:
            ts = time.strftime("%Y%m%d_%H%M%S")
            shutil.copy2(pfad, pfad + ".bak_stream_" + ts)
            with open(pfad, "w", encoding="utf-8") as f:
                f.write(code)
            print("  -> nur MAX_TOKENS gespeichert")
        continue

    # von hinten nach vorne ersetzen (Indizes bleiben gueltig)
    for (start, end, var, args) in reversed(bloecke):
        # Einrueckung der Zeile bestimmen
        zeilen_start = code.rfind("\n", 0, start) + 1
        indent = ""
        k = zeilen_start
        while k < len(code) and code[k] in " \t":
            indent += code[k]; k += 1
        ersatz = (
            "with client.messages.stream(" + args + ") as _stream:\n"
            + indent + "    " + var + " = _stream.get_final_message()"
        )
        code = code[:start] + ersatz + code[end:]
        aenderungen.append("create->stream (%s)" % var)

    ts = time.strftime("%Y%m%d_%H%M%S")
    shutil.copy2(pfad, pfad + ".bak_stream_" + ts)
    with open(pfad, "w", encoding="utf-8") as f:
        f.write(code)
    print("GEPATCHT:", pfad)
    for a in aenderungen:
        print("   -", a)
    gesamt_ok += 1

print("\nFertig. %d Dateien gepatcht." % gesamt_ok)
