File size: 11,732 Bytes
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5c28dc0
 
 
 
 
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0e4f105
 
41ea373
 
dfa9070
 
 
 
 
 
 
 
 
 
 
 
 
 
41ea373
 
 
 
 
dfa9070
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
0e4f105
41ea373
5c28dc0
 
41ea373
 
 
 
5c28dc0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0e4f105
 
 
 
5c28dc0
 
 
 
 
 
 
 
 
 
 
dfa9070
 
 
 
5c28dc0
 
 
41ea373
5c28dc0
 
 
41ea373
 
5c28dc0
 
 
 
 
 
 
 
 
 
 
 
 
 
0e4f105
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
5c28dc0
 
41ea373
 
0e4f105
 
41ea373
 
 
 
 
 
 
 
 
 
 
5c28dc0
 
 
 
41ea373
 
 
 
 
 
 
 
 
5c28dc0
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dfa9070
 
 
0e4f105
dfa9070
 
 
41ea373
5c28dc0
 
0e4f105
dfa9070
 
41ea373
 
 
5c28dc0
dfa9070
5c28dc0
dfa9070
 
 
 
 
 
41ea373
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
#!/usr/bin/env python3
import argparse
import json
import random
import sys
from pathlib import Path

from dotenv import load_dotenv
from rich.console import Console
from rich.panel import Panel
from rich.table import Table
from rich import box

load_dotenv()
sys.path.insert(0, str(Path(__file__).parent.parent.parent))

from viral_script_engine.environment.actions import ActionType
from viral_script_engine.environment.env import ViralScriptEnv

console = Console()
BASE_DIR = Path(__file__).parent.parent


def _bar(score: float, width: int = 8) -> str:
    filled = round(score * width)
    return "#" * filled + "." * (width - filled)


def build_random_action(action_type: ActionType) -> dict:
    labels = {
        ActionType.HOOK_REWRITE: ("hook", "Rewrite the hook to open with a specific number or bold claim."),
        ActionType.SECTION_REORDER: ("body", "Move the strongest point to immediately follow the hook."),
        ActionType.CULTURAL_REF_SUB: ("full", "Replace any generic references with locally relevant ones."),
        ActionType.CTA_PLACEMENT: ("cta", "Move the call-to-action earlier, before the 80% mark."),
    }
    section, instruction = labels[action_type]
    return {
        "action_type": action_type.value,
        "target_section": section,
        "instruction": instruction,
        "critique_claim_id": "C1",
        "reasoning": f"Demo run: applying {action_type.value}",
    }


