jwt.nx source
↩ module page · 231 lines · 7985 B
1// jwt.nx -- JSON Web Token sign + verify (HS256 only).
2//
3// RFC 7519 JWT + RFC 7515 JWS + RFC 7518 algorithms. Compact
4// serialisation:
5//
6// BASE64URL(header) . BASE64URL(payload) . BASE64URL(sig)
7//
8// Header: {\"alg\":\"HS256\",\"typ\":\"JWT\"}
9// Payload: arbitrary JSON -- caller decides claims
10// Sig: HMAC-SHA-256(key, \"header64.payload64\")
11//
12// HS256 only today -- RS256/ES256 require RSA/ECDSA, which we
13// haven't finished. HS256 is fine for single-service auth
14// (web app signs its own tokens, verifies its own tokens). For
15// third-party delegation (OAuth 2) the stack will need RS256
16// later via rsa.nx (pending).
17//
18// Invariants:
19// J1 Verification uses ct_memcmp for the signature compare --
20// timing-safe against forgery attempts.
21// J2 Input parse is tolerant of missing trailing padding
22// (base64url variant omits '=').
23// J3 We do NOT validate payload claims (exp / nbf / iss).
24// Caller parses the returned payload JSON and checks.
25
26import "syscalls.nx"
27import "nx_hmac.nx" // was hmac.nx -- CODE-IDENTICAL twin (49/49 stmts) on the LEGACY syscalls.nx+sha256.nx family.
28// Two files defining hmac_sha256 + main, with the expander deduping BY PATH NOT BY SYMBOL, made
29// every legacy importer a duplicate-symbol landmine for the nx_ family (debt 1785524913).
30import "base64.nx"
31import "ct.nx"
32
33const JWT_ERR_FORMAT: i64 = -1
34const JWT_ERR_SHORT: i64 = -2
35const JWT_ERR_SIG: i64 = -3
36
37// Encode raw bytes as base64url (RFC 4648 ยง5): +/ -> -_ and no '='.
38// Writes to out; returns bytes written.
39func jwt_b64url_encode(data: *u8, n: i64, out: *u8) -> i64 {
40 // Standard base64 -> then patch + to -, / to _, and strip '='.
41 let b64_len: i64 = b64_encode(data, n, out)
42 var stripped: i64 = b64_len
43 // Strip trailing '='.
44 while stripped > 0 {
45 if out[stripped - 1] != 0x3D { break }
46 stripped = stripped - 1
47 }
48 var i: i64 = 0
49 while i < stripped {
50 if out[i] == 0x2B { out[i] = 0x2D } // '+' -> '-'
51 if out[i] == 0x2F { out[i] = 0x5F } // '/' -> '_'
52 i = i + 1
53 }
54 return stripped
55}
56
57// Decode base64url. Uncensored wrapper: copy input, pad with '=',
58// swap -/_ back to +/, then call b64_decode. Not zero-alloc
59// (we need a scratch buffer for the massaged input).
60func jwt_b64url_decode(chars: *u8, n: i64, out: *u8) -> i64 {
61 // Copy with alphabet swap.
62 let scratch_len: i64 = n + 4
63 let scratch: *u8 = sys_mmap(scratch_len + 16)
64 var i: i64 = 0
65 while i < n {
66 var c: i64 = chars[i]
67 if c == 0x2D { c = 0x2B } // '-' -> '+'
68 if c == 0x5F { c = 0x2F } // '_' -> '/'
69 scratch[i] = c
70 i = i + 1
71 }
72 // Pad to multiple of 4 with '='.
73 var pad_n: i64 = n
74 while pad_n % 4 != 0 {
75 scratch[pad_n] = 0x3D
76 pad_n = pad_n + 1
77 }
78 return b64_decode(scratch, pad_n, out)
79}
80
81// Produce a JWT signed with HS256. `header_json` + `payload_json`
82// are the raw JSON texts -- caller builds them (json_emit.nx is
83// available). Returns bytes written or JWT_ERR_SHORT.
84func jwt_sign_hs256(out: *u8, cap: i64,
85 header_json: *u8, header_len: i64,
86 payload_json: *u8, payload_len: i64,
87 key: *u8, key_len: i64) -> i64 {
88 // Encode header + payload as base64url.
89 // Each base64 chunk is at most ceil(n/3)*4 bytes.
90 let h_b64_cap: i64 = header_len * 2 + 8
91 let p_b64_cap: i64 = payload_len * 2 + 8
92 let h_b64: *u8 = sys_mmap(h_b64_cap)
93 let p_b64: *u8 = sys_mmap(p_b64_cap)
94 let h_b64_len: i64 = jwt_b64url_encode(header_json, header_len, h_b64)
95 let p_b64_len: i64 = jwt_b64url_encode(payload_json, payload_len, p_b64)
96
97 // Build signing input = h_b64 + \".\" + p_b64 into a scratch.
98 let sign_input_len: i64 = h_b64_len + 1 + p_b64_len
99 let sign_input: *u8 = sys_mmap(sign_input_len + 16)
100 var i: i64 = 0
101 while i < h_b64_len {
102 sign_input[i] = h_b64[i]
103 i = i + 1
104 }
105 sign_input[h_b64_len] = 0x2E
106 i = 0
107 while i < p_b64_len {
108 sign_input[h_b64_len + 1 + i] = p_b64[i]
109 i = i + 1
110 }
111
112 // HMAC-SHA-256 -> 32-byte tag.
113 let mac: *u8 = sys_mmap(64)
114 hmac_sha256(key, key_len, sign_input, sign_input_len, mac)
115
116 // Base64url-encode the tag.
117 let sig_b64: *u8 = sys_mmap(128)
118 let sig_b64_len: i64 = jwt_b64url_encode(mac, 32, sig_b64)
119
120 // Concatenate "h.p.s" into out.
121 let total: i64 = sign_input_len + 1 + sig_b64_len
122 if cap < total { return JWT_ERR_SHORT }
123 i = 0
124 while i < sign_input_len {
125 out[i] = sign_input[i]
126 i = i + 1
127 }
128 out[sign_input_len] = 0x2E
129 i = 0
130 while i < sig_b64_len {
131 out[sign_input_len + 1 + i] = sig_b64[i]
132 i = i + 1
133 }
134 return total
135}
136
137// Verify an HS256 JWT. Returns 0 on OK, negative on failure.
138// Writes the payload JSON bytes to payload_out + payload_len_out
139// so the caller can parse claims.
140func jwt_verify_hs256(token: *u8, n: i64,
141 key: *u8, key_len: i64,
142 payload_out: *u8, payload_cap: i64,
143 payload_len_out: *i64) -> i64 {
144 // Find the two dots.
145 var dot1: i64 = -1
146 var dot2: i64 = -1
147 var i: i64 = 0
148 while i < n {
149 if token[i] == 0x2E {
150 if dot1 < 0 {
151 dot1 = i
152 } else {
153 dot2 = i
154 break
155 }
156 }
157 i = i + 1
158 }
159 if dot1 < 0 { return JWT_ERR_FORMAT }
160 if dot2 < 0 { return JWT_ERR_FORMAT }
161
162 let sign_input_len: i64 = dot2
163
164 // Recompute HMAC over header64.payload64.
165 let mac: *u8 = sys_mmap(64)
166 hmac_sha256(key, key_len, token, sign_input_len, mac)
167 let sig_b64: *u8 = sys_mmap(128)
168 let sig_b64_len: i64 = jwt_b64url_encode(mac, 32, sig_b64)
169
170 // Provided signature on the wire.
171 let provided_off: i64 = dot2 + 1
172 let provided_len: i64 = n - provided_off
173 if provided_len != sig_b64_len { return JWT_ERR_SIG }
174 // Constant-time compare.
175 if ct_memcmp(token + provided_off, sig_b64, sig_b64_len) != 0 {
176 return JWT_ERR_SIG
177 }
178
179 // Decode payload base64url into out.
180 let payload_b64_off: i64 = dot1 + 1
181 let payload_b64_len: i64 = dot2 - payload_b64_off
182 // Upper bound on decoded bytes = payload_b64_len * 3/4.
183 if payload_cap < payload_b64_len { return JWT_ERR_SHORT }
184 let decoded_len: i64 = jwt_b64url_decode(
185 token + payload_b64_off, payload_b64_len, payload_out)
186 *payload_len_out = decoded_len
187 return 0
188}
189
190// Compile-only smoke: sign + verify round trip.
191func main() -> i64 {
192 let key: *u8 = "nishi-secret-key"
193 let header: *u8 = "{\"alg\":\"HS256\",\"typ\":\"JWT\"}"
194 let payload: *u8 = "{\"sub\":\"elder\",\"exp\":9999999999}"
195
196 let token: *u8 = sys_mmap(512)
197 let token_len: i64 = jwt_sign_hs256(token, 512,
198 header, 27,
199 payload, 32,
200 key, 16)
201 if token_len <= 0 { return 1 }
202
203 // Should contain two dots.
204 var dots: i64 = 0
205 var i: i64 = 0
206 while i < token_len {
207 if token[i] == 0x2E { dots = dots + 1 }
208 i = i + 1
209 }
210 if dots != 2 { return 2 }
211
212 // Verify.
213 let payload_out: *u8 = sys_mmap(256)
214 let p_len_out: *i64 = (sys_mmap(16)) as *i64
215 let rc: i64 = jwt_verify_hs256(token, token_len,
216 key, 16,
217 payload_out, 256, p_len_out)
218 if rc != 0 { return 3 }
219 if *p_len_out != 32 { return 4 }
220 // First byte of decoded payload is '{'
221 if payload_out[0] != 0x7B { return 5 }
222
223 // Tamper last byte -> signature fails.
224 token[token_len - 1] = token[token_len - 1] ^ 1
225 if jwt_verify_hs256(token, token_len,
226 key, 16,
227 payload_out, 256, p_len_out) != JWT_ERR_SIG {
228 return 6
229 }
230 return 0
231}