Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e86100176c | ||
|
|
5e05be6d0e | ||
|
|
06ac861c20 | ||
|
|
061b0a44fc | ||
|
|
be60254227 | ||
|
|
c6d02d7db7 |
@@ -19,6 +19,24 @@ packages repository]. The url is
|
|||||||
[nimble configuration]: https://github.com/nim-lang/nimble#configuration
|
[nimble configuration]: https://github.com/nim-lang/nimble#configuration
|
||||||
[JDB Software Nim packages]: https://git.jdb-software.com/jdb/nim-packages
|
[JDB Software Nim packages]: https://git.jdb-software.com/jdb/nim-packages
|
||||||
|
|
||||||
|
## Authentication contexts and threads
|
||||||
|
|
||||||
|
`ApiAuthContext` is a value object. Assigning it copies its configuration and
|
||||||
|
signing-key cache, so each worker can own a context and refresh keys independently.
|
||||||
|
Copy it before handing it to a worker, and synchronize any access to a context
|
||||||
|
that another thread may be modifying.
|
||||||
|
|
||||||
|
Use a `var ApiAuthContext` for `addSigningKeys`, `findSigningKey`, `validateJWT`,
|
||||||
|
and `extractValidJwt`: lookup and validation can fetch and cache issuer keys.
|
||||||
|
`createSignedJWT`, `newApiAccessToken`, and `createSessionCookies` also accept
|
||||||
|
immutable contexts. Existing callers using `let` for a context that validates
|
||||||
|
tokens must switch to `var`; use `Option[ApiAuthContext]` instead of `nil` when
|
||||||
|
absence needs to be represented.
|
||||||
|
|
||||||
|
The cache stores keys directly in a `Table[string, JwkSet]`, using the value
|
||||||
|
semantics provided by `jwt_full` 0.5.0 or later. Context copies own independent
|
||||||
|
key caches, and key lookup and signing use the parsed keys directly.
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
Buffoonery is available under two licenses depending on usage.
|
Buffoonery is available under two licenses depending on usage.
|
||||||
|
|||||||
+3
-3
@@ -1,8 +1,8 @@
|
|||||||
# Package
|
# Package
|
||||||
|
|
||||||
version = "0.4.7"
|
version = "0.5.0"
|
||||||
author = "Jonathan Bernard"
|
author = "Jonathan Bernard"
|
||||||
description = "Jonathan's opinionated extensions and auth layer for Jester."
|
description = "Jonathan's opinionated extensions and auth layer for Mummy."
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
srcDir = "src"
|
srcDir = "src"
|
||||||
|
|
||||||
@@ -15,7 +15,7 @@ requires "nim >= 1.6.2"
|
|||||||
requires @["bcrypt", "mummy", "uuids", "webby"]
|
requires @["bcrypt", "mummy", "uuids", "webby"]
|
||||||
|
|
||||||
# from https://git.jdb-software.com/jdb/nim-packages
|
# from https://git.jdb-software.com/jdb/nim-packages
|
||||||
requires @["jwt_full >= 0.2.0"]
|
requires @["jwt_full >= 0.5.0"]
|
||||||
|
|
||||||
task unittest, "Runs the unit test suite.":
|
task unittest, "Runs the unit test suite.":
|
||||||
exec "nim c -r test/runner"
|
exec "nim c -r test/runner"
|
||||||
|
|||||||
@@ -1,36 +1,63 @@
|
|||||||
from strutils import isEmptyOrWhitespace
|
import std/[httpcore, json, options, strutils]
|
||||||
from httpcore import HttpCode
|
|
||||||
|
|
||||||
type ApiError* = object of CatchableError
|
type ApiError* = object of CatchableError
|
||||||
respMsg*: string
|
respMsg*: string
|
||||||
respCode*: HttpCode
|
respCode*: HttpCode
|
||||||
|
respData*: Option[JsonNode] # Optional data to include in API response
|
||||||
|
logData*: Option[JsonNode] # Optional data to include in server-side logs
|
||||||
|
|
||||||
|
|
||||||
proc newApiError*(parent: ref Exception = nil, respCode: HttpCode, respMsg: string, msg = ""): ref ApiError =
|
proc newApiError*(
|
||||||
|
parent: ref Exception = nil,
|
||||||
|
respCode: HttpCode,
|
||||||
|
respMsg: string,
|
||||||
|
respData = none[JsonNode](),
|
||||||
|
logData = none[JsonNode](),
|
||||||
|
msg = ""): ref ApiError =
|
||||||
result = newException(ApiError, msg, parent)
|
result = newException(ApiError, msg, parent)
|
||||||
result.respCode = respCode
|
result.respCode = respCode
|
||||||
result.respMsg = respMsg
|
result.respMsg = respMsg
|
||||||
|
result.respData = respData
|
||||||
|
result.logData = logData
|
||||||
|
|
||||||
if not parent.isNil:
|
if not parent.isNil:
|
||||||
result.trace &= parent.trace
|
result.trace &= parent.trace
|
||||||
|
|
||||||
|
|
||||||
proc raiseApiError*(respCode: HttpCode, respMsg: string, msg = "") =
|
proc raiseApiError*(
|
||||||
|
respCode: HttpCode,
|
||||||
|
respMsg: string,
|
||||||
|
respData = none[JsonNode](),
|
||||||
|
logData = none[JsonNode](),
|
||||||
|
msg = "") =
|
||||||
var apiError = newApiError(
|
var apiError = newApiError(
|
||||||
parent = nil,
|
parent = nil,
|
||||||
respCode = respCode,
|
respCode = respCode,
|
||||||
respMsg = respMsg,
|
respMsg = respMsg,
|
||||||
|
respData = respData,
|
||||||
|
logData = logData,
|
||||||
msg = if msg.isEmptyOrWhitespace: respMsg
|
msg = if msg.isEmptyOrWhitespace: respMsg
|
||||||
else: msg)
|
else: msg)
|
||||||
raise apiError
|
raise apiError
|
||||||
|
|
||||||
|
|
||||||
proc raiseApiError*(respCode: HttpCode, parent: ref Exception, respMsg: string = "", msg = "") =
|
proc raiseApiError*(
|
||||||
|
respCode: HttpCode,
|
||||||
|
parent: ref Exception,
|
||||||
|
respMsg: string = "",
|
||||||
|
respData = none[JsonNode](),
|
||||||
|
logData = none[JsonNode](),
|
||||||
|
msg = "") =
|
||||||
var apiError = newApiError(
|
var apiError = newApiError(
|
||||||
parent = parent,
|
parent = parent,
|
||||||
respCode = respCode,
|
respCode = respCode,
|
||||||
respMsg =
|
respMsg =
|
||||||
if respMsg.isEmptyOrWhitespace: parent.msg
|
if respMsg.isEmptyOrWhitespace: parent.msg
|
||||||
else: respMsg,
|
else: respMsg,
|
||||||
msg = msg)
|
respData = respData,
|
||||||
|
logData = logData,
|
||||||
|
msg =
|
||||||
|
if msg.isEmptyOrWhitespace: parent.msg
|
||||||
|
else: msg)
|
||||||
|
|
||||||
raise apiError
|
raise apiError
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import std/[json, jsonutils, options, sequtils, strtabs, strutils]
|
import std/[json, jsonutils, options, sequtils, strtabs, strutils]
|
||||||
import mummy, webby
|
import mummy, webby, uuids
|
||||||
|
|
||||||
import std/httpcore except HttpHeaders
|
import std/httpcore except HttpHeaders
|
||||||
|
|
||||||
@@ -31,7 +31,7 @@ func initApiResponse*[T](
|
|||||||
totalItems: totalItems, nextLink: nextLink, prevLink: prevLink)
|
totalItems: totalItems, nextLink: nextLink, prevLink: prevLink)
|
||||||
|
|
||||||
|
|
||||||
func `%`*(r: ApiResponse): JsonNode =
|
proc `%`*(r: ApiResponse): JsonNode =
|
||||||
result = newJObject()
|
result = newJObject()
|
||||||
if r.details.isSome: result["details"] = %r.details
|
if r.details.isSome: result["details"] = %r.details
|
||||||
if r.data.isSome: result["data"] = %r.data
|
if r.data.isSome: result["data"] = %r.data
|
||||||
@@ -47,6 +47,7 @@ func `$`*(r: ApiResponse): string = $(%r)
|
|||||||
proc makeCorsHeaders*(
|
proc makeCorsHeaders*(
|
||||||
allowedMethods: seq[string],
|
allowedMethods: seq[string],
|
||||||
allowedOrigins: seq[string],
|
allowedOrigins: seq[string],
|
||||||
|
allowedHeaders: Option[seq[string]],
|
||||||
reqOrigin = none[string]()): HttpHeaders =
|
reqOrigin = none[string]()): HttpHeaders =
|
||||||
|
|
||||||
result =
|
result =
|
||||||
@@ -54,8 +55,13 @@ proc makeCorsHeaders*(
|
|||||||
@{
|
@{
|
||||||
"Access-Control-Allow-Origin": reqOrigin.get,
|
"Access-Control-Allow-Origin": reqOrigin.get,
|
||||||
"Access-Control-Allow-Credentials": "true",
|
"Access-Control-Allow-Credentials": "true",
|
||||||
"Access-Control-Allow-Methods": allowedMethods.join(", "),
|
"Access-Control-Allow-Methods": allowedMethods.join(","),
|
||||||
"Access-Control-Allow-Headers": "DNT,User-Agent,X-Requested-With,If-Modified-Since,Cache-Control,Content-Type,Range,Authorization,X-CSRF-TOKEN"
|
"Access-Control-Allow-Headers":
|
||||||
|
if allowedHeaders.isSome: allowedHeaders.get.join(",")
|
||||||
|
else:
|
||||||
|
"DNT,User-Agent,X-Requested-With,If-Modified-Since," &
|
||||||
|
"Cache-Control,Content-Type,Range,Authorization,X-CSRF-TOKEN," &
|
||||||
|
"traceparent,tracestate,X-Request-ID,X-Correlation-ID",
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
if reqOrigin.isSome:
|
if reqOrigin.isSome:
|
||||||
@@ -67,8 +73,16 @@ proc makeCorsHeaders*(
|
|||||||
proc makeCorsHeaders*(
|
proc makeCorsHeaders*(
|
||||||
allowedMethods: seq[HttpMethod],
|
allowedMethods: seq[HttpMethod],
|
||||||
allowedOrigins: seq[string],
|
allowedOrigins: seq[string],
|
||||||
|
allowedHeaders: Option[seq[string]],
|
||||||
reqOrigin = none[string]()): HttpHeaders =
|
reqOrigin = none[string]()): HttpHeaders =
|
||||||
makeCorsHeaders(allowedMethods.mapIt($it), allowedOrigins, reqOrigin )
|
makeCorsHeaders(allowedMethods.mapIt($it), allowedOrigins, allowedHeaders, reqOrigin )
|
||||||
|
|
||||||
|
|
||||||
|
proc makeCorsHeaders*[T: HttpMethod or string](
|
||||||
|
allowedMethods: seq[T],
|
||||||
|
allowedOrigins: seq[string],
|
||||||
|
reqOrigin = none[string]()): HttpHeaders =
|
||||||
|
makeCorsHeaders(allowedMethods, allowedOrigins, none[seq[string]](), reqOrigin )
|
||||||
|
|
||||||
|
|
||||||
func origin*(req: Request): Option[string] =
|
func origin*(req: Request): Option[string] =
|
||||||
@@ -76,6 +90,40 @@ func origin*(req: Request): Option[string] =
|
|||||||
else: none[string]()
|
else: none[string]()
|
||||||
|
|
||||||
|
|
||||||
|
func traceparent*(req: Request): Option[string] =
|
||||||
|
## Extract the traceparent from the request headers, if present.
|
||||||
|
if req.headers.contains("traceparent"):
|
||||||
|
return some(req.headers["traceparent"])
|
||||||
|
else:
|
||||||
|
return none[string]()
|
||||||
|
|
||||||
|
|
||||||
|
proc makeTraceContextHeaders*(req: Request, traceParentId: string): HttpHeaders =
|
||||||
|
var headers = HttpHeaders(@[])
|
||||||
|
|
||||||
|
if req.headers.contains("traceparent"):
|
||||||
|
# If the traceparent header is present, we should update it with our
|
||||||
|
# parent-id.
|
||||||
|
let traceparentParts = req.headers["traceparent"].split("-")
|
||||||
|
if traceparentParts.len != 4:
|
||||||
|
headers["traceparent"] = "00-$#-$#-00" % [
|
||||||
|
replace($genUUID(), "-", ""), # trace-id
|
||||||
|
traceParentId, # parent-id
|
||||||
|
]
|
||||||
|
|
||||||
|
else:
|
||||||
|
headers["traceparent"] = "00-$#-$#-$#" % [
|
||||||
|
traceparentParts[1], # trace-id
|
||||||
|
traceParentId, # parent-id
|
||||||
|
traceparentParts[3], # flags
|
||||||
|
]
|
||||||
|
|
||||||
|
if req.headers.contains("tracestate"):
|
||||||
|
headers["tracestate"] = req.headers["tracestate"]
|
||||||
|
|
||||||
|
return headers
|
||||||
|
|
||||||
|
|
||||||
proc respondWithRawJson*(
|
proc respondWithRawJson*(
|
||||||
req: Request,
|
req: Request,
|
||||||
body: JsonNode,
|
body: JsonNode,
|
||||||
|
|||||||
+10
-10
@@ -13,7 +13,10 @@ type
|
|||||||
AuthError* = object of CatchableError
|
AuthError* = object of CatchableError
|
||||||
additionalInfo*: Option[TableRef[string, JsonNode]]
|
additionalInfo*: Option[TableRef[string, JsonNode]]
|
||||||
|
|
||||||
ApiAuthContext* = ref object
|
ApiAuthContext* = object
|
||||||
|
## An owned authentication configuration and signing-key cache. Copy before
|
||||||
|
## handing to a worker; each worker's cache updates stay local to its copy.
|
||||||
|
## Concurrent access to the same mutable context still needs synchronization.
|
||||||
appDomain*: string ## Application domain for session cookies
|
appDomain*: string ## Application domain for session cookies
|
||||||
cookiePrefix*: string ## Prefix for the user and session cookies
|
cookiePrefix*: string ## Prefix for the user and session cookies
|
||||||
validAudiences*: seq[string] ## Expected audience values for for `aud` JWT check
|
validAudiences*: seq[string] ## Expected audience values for for `aud` JWT check
|
||||||
@@ -26,7 +29,7 @@ type
|
|||||||
## must be provided either when the ApiAuthContext is initialized (see
|
## must be provided either when the ApiAuthContext is initialized (see
|
||||||
## `initApiAuthContext` or via `addSigningKeys`
|
## `initApiAuthContext` or via `addSigningKeys`
|
||||||
|
|
||||||
issuerKeys: TableRef[string, JwkSet]
|
issuerKeys: Table[string, JwkSet]
|
||||||
|
|
||||||
|
|
||||||
proc failAuth*[T](
|
proc failAuth*[T](
|
||||||
@@ -71,7 +74,7 @@ proc initApiAuthContext*(
|
|||||||
issuer: issuer,
|
issuer: issuer,
|
||||||
trustedIssuers: trustedIssuers,
|
trustedIssuers: trustedIssuers,
|
||||||
signingKid: signingKid,
|
signingKid: signingKid,
|
||||||
issuerKeys: newTable[string, JwkSet]([(issuer, signingKeys)]))
|
issuerKeys: toTable([(issuer, signingKeys)]))
|
||||||
|
|
||||||
|
|
||||||
proc fetchJWKs(openIdConfigUrl: string): JwkSet {.gcsafe.} =
|
proc fetchJWKs(openIdConfigUrl: string): JwkSet {.gcsafe.} =
|
||||||
@@ -101,17 +104,16 @@ proc fetchJWKs(openIdConfigUrl: string): JwkSet {.gcsafe.} =
|
|||||||
parentException = getCurrentException())
|
parentException = getCurrentException())
|
||||||
|
|
||||||
|
|
||||||
proc addSigningKeys*(ctx: ApiAuthContext, issuer: string, keySet: JwkSet): void =
|
proc addSigningKeys*(ctx: var ApiAuthContext, issuer: string, keySet: JwkSet): void =
|
||||||
## Manually add a set of signing keys associated with a given issuer.
|
## Manually add a set of signing keys associated with a given issuer.
|
||||||
try:
|
try:
|
||||||
for k in keySet: validateSigningKey(k)
|
for k in keySet: validateSigningKey(k)
|
||||||
if ctx.issuerKeys.isNil: ctx.issuerKeys = newTable[string, JwkSet]()
|
|
||||||
ctx.issuerKeys[issuer] = keySet
|
ctx.issuerKeys[issuer] = keySet
|
||||||
except:
|
except:
|
||||||
raise getCurrentException()
|
raise getCurrentException()
|
||||||
|
|
||||||
|
|
||||||
proc findSigningKey*(ctx: ApiAuthContext, jwt: JWT, allowFetch = true): JWK {.gcsafe.} =
|
proc findSigningKey*(ctx: var ApiAuthContext, jwt: JWT, allowFetch = true): JWK {.gcsafe.} =
|
||||||
## Lookup the signing key for a given JWT. This method assumes that you trust
|
## Lookup the signing key for a given JWT. This method assumes that you trust
|
||||||
## the issuer named in the JWT.
|
## the issuer named in the JWT.
|
||||||
##
|
##
|
||||||
@@ -124,8 +126,6 @@ proc findSigningKey*(ctx: ApiAuthContext, jwt: JWT, allowFetch = true): JWK {.gc
|
|||||||
if jwt.claims.iss.isNone: failAuth "JWT is missing 'iss' claim."
|
if jwt.claims.iss.isNone: failAuth "JWT is missing 'iss' claim."
|
||||||
if jwt.header.kid.isNone: failAuth "JWT is missing 'kid' header."
|
if jwt.header.kid.isNone: failAuth "JWT is missing 'kid' header."
|
||||||
|
|
||||||
if ctx.issuerKeys.isNil: ctx.issuerKeys = newTable[string, JwkSet]()
|
|
||||||
|
|
||||||
let jwtIssuer = jwt.claims.iss.get
|
let jwtIssuer = jwt.claims.iss.get
|
||||||
|
|
||||||
if ctx.issuerKeys.hasKey(jwtIssuer):
|
if ctx.issuerKeys.hasKey(jwtIssuer):
|
||||||
@@ -151,7 +151,7 @@ proc findSigningKey*(ctx: ApiAuthContext, jwt: JWT, allowFetch = true): JWK {.gc
|
|||||||
failAuth("unable to find JWT signing key", getCurrentException())
|
failAuth("unable to find JWT signing key", getCurrentException())
|
||||||
|
|
||||||
|
|
||||||
proc validateJWT*(ctx: ApiAuthContext, jwt: JWT) =
|
proc validateJWT*(ctx: var ApiAuthContext, jwt: JWT) =
|
||||||
## Given a JWT, validate that it is a well-formed JWT, validate the issuer's
|
## Given a JWT, validate that it is a well-formed JWT, validate the issuer's
|
||||||
## signature on the token, and validate all the claims that it preesnts.
|
## signature on the token, and validate all the claims that it preesnts.
|
||||||
try:
|
try:
|
||||||
@@ -195,7 +195,7 @@ proc validateJWT*(ctx: ApiAuthContext, jwt: JWT) =
|
|||||||
failAuth(getCurrentExceptionMsg(), getCurrentException())
|
failAuth(getCurrentExceptionMsg(), getCurrentException())
|
||||||
|
|
||||||
|
|
||||||
proc extractValidJwt*(ctx: ApiAuthContext, req: Request, validateCsrf = true): JWT =
|
proc extractValidJwt*(ctx: var ApiAuthContext, req: Request, validateCsrf = true): JWT =
|
||||||
## Extracts a valid JWT representing the user's authentication and
|
## Extracts a valid JWT representing the user's authentication and
|
||||||
## authorization details, if present. If there are no valid credentials an
|
## authorization details, if present. If there are no valid credentials an
|
||||||
## exception is raised.
|
## exception is raised.
|
||||||
|
|||||||
@@ -1,2 +1,3 @@
|
|||||||
switch("path", "../src")
|
switch("path", "../src")
|
||||||
switch("verbosity", "0")
|
switch("verbosity", "0")
|
||||||
|
switch("threads", "on")
|
||||||
|
|||||||
+131
@@ -0,0 +1,131 @@
|
|||||||
|
import std/[json, options, unittest]
|
||||||
|
import jwt_full
|
||||||
|
import buffoonery/auth
|
||||||
|
|
||||||
|
const testIssuer = "https://issuer.example"
|
||||||
|
|
||||||
|
proc keyNode(kid = "original"): JsonNode =
|
||||||
|
%*{"kty": "oct", "alg": "HS256", "kid": kid,
|
||||||
|
"k": "c2VjcmV0", "metadata": {"labels": ["original"]}}
|
||||||
|
|
||||||
|
proc newContext(): ApiAuthContext =
|
||||||
|
initApiAuthContext("example.com", "test", @[testIssuer], testIssuer,
|
||||||
|
@[testIssuer], "original", @[initJwk(keyNode())])
|
||||||
|
|
||||||
|
proc tokenFor(kid = "original", issuer = testIssuer): JWT =
|
||||||
|
# Lookup only needs the header and issuer; signing is tested separately.
|
||||||
|
createSignedJwt(initJoseHeader(%*{"alg": "HS256", "kid": kid}),
|
||||||
|
initJwtClaims(%*{"iss": issuer}), initJwk(keyNode(kid)))
|
||||||
|
|
||||||
|
when compileOption("threads"):
|
||||||
|
type WorkerInput = tuple[ctx: ApiAuthContext, ok: ptr bool]
|
||||||
|
|
||||||
|
proc useContext(input: WorkerInput) {.thread.} =
|
||||||
|
var ctx = input.ctx
|
||||||
|
try:
|
||||||
|
ctx.validAudiences[0] = "worker"
|
||||||
|
let key = ctx.findSigningKey(tokenFor(), false)
|
||||||
|
key["metadata"].get["labels"][0].str = "worker"
|
||||||
|
ctx.addSigningKeys(testIssuer, @[initJwk(keyNode("worker"))])
|
||||||
|
input.ok[] = ctx.findSigningKey(tokenFor("worker"), false).kid.get == "worker"
|
||||||
|
except CatchableError:
|
||||||
|
input.ok[] = false
|
||||||
|
|
||||||
|
suite "auth context value semantics":
|
||||||
|
test "copies own their configuration and signing-key cache":
|
||||||
|
var original = newContext()
|
||||||
|
var copied = original
|
||||||
|
copied.appDomain[0] = 'E'
|
||||||
|
copied.validAudiences[0][0] = 'H'
|
||||||
|
copied.trustedIssuers.add "https://other.example"
|
||||||
|
copied.addSigningKeys(testIssuer, @[initJwk(keyNode("replacement"))])
|
||||||
|
|
||||||
|
check original.appDomain == "example.com"
|
||||||
|
check original.validAudiences == @[testIssuer]
|
||||||
|
check original.trustedIssuers == @[testIssuer]
|
||||||
|
check original.findSigningKey(tokenFor(), false).kid.get == "original"
|
||||||
|
check copied.findSigningKey(tokenFor("replacement"), false).kid.get == "replacement"
|
||||||
|
expect AuthError:
|
||||||
|
discard copied.findSigningKey(tokenFor(), false)
|
||||||
|
expect AuthError:
|
||||||
|
discard original.findSigningKey(tokenFor("replacement"), false)
|
||||||
|
|
||||||
|
test "constructor and added keys do not retain caller-owned JSON":
|
||||||
|
let node = keyNode()
|
||||||
|
var ctx = initApiAuthContext("example.com", "test", @[testIssuer],
|
||||||
|
testIssuer, @[testIssuer], "original", @[initJwk(node)])
|
||||||
|
node["metadata"]["labels"][0].str = "changed"
|
||||||
|
check ctx.findSigningKey(tokenFor(), false)["metadata"].get["labels"][0].getStr == "original"
|
||||||
|
|
||||||
|
let added = keyNode("added")
|
||||||
|
ctx.addSigningKeys(testIssuer, @[initJwk(added)])
|
||||||
|
added["metadata"]["labels"][0].str = "changed"
|
||||||
|
check ctx.findSigningKey(tokenFor("added"), false)["metadata"].get["labels"][0].getStr == "original"
|
||||||
|
|
||||||
|
test "returned key metadata is independent of both context copies":
|
||||||
|
var original = newContext()
|
||||||
|
var copied = original
|
||||||
|
let key = copied.findSigningKey(tokenFor(), false)
|
||||||
|
key["metadata"].get["labels"][0].str = "changed"
|
||||||
|
check original.findSigningKey(tokenFor(), false)["metadata"].get["labels"][0].getStr == "original"
|
||||||
|
check copied.findSigningKey(tokenFor(), false)["metadata"].get["labels"][0].getStr == "original"
|
||||||
|
|
||||||
|
test "default contexts accept keys and reject missing keys without fetching":
|
||||||
|
var ctx: ApiAuthContext
|
||||||
|
expect AuthError:
|
||||||
|
discard ctx.findSigningKey(tokenFor(), false)
|
||||||
|
ctx.addSigningKeys(testIssuer, @[initJwk(keyNode())])
|
||||||
|
check ctx.findSigningKey(tokenFor(), false).kid.get == "original"
|
||||||
|
|
||||||
|
test "invalid key updates leave the existing cache intact":
|
||||||
|
var ctx = newContext()
|
||||||
|
let invalid = keyNode("invalid")
|
||||||
|
invalid.delete("alg")
|
||||||
|
expect AuthError:
|
||||||
|
ctx.addSigningKeys(testIssuer, @[initJwk(invalid)])
|
||||||
|
check ctx.findSigningKey(tokenFor(), false).kid.get == "original"
|
||||||
|
|
||||||
|
test "immutable contexts can sign tokens that mutable copies validate":
|
||||||
|
let original = newContext()
|
||||||
|
let token = original.newApiAccessToken("user")
|
||||||
|
var copied = original
|
||||||
|
copied.validateJWT(token)
|
||||||
|
check token.claims.sub.get == "user"
|
||||||
|
|
||||||
|
test "cached keys preserve RSA and EC variants":
|
||||||
|
var ctx: ApiAuthContext
|
||||||
|
let nodes = @[
|
||||||
|
%*{"kty": "RSA", "alg": "RS256", "kid": "rsa-public",
|
||||||
|
"n": "AQAB", "e": "AQAB"},
|
||||||
|
%*{"kty": "RSA", "alg": "RS256", "kid": "rsa-private",
|
||||||
|
"n": "AQAB", "e": "AQAB", "d": "Ag", "p": "Aw", "q": "BQ",
|
||||||
|
"oth": [{"r": "Bw", "d": "Ag", "t": "Aw"}]},
|
||||||
|
%*{"kty": "EC", "alg": "ES256", "kid": "ec-public",
|
||||||
|
"crv": "P-256", "x": "AQ", "y": "Ag"},
|
||||||
|
%*{"kty": "EC", "alg": "ES256", "kid": "ec-private",
|
||||||
|
"crv": "P-256", "x": "AQ", "y": "Ag", "d": "Aw"}]
|
||||||
|
for node in nodes:
|
||||||
|
let key = initJwk(node)
|
||||||
|
ctx.addSigningKeys(testIssuer, @[key])
|
||||||
|
let restored = ctx.findSigningKey(tokenFor(key.kid.get), false)
|
||||||
|
check restored.keyKind == key.keyKind
|
||||||
|
for name, value in node.pairs:
|
||||||
|
check restored[name].get == value
|
||||||
|
case key.keyKind
|
||||||
|
of RsaPublic: check restored.rsaPub == key.rsaPub
|
||||||
|
of RsaPrivate: check restored.rsaPrv == key.rsaPrv
|
||||||
|
of EcPublic: check restored.ecPub == key.ecPub
|
||||||
|
of EcPrivate: check restored.ecPrv == key.ecPrv
|
||||||
|
of Octet: discard
|
||||||
|
|
||||||
|
when compileOption("threads"):
|
||||||
|
test "workers mutate independent context copies":
|
||||||
|
var original = newContext()
|
||||||
|
var workers: array[2, Thread[WorkerInput]]
|
||||||
|
var succeeded: array[2, bool]
|
||||||
|
for i in 0 ..< workers.len:
|
||||||
|
createThread(workers[i], useContext, (original, addr succeeded[i]))
|
||||||
|
joinThreads(workers)
|
||||||
|
check succeeded == [true, true]
|
||||||
|
check original.validAudiences == @[testIssuer]
|
||||||
|
check original.findSigningKey(tokenFor(), false)["metadata"].get["labels"][0].getStr == "original"
|
||||||
|
|||||||
Reference in New Issue
Block a user