def run_episode(difficulty: str, steps: int, verbose: bool) -> dict:
    scripts_path = str(BASE_DIR / "data" / "test_scripts" / "scripts.json")
    cultural_kb_path = str(BASE_DIR / "data" / "cultural_kb.json")
    env = ViralScriptEnv(scripts_path=scripts_path, cultural_kb_path=cultural_kb_path, max_steps=steps, difficulty=difficulty)

    obs, _ = env.reset()

    # Phase 8: show creator profile panel
    cp = obs.get("creator_profile") or {}
    if cp:
        console.print(Panel(
            f"Tier:        {cp.get('tier','?').capitalize()} ({cp.get('follower_count','?')} followers)\n"
            f"Frequency:   {cp.get('posting_frequency','?')}\n"
            f"Niche:       {cp.get('niche','?')}\n"
            f"Weak points: {', '.join(cp.get('past_weak_points', []))}\n"
            f"Voice:       {', '.join(cp.get('voice_descriptors', []))}",
            title="[bold cyan]CREATOR PROFILE[/bold cyan]",
            border_style="cyan",
        ))

    console.print(Panel(
        f"[bold]Episode started[/bold]\n"
        f"Difficulty: {difficulty}  |  Max steps: {steps}\n"
        f"Region: {obs['region']}  |  Platform: {obs['platform']}  |  Niche: {obs['niche']}\n"
        f"Episode ID: {obs['episode_id']}",
        title="[bold blue]Phase 8 Demo Episode[/bold blue]",
        border_style="blue",
    ))

    episode_log = {
        "episode_id": obs["episode_id"],
        "difficulty": difficulty,
        "steps": [],
        "final_state": None,
    }

    for step_num in range(steps):
        action_type = random.choice(list(ActionType))
        action = build_random_action(action_type)

        obs, reward, terminated, truncated, info = env.step(action, raw_output=None)
        rc = info["reward_components"]
        mod_out = info.get("moderation_output", {})
        orig_out = info.get("originality_output", {})

        if verbose:
            t = Table(title=f"Step {step_num + 1} β€” {action_type.value}", box=box.SIMPLE_HEAD)
            t.add_column("Metric", style="cyan", min_width=22)
            t.add_column("Score", min_width=12)
            t.add_column("Bar", min_width=10)

            def _row(label, key, suffix=""):
                val = rc.get(key)
                score_str = f"{val:.3f}" if val is not None else "N/A"
                bar_str = _bar(val) if val is not None else ""
                t.add_row(label, score_str + suffix, bar_str)

            _row("R1 Hook Strength", "r1_hook_strength")
            _row("R2 Coherence", "r2_coherence")
            _row("R3 Cultural", "r3_cultural_alignment")
            _row("R4 Resolution", "r4_debate_resolution")
            _row("R5 Preservation", "r5_defender_preservation")

            pr_val = rc.get("process_reward")
            pr_str = f"{pr_val:.3f}" if pr_val is not None else "N/A"
            t.add_row("Process Reward", pr_str, _bar(pr_val) if pr_val is not None else "")

            r6_val = rc.get("r6_safety")
            r6_suffix = "  [OK] No flags" if mod_out.get("total_flags", 0) == 0 else f"  [!] {mod_out.get('total_flags', 0)} flag(s)"
            r6_str = (f"{r6_val:.3f}{r6_suffix}" if r6_val is not None else "N/A")
            t.add_row("R6 Safety", r6_str, _bar(r6_val) if r6_val is not None else "")

            r7_val = rc.get("r7_originality")
            orig_flags = len(orig_out.get("flags", []))
            r7_suffix = f"  [!] {orig_flags} template match(es)" if orig_flags > 0 else "  [OK] Original"
            r7_str = (f"{r7_val:.3f}{r7_suffix}" if r7_val is not None else "N/A")
            t.add_row("R7 Originality", r7_str, _bar(r7_val) if r7_val is not None else "")

            r8_val = rc.get("r8_persona_fit")
            r8_str = f"{r8_val:.3f}" if r8_val is not None else "N/A"
            t.add_row("R8 Persona Fit", r8_str, _bar(r8_val) if r8_val is not None else "")

            t.add_row("-" * 22, "-" * 12, "-" * 10)
            t.add_row("[bold]Total[/bold]", f"[bold]{reward:.3f}[/bold]", _bar(reward))

            if info.get("anti_gaming_triggered"):
                t.add_row("Anti-Gaming Penalty", f"[red]{rc.get('anti_gaming_penalty', 0):.3f}[/red]", "")
                t.add_row("Penalty Reason", f"[red]{info.get('penalty_reason', '')}[/red]", "")
            t.add_row("Terminated", str(terminated), "")
            console.print(t)

            # Show moderation flags in red panel if any
            mod_flags = mod_out.get("flags", [])
            if mod_flags:
                flag_lines = []
                for fl in mod_flags:
                    flag_lines.append(
                        f"  [{fl['severity']}] {fl['category']} in {fl['position']}: \"{fl['trigger_phrase']}\" β†’ {fl['suggestion']}"
                    )
                console.print(Panel(
                    "\n".join(flag_lines),
                    title="[bold red]!! MODERATION FLAGS DETECTED[/bold red]",
                    border_style="red",
                ))

            # Show reasoning chain if present
            rc_chain = info.get("reasoning_chain")
            if rc_chain:
                chain_lines = []
                if rc_chain.get("priority_assessment"):
                    chain_lines.append(f"[cyan]Priority:[/cyan] {rc_chain['priority_assessment']}")
                cf = rc_chain.get("conflict_check_answer", "")
                if cf:
                    chain_lines.append(
                        f"[yellow]Conflict:[/yellow] {cf} β€” {rc_chain.get('conflict_check_reason', '')}"
                    )
                df = rc_chain.get("defender_consideration_answer", "")
                if df:
                    chain_lines.append(
                        f"[green]Defender:[/green] {df} β€” {rc_chain.get('defender_consideration_reason', '')}"
                    )
                pr_res = info.get("process_reward_result")
                if pr_res:
                    chain_lines.append(
                        f"[magenta]Process Scores:[/magenta] "
                        f"priority={pr_res['priority_score']:.2f}  "
                        f"conflict={pr_res['conflict_score']:.2f}  "
                        f"defender={pr_res['defender_score']:.2f}  "
                        f"total={pr_res['process_score']:.2f}"
                    )
                console.print(Panel(
                    "\n".join(chain_lines) if chain_lines else "[dim]No reasoning chain[/dim]",
                    title="[bold magenta]Reasoning Chain[/bold magenta]",
                    border_style="magenta",
                ))
            else:
                console.print(Panel(
                    "[dim]No reasoning chain β€” zero-shot decision[/dim]",
                    title="[bold magenta]Reasoning Chain[/bold magenta]",
                    border_style="dim",
                ))

            if obs.get("debate_history"):
                latest = obs["debate_history"][-1]
                if latest.get("rewrite_diff"):
                    console.print(Panel(
                        latest["rewrite_diff"][:600] or "(no diff)",
                        title="Script Diff",
                        border_style="yellow",
                    ))

        episode_log["steps"].append({
            "step": step_num + 1,
            "action": action,
            "reward": reward,
            "reward_components": rc,
            "moderation_output": mod_out,
            "originality_output": orig_out,
            "anti_gaming": info.get("anti_gaming_triggered", False),
            "terminated": terminated,
            "process_reward_result": info.get("process_reward_result"),
            "reasoning_chain": info.get("reasoning_chain"),
        })

        if terminated:
            break

    final_state = env.state()
    episode_log["final_state"] = final_state
    final_rc = final_state["reward_components"]

    console.print(Panel(
        f"[bold green]Final Reward:[/bold green] {final_rc.get('total', 0):.3f}\n"
        f"R1 Hook Strength:    {final_rc.get('r1_hook_strength', 'N/A')}\n"
        f"R2 Coherence:        {final_rc.get('r2_coherence', 'N/A')}\n"
        f"R6 Safety:           {final_rc.get('r6_safety', 'N/A')}\n"
        f"R7 Originality:      {final_rc.get('r7_originality', 'N/A')}\n"
        f"Steps completed: {final_state['step_num']}",
        title="Episode Summary",
        border_style="green",
    ))

    return episode_log


