File size: 15,629 Bytes
dec5552
 
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dec5552
 
 
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dec5552
 
34bec90
 
 
 
 
 
dec5552
 
34bec90
dec5552
 
 
34bec90
 
dec5552
 
34bec90
 
 
 
 
dec5552
34bec90
 
 
dec5552
34bec90
dec5552
34bec90
dec5552
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dec5552
34bec90
 
dec5552
 
 
34bec90
 
 
 
 
 
 
 
 
dec5552
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dec5552
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dec5552
34bec90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// Talks to a local OpenAI-compatible server (llama.cpp llama-server with the Maverick GGUF, Ollama, ...),
// executes the returned tool calls on the SceneWorld
// and feeds tool results back until the model answers in text, asks the user a clarifying question
// (ask_clarification hands the turn back: the answer arrives as the next request), or runs out of turns.
//
// The loop is a coroutine: HTTP runs on a worker thread, tool calls run on the main thread
// (Unity APIs are main-thread only). Play mode: StartCoroutine(agent.Run(...)).
// Batch/editor code can drive it synchronously with RunBlocking(...).
using System;
using System.Collections;
using System.Collections.Generic;
using System.Globalization;
using System.Linq;
using System.Net.Http;
using System.Text;
using System.Text.RegularExpressions;
using System.Threading;
using System.Threading.Tasks;
using UnityEngine;

namespace SceneAgent
{
    [Serializable]
    public class AgentTurn
    {
        public string modelOutput;
        public List<string> calls = new List<string>();
        public List<string> results = new List<string>();
        /// <summary>The output check's verdict on this answer (SceneGrammar.Check): ok, invalid_tool, invalid_value,
        /// invalid_id or malformed. Null when the check is off.</summary>
        public string check;
    }

    public class AgentResult
    {
        public List<AgentTurn> turns = new List<AgentTurn>();
        public string finalText;
        /// <summary>Set when the model asked the user a question (ask_clarification); the turn ends there.</summary>
        public string clarificationQuestion;
        public string error;
        public double seconds;
    }

    public class SceneAgent : MonoBehaviour
    {
        public string endpoint = "http://127.0.0.1:8081/v1/chat/completions";
        public string model = "unity-scene-agent";
        [Tooltip("The model's system prompt. Leave empty to use the prompt the model was trained with " +

                 "(Resources/SceneAgent/system_prompt). Change it only for a model trained on your prompt.")]
        public TextAsset systemPromptAsset;
        [Tooltip("Optional: the LocalModelServer that runs the model. Run() then waits until it is ready and uses its endpoint.")]
        public LocalModelServer server;
        [NonSerialized] public string systemPrompt;
        public int maxTurns = 4;
        public SceneWorld world;
        [Tooltip("Constrain every answer with a grammar: only the 11 tools and ids that exist. Needs llama-server. " +

                 "Not recommended: it can force an action where the model should decline. Use the output check below.")]
        public bool constrainedDecoding = false;
        [Tooltip("Output check (recommended): the model answers freely, then a call to a tool or value that does not exist " +

                 "becomes a refusal, and an unknown object id or a broken call is generated again under a grammar. " +

                 "Used when constrainedDecoding is off.")]
        public bool fallbackDecoding = true;
        [Tooltip("Safety: when the request names an object that is only in another room and the model acts on a different " +

                 "object here instead, that call is not executed; the model is told where the named object is. Example: " +

                 "'put the drill on the workbench' with the drill in another room must not move the screwdriver.")]
        public bool guardSubstitutes = true;

        static readonly HttpClient Http = new HttpClient { Timeout = TimeSpan.FromSeconds(120) };
        static readonly Regex ToolCallRe = new Regex(@"<tool_call>\s*(.*?)\s*</tool_call>", RegexOptions.Singleline);
        static readonly Regex ThinkRe = new Regex(@"<think>.*?</think>\s*", RegexOptions.Singleline);
        // These calls end the agent's turn: the user's answer arrives as the next request.
        static readonly HashSet<string> Handoff = new HashSet<string> { "ask_clarification" };

        static string defaultPrompt;
        static string DefaultPrompt => defaultPrompt ??= Resources.Load<TextAsset>("SceneAgent/system_prompt")?.text.Replace("\r\n", "\n");
        // no ?. on systemPromptAsset: in the Editor an unassigned field is a "fake null" Unity object, and ?. would call
        // .text on it (UnassignedReferenceException on the first command in Play mode)
        string Prompt => systemPrompt ?? (systemPromptAsset != null ? systemPromptAsset.text.Replace("\r\n", "\n") : null) ?? DefaultPrompt
                         ?? throw new InvalidOperationException("no system prompt (the package's Resources/SceneAgent/system_prompt is missing)");

