diff --git a/src/json.c b/src/json.c index 98a72bf..cc45a8e 100644 --- a/src/json.c +++ b/src/json.c @@ -27,7 +27,6 @@ static char *dupn(const char *s, size_t n) { static const char *parse_str(const char *s, char **out) { if (*s != '"') return NULL; s++; - const char *start = s; char buf[4096]; size_t n = 0; while (*s && *s != '"') { @@ -38,13 +37,20 @@ static const char *parse_str(const char *s, char **out) { if (c == 'n') c = '\n'; else if (c == 't') c = '\t'; else if (c == 'r') c = '\r'; + else if (c == 'b') c = '\b'; + else if (c == 'f') c = '\f'; + else if (c == 'u') { + for (int i = 0; i < 4; i++) { + if (!isxdigit((unsigned char)s[i])) return NULL; + } + } else if (c != '"' && c != '\\' && c != '/') return NULL; buf[n++] = c; } else { + if ((unsigned char)*s < 0x20) return NULL; buf[n++] = *s++; } } if (*s != '"') return NULL; - (void)start; *out = dupn(buf, n); return *out ? s + 1 : NULL; } @@ -53,9 +59,17 @@ static const char *parse_num(const char *s, JVal *v) { const char *b = s; if (*s == '-') s++; if (!isdigit((unsigned char)*s)) return NULL; - while (isdigit((unsigned char)*s)) s++; + if (*s == '0') s++; + else while (isdigit((unsigned char)*s)) s++; if (*s == '.') { s++; + if (!isdigit((unsigned char)*s)) return NULL; + while (isdigit((unsigned char)*s)) s++; + } + if (*s == 'e' || *s == 'E') { + s++; + if (*s == '+' || *s == '-') s++; + if (!isdigit((unsigned char)*s)) return NULL; while (isdigit((unsigned char)*s)) s++; } v->s = dupn(b, (size_t)(s - b)); @@ -117,11 +131,16 @@ static const char *parse_arr(const char *s, JVal *v) { static const char *parse_val(const char *s, JVal **out) { s = skip(s); + const char *start = s; JVal *v = NULL; if (*s == '"') { v = node(JSTR); if (!v) return NULL; s = parse_str(s, &v->s); + if (s) { + v->raw = dupn(start, (size_t)(s - start)); + if (!v->raw) s = NULL; + } } else if (*s == '{') { v = node(JOBJ); if (!v) return NULL; @@ -178,6 +197,7 @@ void jfree(JVal *v) { JVal *n = v->next; jfree(v->head); free(v->s); + free(v->raw); free(v->k); free(v); v = n; diff --git a/src/json.h b/src/json.h index 94546dc..0fc64b2 100644 --- a/src/json.h +++ b/src/json.h @@ -10,6 +10,7 @@ struct JVal { JType t; int b; char *s; /* JSTR, and JNUM lexeme */ + char *raw; /* owned JSTR JSON lexeme, for lossless protocol ID echo */ JVal *head; /* JARR/JOBJ children */ JVal *next; char *k; /* object member key */ diff --git a/src/mcp.c b/src/mcp.c index 84031e2..cca8208 100644 --- a/src/mcp.c +++ b/src/mcp.c @@ -37,11 +37,14 @@ static void reply_ok(const char *id, const char *body) { } static void reply_err(const char *id, int code, const char *msg) { - char line[1024]; - snprintf(line, sizeof(line), + size_t n = strlen(id) + strlen(msg) + 96; + char *line = malloc(n); + if (!line) return; + snprintf(line, n, "{\"jsonrpc\":\"2.0\",\"id\":%s,\"error\":{\"code\":%d,\"message\":\"%s\"}}", id, code, msg); emit(line); + free(line); } static char *b64enc(const uint8_t *src, size_t n, size_t *outn) { @@ -86,22 +89,10 @@ static void json_escape(const char *in, char *out, size_t n) { out[i] = 0; } -static void id_emit(const JVal *idv, char *out, size_t n) { - if (!idv || idv->t == JNULL) { - snprintf(out, n, "null"); - return; - } - if (idv->t == JNUM && idv->s) { - snprintf(out, n, "%s", idv->s); - return; - } - if (idv->t == JSTR && idv->s) { - char esc[128]; - json_escape(idv->s, esc, sizeof(esc)); - snprintf(out, n, "\"%s\"", esc); - return; - } - snprintf(out, n, "null"); +static const char *id_emit(const JVal *idv) { + if (idv && idv->t == JNUM && idv->s) return idv->s; + if (idv && idv->t == JSTR && idv->raw) return idv->raw; + return "null"; } static void tool_text(const char *id, const char *text, int is_err) { @@ -218,8 +209,7 @@ int gaze_mcp(uint16_t vid, uint16_t pid) { reply_err("null", -32700, "parse error"); continue; } - char id[64]; - id_emit(jobj(root, "id"), id, sizeof(id)); + const char *id = id_emit(jobj(root, "id")); const char *method = jstr(jobj(root, "method")); if (!method) { if (jobj(root, "id")) reply_err(id, -32600, "no method"); diff --git a/tests/test_protocol.py b/tests/test_protocol.py index 467b44f..50983ca 100644 --- a/tests/test_protocol.py +++ b/tests/test_protocol.py @@ -61,6 +61,53 @@ def call(self, method, params=None): class Protocol(unittest.TestCase): + def test_request_ids_round_trip(self): + # Success and error responses must remain valid JSON and echo the full ID. + ids = [ + "", + "request-" + "x" * 2000, + 'quoted"\\request', + "line\nreturn\rtab\t", + "controls\b\f\x00\x1f", + "café-東京-🎥", + 1234567890123456789012345678901234567890123456789012345678901234567890, + 1.25e30, + -1.25e-30, + None, + ] + mcp = Mcp() + try: + for request_id in ids: + for method in ("ping", "unknown-method", None): + with self.subTest(request_id=repr(request_id)[:80], method=method): + request = {"jsonrpc": "2.0", "id": request_id} + if method is not None: + request["method"] = method + reply = mcp.send(request) + self.assertEqual(reply["id"], request_id) + if method == "ping": + self.assertEqual(reply["result"], {}) + else: + self.assertEqual(reply["error"]["code"], + -32600 if method is None else -32601) + self.assertEqual(mcp.call("ping")["result"], {}) + finally: + mcp.close() + + def test_invalid_id_tokens_are_parse_errors(self): + mcp = Mcp() + try: + for token in ('"\\q"', '"\\u12"', '"\\uZZZZ"', '"raw\tcontrol"', + '01', '-01', '1.', '1e', '1e+'): + with self.subTest(token=token): + reply = mcp.send('{"jsonrpc":"2.0","id":' + token + + ',"method":"ping"}') + self.assertIsNone(reply["id"]) + self.assertEqual(reply["error"]["code"], -32700) + self.assertEqual(mcp.call("ping")["result"], {}) + finally: + mcp.close() + def test_version(self): p = subprocess.run([str(GAZE), "--version"], capture_output=True, text=True) self.assertEqual(p.returncode, 0)