"""Interactive mode: step-through generation with session-based state. The frontend carries only the session id; the live tracer lives in the core SessionManager. Every handler resolves the session, takes its lock, mutates the tracer, and returns the canonical payload via ``serialize.render_state``. """ from __future__ import annotations from collections.abc import Iterator from api.helpers import ( ChatValidationError, layer_selection, parse_chat_messages, ui_sampling_params, ) from api.models import model_manager from api.serialize import error_payload, fig_json, render_state from api.state import get_active_interventions from miru_tracer.core.lens import compute_lens_slice, get_lens_store from miru_tracer.core.logging_config import get_logger from miru_tracer.core.session_manager import get_session_manager from miru_tracer.visualization.plots import plot_lens_heatmap logger = get_logger(__name__) def _resolve_session(session_id: str | None): """Common session lookup; returns (session, error_payload_or_None).""" if not session_id: return None, error_payload("Not initialized. Click 'Initialize' first.") session = get_session_manager().get_session(session_id) if session is None: return None, error_payload("Session not found. Please reinitialize.") return session, None def _reset_tracer_for_mode( tracer, mode: str, prompt: str, chat_json: str, raw_text: str, thinking: str, think_prefill: str, ) -> None: """Shared reset dispatch (raises ChatValidationError / ValueError).""" if mode == "chat": tracer.reset( messages=parse_chat_messages(chat_json), mode="chat", thinking=thinking or "auto", think_prefill=think_prefill or "", ) elif mode == "raw": tracer.reset(prompt=raw_text, mode="raw") else: tracer.reset(prompt=prompt, mode="completion") def interactive_init( mode: str, prompt: str, chat_json: str, raw_text: str, thinking: str, think_prefill: str, strategy: str, temperature: float, top_k: int, top_p: float, log_top_k: int, ) -> dict: model = model_manager.get_model() tokenizer = model_manager.get_tokenizer() device = model_manager.get_device() if model is None or tokenizer is None: return error_payload("No model loaded") try: session_manager = get_session_manager() session_id = session_manager.create_session(model, tokenizer, device) session = session_manager.get_session(session_id) with session.lock: _reset_tracer_for_mode( session.tracer, mode, prompt, chat_json, raw_text, thinking, think_prefill, ) params = ui_sampling_params(strategy, temperature, top_k, top_p) logger.info(f"Interactive session initialized: {session_id} (mode={mode})") return render_state( session_id, session.tracer, f"Initialized in {mode} mode", params, log_top_k, ) except ChatValidationError as e: return error_payload(str(e)) except Exception as e: logger.error(f"Initialize error: {e}", exc_info=True) return error_payload(str(e), trace=True) def interactive_reset(session_id: str) -> dict: if session_id: get_session_manager().delete_session(session_id) logger.info(f"Interactive session reset: {session_id}") return { "ok": True, "status": "Reset complete. Click 'Initialize' to start a new generation.", "session_id": None, "text": "", "step": 0, "candidates": [], "preview_id": None, "eos": False, } def interactive_step( session_id: str, strategy: str, temperature: float, top_k: int, top_p: float, selected_token_id: int | None, override_enabled: bool, override_id: int | None, log_top_k: int, log_full_probs: bool, stop_at_eos: bool, ) -> dict: session, error = _resolve_session(session_id) if error: return error with session.lock: tracer = session.tracer try: params = ui_sampling_params(strategy, temperature, top_k, top_p) if override_enabled: if override_id is None: return error_payload("Override enabled but no token ID provided") token_id = int(override_id) if not 0 <= token_id < len(tracer.tokenizer): return error_payload( f"Token ID {token_id} is out of range " f"(vocab size: {len(tracer.tokenizer)})" ) else: token_id = ( int(selected_token_id) if selected_token_id is not None else None ) step_data = tracer.step( params, token_id=token_id, log_top_k=max(int(log_top_k or 10), 1), log_full_probs=bool(log_full_probs), ) if stop_at_eos and tracer.is_eos(step_data.token_id): logger.info( f"EOS reached: session={session_id}, steps={len(tracer.history)}" ) return { "ok": True, "status": ( f"Generation complete (EOS reached)\n" f"Total steps: {len(tracer.history)}" ), "session_id": session_id, "text": tracer.get_full_text(), "step": len(tracer.history), "candidates": [], "preview_id": None, "eos": True, } status = ( f"Step {len(tracer.history)} complete\n" f"Generated: {step_data.token_text} (p={step_data.probability:.4f})" ) return render_state(session_id, tracer, status, params, log_top_k) except Exception as e: logger.error(f"Step error: {e}", exc_info=True) return error_payload(str(e), trace=True) def interactive_undo( session_id: str, strategy: str, temperature: float, top_k: int, top_p: float, log_top_k: int, ) -> dict: session, error = _resolve_session(session_id) if error: return error with session.lock: tracer = session.tracer try: if not tracer.undo(): return error_payload("No steps to undo") params = ui_sampling_params(strategy, temperature, top_k, top_p) return render_state( session_id, tracer, f"Undone last step. Current steps: {len(tracer.history)}", params, log_top_k, ) except Exception as e: logger.error(f"Undo error: {e}", exc_info=True) return error_payload(str(e), trace=True) def interactive_goto( session_id: str, target_step: int, strategy: str, temperature: float, top_k: int, top_p: float, log_top_k: int, ) -> dict: session, error = _resolve_session(session_id) if error: return error with session.lock: tracer = session.tracer current_steps = len(tracer.history) try: if target_step is None or target_step < 0: return error_payload( f"Target step must be 0 or greater. Current step: {current_steps}" ) if target_step > current_steps: return error_payload( f"Target step {int(target_step)} is beyond current step " f"{current_steps}" ) tracer.goto_step(int(target_step)) params = ui_sampling_params(strategy, temperature, top_k, top_p) undone = current_steps - int(target_step) status = ( f"Already at step {int(target_step)}" if undone == 0 else f"Went back to step {int(target_step)} (undid {undone} steps)" ) return render_state(session_id, tracer, status, params, log_top_k) except Exception as e: logger.error(f"Go-to-step error: {e}", exc_info=True) return error_payload(str(e), trace=True) def interactive_continue( session_id: str, strategy: str, temperature: float, top_k: int, top_p: float, n_tokens: int, log_top_k: int, log_full_probs: bool, stop_at_eos: bool, ) -> Iterator[dict]: """Run N steps, streaming progress; stops cooperatively via request_stop.""" session, error = _resolve_session(session_id) if error: yield error return if n_tokens is None or n_tokens < 1: yield error_payload("Number of tokens must be at least 1") return with session.lock: tracer = session.tracer try: params = ui_sampling_params(strategy, temperature, top_k, top_p) tracer.clear_stop_flag() logger.info(f"Continue generation: session={session_id}, n_tokens={n_tokens}") stopped_reason = None for i in range(int(n_tokens)): if tracer._stop_requested: stopped_reason = f"Generation stopped by user after {i} tokens" break step_data = tracer.step( params, log_top_k=max(int(log_top_k or 10), 1), log_full_probs=bool(log_full_probs), ) if stop_at_eos and tracer.is_eos(step_data.token_id): stopped_reason = "Generation complete (EOS reached)" break yield { "ok": True, "type": "progress", "status": ( f"Generating... Step {len(tracer.history)} " f"({i + 1}/{int(n_tokens)})" ), "text": tracer.get_full_text(), "step": len(tracer.history), } status = stopped_reason or "Continue complete" status += f"\nTotal steps: {len(tracer.history)}" final = render_state(session_id, tracer, status, params, log_top_k) final["type"] = "final" yield final except Exception as e: logger.error(f"Continue error: {e}", exc_info=True) payload = error_payload(str(e), trace=True) payload["type"] = "final" yield payload def interactive_stop(session_id: str) -> dict: """Request stop; the running Continue stream finalizes on its own.""" if session_id: session = get_session_manager().get_session(session_id) if session is not None: session.tracer.request_stop() logger.info(f"Stop requested for session {session_id}") return {"ok": True} def interactive_export( session_id: str, strategy: str, temperature: float, top_k: int, top_p: float, ) -> dict: """The full session log; the frontend saves it as a JSON download.""" session, error = _resolve_session(session_id) if error: return error with session.lock: params = ui_sampling_params(strategy, temperature, top_k, top_p) return {"ok": True, "export": session.tracer.export_to_dict(params)} def interactive_lens( session_id: str, mode: str, stride: int, top_k: int, ) -> dict: """Per-layer lens readout of the next-token position (Plotly heatmap).""" session, error = _resolve_session(session_id) if error: return error with session.lock: tracer = session.tracer if tracer.input_ids is None: return error_payload("Initialize a prompt first.") model_name = model_manager.get_model_name() jlens = get_lens_store().get(model_name) if mode in ("jacobian", "diff") and jlens is None: return error_payload( f"No fitted Jacobian lens for {model_name}. Upload one in the " f"Lens view, or fit one on a GPU box: miru-tracer-fit-lens {model_name}" ) try: n_layers = tracer.model.config.get_text_config().num_hidden_layers layers = layer_selection(n_layers, 0, -1, stride) if mode in ("jacobian", "diff") and jlens is not None: fitted = set(jlens.source_layers) | {n_layers - 1} layers = [layer for layer in layers if layer in fitted] slice_ = compute_lens_slice( tracer.model, tracer.tokenizer, tracer.input_ids, layers=layers, positions=[tracer.seq_len - 1], mode=mode, jlens=jlens, top_k=int(top_k), interventions=tracer._intervention_set, ) active = len(tracer.interventions) status = ( f"{mode} lens over {len(layers)} layers at position " f"{tracer.seq_len - 1}." ) if active: status += f" {active} intervention(s) active on this session." return {"ok": True, "figure": fig_json(plot_lens_heatmap(slice_)), "status": status} except Exception as e: logger.error(f"Interactive lens error: {e}", exc_info=True) return error_payload(str(e)) def interactive_apply_interventions(session_id: str) -> dict: """Apply the Lens view's active interventions to this session.""" session, error = _resolve_session(session_id) if error: return error interventions = get_active_interventions() with session.lock: jlens = get_lens_store().get(model_manager.get_model_name()) try: session.tracer.set_interventions(interventions or None, jlens=jlens) except ValueError as e: return error_payload(str(e)) if not interventions: return { "ok": True, "status": "No active interventions in the Lens view — session cleared.", } return { "ok": True, "status": ( f"Applied {len(interventions)} intervention(s) to this session. " "They affect all subsequent steps (KV cache was rebuilt)." ), }