        public IEnumerator Run(string utterance, Action<AgentResult> done)
        {
            var t0 = DateTime.UtcNow;
            var res = new AgentResult();
            if (server != null)
            {
                while (!server.Ready && !server.Failed) yield return null;
                if (server.Failed)
                {
                    res.error = "model server not available: " + server.status;
                    res.seconds = (DateTime.UtcNow - t0).TotalSeconds;
                    done(res);
                    yield break;
                }
                endpoint = server.Endpoint;
            }
            string sceneJson = world.BuildSceneJson();
            var messages = new List<object> {
                Msg("system", Prompt),
                Msg("user", "Scene: " + sceneJson + "\nRequest: " + utterance) };
            var toolResults = new List<string>();  // ids returned here become nameable (grammar)
            for (int turn = 0; turn < maxTurns && res.error == null; turn++)
            {
                var req = new Dictionary<string, object> {
                    { "model", model }, { "messages", messages }, { "temperature", 0.0 }, { "max_tokens", 192 } };
                if (constrainedDecoding) req["grammar"] = SceneGrammar.Build(sceneJson, toolResults);
                string body = MiniJson.Serialize(req);
                var task = Task.Run(() => Post(body));  // worker thread: no Unity API inside
                while (!task.IsCompleted) yield return null;
                if (task.IsFaulted) { res.error = task.Exception?.GetBaseException().Message; break; }
                string output = task.Result, check = null;
                if (!constrainedDecoding && fallbackDecoding)
                {
                    check = SceneGrammar.Check(output, sceneJson, toolResults);
                    if (check == "invalid_tool" || check == "invalid_value") output = SceneGrammar.Refusal;
                    else if (check == "invalid_id" || check == "malformed")
                    {
                        req["grammar"] = SceneGrammar.Build(sceneJson, toolResults);
                        string again = MiniJson.Serialize(req);
                        var retry = Task.Run(() => Post(again));
                        while (!retry.IsCompleted) yield return null;
                        if (retry.IsFaulted) { res.error = retry.Exception?.GetBaseException().Message; break; }
                        output = retry.Result;
                    }
                }
                var at = new AgentTurn { modelOutput = output, check = check };
                res.turns.Add(at);
                var blocks = ToolCallRe.Matches(output).Cast<Match>().Select(m => m.Groups[1].Value).ToList();
                if (blocks.Count == 0)
                {
                    res.finalText = Clean(output);
                    break;
                }
                var toolMsgs = new List<object>();
                string question = null;
                foreach (var block in blocks)
                {
                    Dictionary<string, object> c;
                    try { c = (Dictionary<string, object>)MiniJson.Parse(block); }
                    catch (Exception e) { res.error = "unparseable tool call: " + e.Message; break; }
                    string name = Convert.ToString(c.TryGetValue("name", out var n) ? n : "");
                    var args = c.TryGetValue("arguments", out var a) && a is Dictionary<string, object> d ? d : new Dictionary<string, object>();
                    object result;
                    string guard = SubstituteGuard(utterance, name, args);
                    if (guard != null)
                    {
                        Debug.Log("[SceneAgent] call not executed: " + guard);
                        result = new Dictionary<string, object> { { "ok", false }, { "error", guard } };
                    }
                    else
                    {
                        try { result = world.Execute(name, args); }
                        catch (ToolError e) { result = new Dictionary<string, object> { { "ok", false }, { "error", e.Message } }; }
                    }
                    if (Handoff.Contains(name))
                        question = Convert.ToString(args.TryGetValue("question", out var q) ? q : "") ?? "";
                    string resultJson = PyJson(result);
                    at.calls.Add(block);
                    at.results.Add(resultJson);
                    toolResults.Add(resultJson);
                    toolMsgs.Add(new Dictionary<string, object> { { "role", "tool" }, { "name", name }, { "content", resultJson } });
                }
                // history: the model's own call text, then tool results (Qwen3 renders them as <tool_response>)
                messages.Add(Msg("assistant", Clean(output)));
                messages.AddRange(toolMsgs);
                if (question != null) { res.clarificationQuestion = question; break; }
            }
            res.seconds = (DateTime.UtcNow - t0).TotalSeconds;
            done(res);
        }

        /// <summary>Synchronous driver for editor/batch code (main thread).</summary>
        public AgentResult RunBlocking(string utterance)
        {
            AgentResult result = null;
            var it = Run(utterance, r => result = r);
            while (it.MoveNext()) Thread.Sleep(2);
            return result;
        }

        static readonly HashSet<string> ActingTools = new HashSet<string> { "grab", "place", "move_by", "rotate", "set_state", "press" };

