File size: 5,845 Bytes
be6c5ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Regression tests for code execution shell lifecycle behavior.

Pagers (more/less) must be disabled in the non-interactive shells created by the
code execution tool: without user input they block forever and spin at 100% CPU.
"""

import asyncio
import importlib
from types import SimpleNamespace

from plugins._code_execution.helpers import shell_local, shell_ssh
from plugins._code_execution.helpers.tty_session import TTYSession
from plugins._code_execution.tools.code_execution_tool import (
    CodeExecution,
    ShellWrap,
    State,
    _group_multiline_command,
    _is_closed_pty_error,
)


def test_local_env_disables_pagers_and_preserves_existing():
    env = shell_local.disable_pagers_in_env({"PATH": "/usr/bin", "PAGER": "less"})
    assert env["PAGER"] == "cat"
    assert env["GIT_PAGER"] == "cat"
    # pre-existing keys are preserved
    assert env["PATH"] == "/usr/bin"


def test_local_env_defaults_to_environ():
    env = shell_local.disable_pagers_in_env()
    assert env["PAGER"] == "cat"
    assert env["GIT_PAGER"] == "cat"


def test_local_env_does_not_mutate_input():
    src = {"PATH": "/usr/bin"}
    shell_local.disable_pagers_in_env(src)
    assert src == {"PATH": "/usr/bin"}


def test_ssh_command_disables_pagers():
    assert "GIT_PAGER=cat" in shell_ssh.PAGER_DISABLE_COMMAND
    assert "PAGER=cat" in shell_ssh.PAGER_DISABLE_COMMAND


def test_paramiko_import_error_does_not_retain_tool_loading_stack(monkeypatch):
    try:
        raise ImportError("invoke")
    except ImportError as error:
        saved_error = error
        monkeypatch.setattr(shell_ssh.paramiko.config, "invoke_import_error", error)

    importlib.reload(shell_ssh)

    assert shell_ssh.paramiko.config.invoke_import_error is saved_error
    assert saved_error.__traceback__ is None


def test_multiline_terminal_commands_are_one_current_shell_compound():
    assert _group_multiline_command("pwd") == "pwd"
    assert _group_multiline_command("cd /tmp\npwd") == "{\ncd /tmp\npwd\n}"
    assert _group_multiline_command("$env:FOO='bar'\n$env:FOO", powershell=True) == (
        ". {\n$env:FOO='bar'\n$env:FOO\n}"
    )


def test_exited_tty_process_is_a_recoverable_closed_session():
    assert _is_closed_pty_error(RuntimeError("TTYSpawn process has exited"))


def test_tty_close_kills_term_resistant_process():
    async def run():
        session = TTYSession("bash -lc 'trap \"\" TERM; sleep 30'")
        await session.start()
        await asyncio.wait_for(session.close(), timeout=6)
        assert session._proc is None

    asyncio.run(run())


def test_tty_reports_strict_mode_shell_exit():
    async def run():
        session = TTYSession("/bin/bash --noprofile --norc -i")
        await session.start()
        await session.read_full_until_idle(idle_timeout=0.05, total_timeout=1)
        await session.sendline("{\nset -euo pipefail\nfalse\nprintf 'unreachable\\n'\n}")

        exit_code = await asyncio.wait_for(session.wait(), timeout=5)

        assert exit_code != 0
        assert session.is_terminated()
        assert session.get_exit_code() == exit_code
        await session.close()

    asyncio.run(run())


def test_ssh_session_reports_channel_exit_status():
    class FakeChannel:
        closed = False

        @staticmethod
        def exit_status_ready():
            return True

        @staticmethod
        def recv_exit_status():
            return 7

    session = object.__new__(shell_ssh.SSHInteractiveSession)
    session.shell = FakeChannel()
    session.client = SimpleNamespace(
        get_transport=lambda: SimpleNamespace(is_active=lambda: True)
    )
    session._exit_code = None

    assert session.is_terminated()
    assert session.get_exit_code() == 7


def test_code_execution_returns_immediately_when_shell_exits():
    class FinishedSession:
        async def read_output(self, timeout=0, reset_full_output=False):
            return "nothing to commit, working tree clean\n", "nothing to commit, working tree clean\n"

        @staticmethod
        def is_terminated():
            return True

        @staticmethod
        def get_exit_code():
            return 1

    class FakeAgent:
        agent_name = "test"

        async def handle_intervention(self):
            return None

        @staticmethod
        def read_prompt(name, **kwargs):
            if name == "fw.code.shell_exit.md":
                return f"Terminal shell exited{kwargs['status']}. The command has finished."
            if name == "fw.code.info.md":
                return f"[SYSTEM: {kwargs['info']}]"
            raise AssertionError(f"Unexpected prompt: {name}")

    async def run():
        session = FinishedSession()
        state = State(
            ssh_enabled=False,
            shells={0: ShellWrap(id=0, session=session, running=True)},
        )
        tool = CodeExecution(
            FakeAgent(),
            "code_execution_tool",
            "",
            {"runtime": "terminal", "session": 0},
            "",
            None,
        )
        updates = []
        tool.log = SimpleNamespace(update=lambda **kwargs: updates.append(kwargs))

        async def prepare_state(*args, **kwargs):
            return state

        async def set_progress(content):
            return None

        tool.prepare_state = prepare_state
        tool.set_progress = set_progress
        tool.fix_full_output = lambda output: output

        response = await tool.get_terminal_output(
            {"prompt_patterns": [], "dialog_patterns": []},
            session=0,
            sleep_time=0,
        )

        assert "nothing to commit" in response
        assert "exit code 1" in response
        assert "command has finished" in response
        assert not state.shells[0].running
        assert updates[-1]["heading"].endswith(" icon://done_all")

    asyncio.run(run())