def main():
    parser = argparse.ArgumentParser(description="Run Phase 6 dummy episode")
    parser.add_argument("--difficulty", default="easy", choices=["easy", "medium", "hard"])
    parser.add_argument("--steps", type=int, default=3)
    parser.add_argument("--verbose", action="store_true")
    args = parser.parse_args()

    episode_log = run_episode(args.difficulty, args.steps, args.verbose)

    logs_dir = BASE_DIR / "logs"
    logs_dir.mkdir(exist_ok=True)
    log_path = logs_dir / f"episode_{episode_log['episode_id']}.json"
    with open(log_path, "w") as f:
        json.dump(episode_log, f, indent=2, default=str)
    console.print(f"[dim]Episode log saved -> {log_path}[/dim]")

    final_rc = episode_log["final_state"]["reward_components"]
    final_profile = episode_log["final_state"].get("creator_profile") or {}
    profile_tier = final_profile.get("tier", "")

    has_process_reward_key = "process_reward" in final_rc
    has_r8_key = "r8_persona_fit" in final_rc
    has_profile = bool(final_profile)

    gate_pass = (
        final_rc.get("r6_safety") is not None
        and final_rc.get("r7_originality") is not None
        and has_process_reward_key
        and has_r8_key
        and has_profile
        and log_path.exists()
    )
    style = "bold green" if gate_pass else "bold red"
    if gate_pass:
        label = f"PHASE 8 GATE: PASS β€” Creator persona active. R8 (persona fit) firing. Profile tier: {profile_tier}."
    else:
        missing = []
        if not has_r8_key:
            missing.append("r8_persona_fit missing from reward output")
        if not has_profile:
            missing.append("creator_profile missing from episode state")
        label = "PHASE 8 GATE: FAIL β€” " + "; ".join(missing) if missing else "PHASE 8 GATE: FAIL"
    console.print(Panel(f"[{style}]{label}[/{style}]", border_style="green" if gate_pass else "red"))


if __name__ == "__main__":
    main()