        /// <summary>Why this call must not run (see guardSubstitutes), or null. It fires only when the acted-on object is
        /// not named in the request while the request names an object that exists only in other rooms.</summary>
        public string SubstituteGuard(string utterance, string tool, Dictionary<string, object> args)
        {
            if (!guardSubstitutes || world == null || !ActingTools.Contains(tool ?? "")) return null;
            string id = args.TryGetValue("object_id", out var v) && v != null ? Convert.ToString(v, CultureInfo.InvariantCulture) : null;
            var item = id == null ? null : world.Items.FirstOrDefault(x => x.itemId == id);
            if (item == null) return null;
            if (string.IsNullOrEmpty(item.label)) return null;
            // tested on 13,438 model calls: it fired 4 times, each time on a wrong action, with no false alarm
            string text = " " + Regex.Replace(utterance.ToLowerInvariant(), "[^a-z0-9]+", " ") + " ";
            var words = text.Split(new[] { ' ' }, StringSplitOptions.RemoveEmptyEntries);
            string joined = string.Concat(words);
            IEnumerable<string> WordsOf(string label) => label.Split(' ').Where(w => w.Length >= 3);
            bool Phrase(string label) => text.Contains(" " + label + " ") || text.Contains(" " + label + "s ") || text.Contains(" " + label + "es ");
            // the acted-on object: any sign the user meant it ("box" -> toolbox, "note book" -> notebook, "keys" -> key chain)
            bool Loosely(SceneItem o)
            {
                string h = o.label.Split(' ').Last();
                return Phrase(o.label) || WordsOf(o.label).Any(w => words.Contains(w) || words.Contains(w + "s") || words.Contains(w + "es"))
                       || words.Any(w => w.Length >= 3 && h.EndsWith(w)) || joined.Contains(o.label.Replace(" ", ""));
            }
            if (Loosely(item)) return null;
            // an object elsewhere counts only when named by its full label, and only if nothing here shares a word with it
            var hereWords = new HashSet<string>(world.Items.Where(y => y.room == world.playerRoom && !string.IsNullOrEmpty(y.label))
                                                           .SelectMany(y => WordsOf(y.label)));
            var elsewhere = world.Items.FirstOrDefault(x => x.room != world.playerRoom && !string.IsNullOrEmpty(x.label)
                                                            && Phrase(x.label) && !WordsOf(x.label).Any(hereWords.Contains));
            if (elsewhere == null) return null;
            return $"not done: '{id}' is a {item.label}, but the request names a {elsewhere.label}; " +
                   $"the {elsewhere.label} ({elsewhere.itemId}) is in the {elsewhere.room}, not in this room";
        }

        static string Clean(string s) => ThinkRe.Replace(s, "").Replace("<|im_end|>", "").Trim();

        static Dictionary<string, object> Msg(string role, string content) =>
            new Dictionary<string, object> { { "role", role }, { "content", content } };

        string Post(string body)
        {
            var resp = Http.PostAsync(endpoint, new StringContent(body, Encoding.UTF8, "application/json")).ConfigureAwait(false).GetAwaiter().GetResult();
            string text = resp.Content.ReadAsStringAsync().ConfigureAwait(false).GetAwaiter().GetResult();
            if (!resp.IsSuccessStatusCode)
            {
                var m = Regex.Match(text, @"request \((\d+) tokens\) exceeds the available context size \((\d+) tokens\)");
                if (m.Success)
                    throw new Exception($"the scene and request need {m.Groups[1].Value} tokens but the model's context holds " +
                                        $"{m.Groups[2].Value}: put fewer objects in this room, or raise LocalModelServer.contextSize");
                throw new Exception($"HTTP {(int)resp.StatusCode}: {text}");
            }
            var root = (Dictionary<string, object>)MiniJson.Parse(text);
            var msg = (Dictionary<string, object>)((Dictionary<string, object>)((List<object>)root["choices"])[0])["message"];
            return Convert.ToString(msg.TryGetValue("content", out var c) ? c : "") ?? "";
        }

        // Tool results in the format the model was trained on (", " and ": " separators).
        public static string PyJson(object v)
        {
            switch (v)
            {
                case null: return "null";
                case bool b: return b ? "true" : "false";
                case string s: return MiniJson.Serialize(s);
                case double d: return MiniJson.FormatNumber(d);
                case IDictionary<string, object> dict:
                    return "{" + string.Join(", ", dict.Select(kv => MiniJson.Serialize(kv.Key) + ": " + PyJson(kv.Value))) + "}";
                case IEnumerable e:
                    return "[" + string.Join(", ", e.Cast<object>().Select(PyJson)) + "]";
                default: return MiniJson.Serialize(Convert.ToString(v, CultureInfo.InvariantCulture));
            }
        }
    }
}