File size: 11,626 Bytes
98acb70
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
-- tools.lua - Tool registry and the built-in set.
--
-- A tool is a plain table:
--   name              string, matched verbatim against the model's output
--   description       one line; this is prompt real estate, keep it short
--   args              { {name=, type="string"|"number", required=bool}, ... }
--   run               function(args) -> result_string   (may raise)
--   requires_approval bool, gates the tool behind the agent's approve callback
--   timeout_ms        number|nil

local protocol = require('protocol')

local tools = {}

--------------------------------------------------------------------------
-- Registry
--------------------------------------------------------------------------

local Registry = {}
Registry.__index = Registry

-- The answer pseudo-tool is always present: it is how the loop terminates,
-- so it must appear in the catalogue and in the grammar's enum of names.
function tools.registry()
    local self = setmetatable({ by_name = {}, order = {} }, Registry)
    self:add(protocol.answer_tool)
    return self
end

function Registry:add(tool)
    if type(tool) ~= "table" or type(tool.name) ~= "string" then
        error("tools: a tool needs a name")
    end
    if not self.by_name[tool.name] then
        self.order[#self.order + 1] = tool.name
    end
    self.by_name[tool.name] = tool
    return self
end

function Registry:get(name)
    return self.by_name[name]
end

-- Stable order, so the rendered catalogue and the compiled grammar are
-- byte-identical across runs. Traces stop diffing otherwise.
function Registry:list()
    local out = {}
    for _, n in ipairs(self.order) do out[#out + 1] = self.by_name[n] end
    return out
end

function Registry:names()
    local out = {}
    for _, n in ipairs(self.order) do out[#out + 1] = n end
    return out
end

-- Validate a parsed call before running it. Returns ok, err_or_args.
-- A validation failure is an observation, not a crash: the message goes back
-- to the model, so it is phrased for the model to act on.
function Registry:validate(name, args)
    local tool = self.by_name[name]
    if not tool then
        local avail = table.concat(self:names(), ", ")
        return false, string.format(
            "No tool named %q. Available tools: %s.", name, avail)
    end
    args = args or {}
    local clean = {}
    for _, spec in ipairs(tool.args or {}) do
        local v = args[spec.name]
        if v == nil or v == "" then
            if spec.required then
                return false, string.format(
                    "The %s tool needs a %q argument.", name, spec.name)
            end
        else
            if spec.type == "number" then
                local n = tonumber(v)
                if not n then
                    return false, string.format(
                        "The %q argument of %s must be a number, got %q.",
                        spec.name, name, tostring(v))
                end
                clean[spec.name] = n
            else
                clean[spec.name] = tostring(v)
            end
        end
    end
    return true, clean
end

-- Execute with validation, the approval gate, and error capture applied.
-- Returns ok, result_or_error, status  where status is one of
-- "continue" | "tool_error" | "timeout" | "denied".
function Registry:invoke(name, args, approve)
    local ok, clean = self:validate(name, args)
    if not ok then
        return false, "Error: " .. clean, "tool_error"
    end

    local tool = self.by_name[name]
    if tool.requires_approval then
        -- Default-deny. An agent that writes to disk because nobody wired up
        -- an approval callback is the failure mode this guards against.
        if not approve or not approve({ tool = name, args = clean }) then
            return false, string.format(
                "Denied: the user did not approve calling %s.", name), "denied"
        end
    end

    local started = os.clock()
    local ran, result = pcall(tool.run, clean)
    local elapsed_ms = (os.clock() - started) * 1000

    if not ran then
        -- Strip Lua's file:line prefix; the model cannot act on it and it
        -- costs tokens on every retry.
        local msg = tostring(result):gsub("^.-:%d+:%s*", "")
        return false, "Error: " .. msg, "tool_error"
    end
    if tool.timeout_ms and elapsed_ms > tool.timeout_ms then
        return false, string.format(
            "Error: %s timed out after %.0fms.", name, elapsed_ms), "timeout"
    end
    return true, tostring(result), "continue"
end

--------------------------------------------------------------------------
-- calc: a real expression evaluator.
--
-- The obvious implementation is load("return " .. expr). Do not do that: the
-- expression is model output, and load() hands it the entire Lua runtime.
-- This is the single most common way a toy agent becomes a remote code
-- execution hole, so the repo shows the ~50-line alternative instead.
--
-- Grammar:  expr := term (('+' | '-') term)*
--           term := power (('*' | '/' | '%') power)*
--           power := unary ('^' power)?          -- right associative
--           unary := '-'? primary
--           primary := number | '(' expr ')'
--------------------------------------------------------------------------

local function tokenize(s)
    local out, i = {}, 1
    while i <= #s do
        local c = s:sub(i, i)
        if c:match("%s") then
            i = i + 1
        elseif c:match("%d") or (c == "." and s:sub(i + 1, i + 1):match("%d")) then
            local num = s:match("^%d*%.?%d+", i) or s:match("^%d+", i)
            out[#out + 1] = { t = "num", v = tonumber(num) }
            i = i + #num
        elseif ("+-*/%^()"):find(c, 1, true) then
            out[#out + 1] = { t = c }
            i = i + 1
        else
            return nil, string.format(
                "%q is not arithmetic. calc evaluates numbers and + - * / %% ^ only.", s)
        end
    end
    if #out == 0 then
        return nil, "calc needs an arithmetic expression, for example 4871 * 209."
    end
    return out
end

local function evaluate(toks)
    local pos = 1
    local function peek() return toks[pos] and toks[pos].t end
    local expr

    local function primary()
        local tk = toks[pos]
        if not tk then error("unexpected end of expression", 0) end
        if tk.t == "num" then pos = pos + 1 return tk.v end
        if tk.t == "(" then
            pos = pos + 1
            local v = expr()
            if peek() ~= ")" then error("missing closing parenthesis", 0) end
            pos = pos + 1
            return v
        end
        error("unexpected " .. tk.t .. " in expression", 0)
    end

    local function unary()
        if peek() == "-" then pos = pos + 1 return -unary() end
        if peek() == "+" then pos = pos + 1 return unary() end
        return primary()
    end

    local function power()
        local base = unary()
        if peek() == "^" then
            pos = pos + 1
            return base ^ power()  -- right associative
        end
        return base
    end

    local function term()
        local v = power()
        while peek() == "*" or peek() == "/" or peek() == "%" do
            local op = peek()
            pos = pos + 1
            local rhs = power()
            if (op == "/" or op == "%") and rhs == 0 then
                error("division by zero", 0)
            end
            if op == "*" then v = v * rhs
            elseif op == "/" then v = v / rhs
            else v = v % rhs end
        end
        return v
    end

    expr = function()
        local v = term()
        while peek() == "+" or peek() == "-" do
            local op = peek()
            pos = pos + 1
            if op == "+" then v = v + term() else v = v - term() end
        end
        return v
    end

    local result = expr()
    if pos <= #toks then
        error("unexpected " .. tostring(toks[pos].t) .. " after the expression", 0)
    end
    return result
end

-- Format without float noise: 1018039, not 1018039.0
local function format_number(n)
    if n ~= n then return "nan" end
    if n == math.huge or n == -math.huge then return tostring(n) end
    if n % 1 == 0 and math.abs(n) < 1e15 then return string.format("%d", n) end
    return (string.format("%.10g", n))
end

tools.calc = {
    name = "calc",
    description = "Evaluate an arithmetic expression.",
    args = { { name = "expr", type = "string", required = true } },
    run = function(args)
        local toks, err = tokenize(args.expr)
        if not toks then error(err, 0) end
        return format_number(evaluate(toks))
    end,
}

--------------------------------------------------------------------------
-- Filesystem tools
--------------------------------------------------------------------------

-- Observations are charged against the context budget, so a tool that can
-- return a whole file truncates itself. context.lua handles the window; this
-- stops one pathological read from blowing it in a single step.
tools.MAX_OBSERVATION = 4000

local function truncate(s)
    if #s <= tools.MAX_OBSERVATION then return s end
    return s:sub(1, tools.MAX_OBSERVATION) .. string.format(
        "\n... [truncated, %d bytes total]", #s)
end

tools.read_file = {
    name = "read_file",
    description = "Read a text file.",
    args = { { name = "path", type = "string", required = true } },
    run = function(args)
        local fh, err = io.open(args.path, "r")
        if not fh then
            -- Name the failure precisely: a model that is told the path does
            -- not exist can correct it; one told "error" retries the same path.
            error(string.format("cannot open %q (%s)", args.path,
                tostring(err):gsub("^.*: ", "")), 0)
        end
        local content = fh:read("a")
        fh:close()
        return truncate(content)
    end,
}

tools.list_dir = {
    name = "list_dir",
    description = "List the entries of a directory.",
    args = { { name = "path", type = "string", required = true } },
    run = function(args)
        -- No POSIX module, so shell out. Quoting matters: the path is model
        -- output. Reject anything that could escape the argument.
        if args.path:match('[`$;|&<>"\n]') then
            error("path contains characters that are not allowed", 0)
        end
        local cmd = package.config:sub(1, 1) == "\\"
            and string.format('dir /b "%s" 2>nul', args.path)
            or string.format("ls -1 '%s' 2>/dev/null", args.path)
        local pipe = io.popen(cmd, "r")
        if not pipe then error("cannot list " .. args.path, 0) end
        local out = pipe:read("a")
        pipe:close()
        if not out or out == "" then
            error(string.format("%q is empty or does not exist", args.path), 0)
        end
        return truncate(out)
    end,
}

tools.write_file = {
    name = "write_file",
    description = "Write text to a file, replacing it.",
    args = { { name = "path", type = "string", required = true },
             { name = "text", type = "string", required = true } },
    requires_approval = true,
    run = function(args)
        local fh, err = io.open(args.path, "w")
        if not fh then error(tostring(err), 0) end
        fh:write(args.text)
        fh:close()
        return string.format("wrote %d bytes to %s", #args.text, args.path)
    end,
}

-- The default set, in catalogue order.
function tools.defaults()
    local reg = tools.registry()
    reg:add(tools.calc)
    reg:add(tools.read_file)
    reg:add(tools.list_dir)
    reg:add(tools.write_file)
    return reg
end

return tools