1 Commits
Author SHA1 Message Date
jdb ac2edf230d Backport support for aud list values into 0.3.x (support live-budget). 2025-12-31 15:07:26 -06:00
9 changed files with 152 additions and 462 deletions
-18
View File
@@ -19,24 +19,6 @@ 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.
+4 -4
View File
@@ -1,8 +1,8 @@
# Package # Package
version = "0.5.0" version = "0.3.1"
author = "Jonathan Bernard" author = "Jonathan Bernard"
description = "Jonathan's opinionated extensions and auth layer for Mummy." description = "Jonathan's opinionated extensions and auth layer for Jester."
license = "MIT" license = "MIT"
srcDir = "src" srcDir = "src"
@@ -12,10 +12,10 @@ srcDir = "src"
requires "nim >= 1.6.2" requires "nim >= 1.6.2"
# from standard nimble repo # from standard nimble repo
requires @["bcrypt", "mummy", "uuids", "webby"] requires @["bcrypt", "jester >= 0.5.0", "uuids"]
# from https://git.jdb-software.com/jdb/nim-packages # from https://git.jdb-software.com/jdb/nim-packages
requires @["jwt_full >= 0.5.0"] requires @["jwt_full >= 0.2.0", "namespaced_logging >= 0.3.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"
+4 -1
View File
@@ -1,2 +1,5 @@
import buffoonery/[apierror, apiutils, auth, jsonutils] import buffoonery/apierror,
buffoonery/apiutils,
buffoonery/auth,
buffoonery/jsonutils
export apierror, apiutils, auth, jsonutils export apierror, apiutils, auth, jsonutils
+7 -42
View File
@@ -1,63 +1,28 @@
import std/[httpcore, json, options, strutils] from strutils import isEmptyOrWhitespace
from httpclient 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: proc raiseApiError*(respCode: HttpCode, respMsg: string, msg = "") =
result.trace &= parent.trace
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*(parent: ref Exception, respCode: HttpCode, 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 = respMsg,
if respMsg.isEmptyOrWhitespace: parent.msg msg = msg)
else: respMsg,
respData = respData,
logData = logData,
msg =
if msg.isEmptyOrWhitespace: parent.msg
else: msg)
raise apiError raise apiError
+87 -129
View File
@@ -1,24 +1,26 @@
import std/[json, jsonutils, options, sequtils, strtabs, strutils] import std/[json, jsonutils, logging, options, strutils, sequtils]
import mummy, webby, uuids import jester, namespaced_logging
import std/httpcore except HttpHeaders
import ./apierror import ./apierror
const CONTENT_TYPE_JSON* = "application/json" const CONTENT_TYPE_JSON* = "application/json"
var logNs {.threadvar.}: LoggingNamespace
template log(): untyped =
if logNs.isNil: logNs = getLoggerForNamespace("buffoonery/apiutils", lvlDebug)
logNs
## Response Utilities ## Response Utilities
## ------------------ ## ------------------
type type ApiResponse*[T] = object
ApiResponse*[T] = object details*: Option[string]
details*: Option[string] data*: Option[T]
data*: Option[T] nextOffset*: Option[int]
nextOffset*: Option[int] totalItems*: Option[int]
totalItems*: Option[int] nextLink*: Option[string]
nextLink*: Option[string] prevLink*: Option[string]
prevLink*: Option[string]
func initApiResponse*[T]( func initApiResponse*[T](
details = none[string](), details = none[string](),
@@ -30,8 +32,7 @@ func initApiResponse*[T](
ApiResponse[T](details: details, data: data, nextOffset: nextOffset, ApiResponse[T](details: details, data: data, nextOffset: nextOffset,
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
@@ -40,132 +41,89 @@ proc `%`*(r: ApiResponse): JsonNode =
if r.nextLink.isSome: result["nextLink"] = %r.nextLink if r.nextLink.isSome: result["nextLink"] = %r.nextLink
if r.prevLink.isSome: result["prevLink"] = %r.prevLink if r.prevLink.isSome: result["prevLink"] = %r.prevLink
template halt*(
code: HttpCode,
headers: RawHeaders,
content: string) =
## Immediately replies with the specified request. This means any further
## code will not be executed after calling this template in the current
## route.
bind TCActionSend, newHttpHeaders
result[0] = CallbackAction.TCActionSend
result[1] = code
result[2] = if isSome(result[2]): some(result[2].get & headers)
else: some(headers)
result[3] = content
result.matched = true
break allRoutes
func `$`*(r: ApiResponse): string = $(%r) template sendJsonResp*(
body: JsonNode,
code: HttpCode = Http200,
knownOrigins: seq[string] = @[],
headersToSend: RawHeaders = @{:}) =
## Immediately send a JSON response and stop processing the request.
let reqOrigin =
if headers(request).hasKey("Origin"): $(headers(request)["Origin"])
else: ""
let corsHeaders =
proc makeCorsHeaders*( if knownOrigins.contains(reqOrigin):
allowedMethods: seq[string],
allowedOrigins: seq[string],
allowedHeaders: Option[seq[string]],
reqOrigin = none[string]()): HttpHeaders =
result =
if reqOrigin.isSome and allowedOrigins.contains(reqOrigin.get):
@{ @{
"Access-Control-Allow-Origin": reqOrigin.get, "Access-Control-Allow-Origin": reqOrigin,
"Access-Control-Allow-Credentials": "true", "Access-Control-Allow-Credentials": "true",
"Access-Control-Allow-Methods": allowedMethods.join(","), "Access-Control-Allow-Methods": $(reqMethod(request)),
"Access-Control-Allow-Headers": "Access-Control-Allow-Headers": "Authorization,X-CSRF-TOKEN"
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: log().debug "Unrecognized Origin '" & reqOrigin & "', excluding CORS headers."
@{"X-Invalid-Origin-Details": "Unrecognized origin '" & reqOrigin.get & "'."} @{:}
else:
@{"X-Invalid-Origin-Details": "Missing Origin."}
halt(
code,
cast[RawHeaders](headersToSend) & corsHeaders & @{
"Content-Type": CONTENT_TYPE_JSON,
"Cache-Control": "no-cache"
},
$body
)
proc makeCorsHeaders*( template sendResp*[T](
allowedMethods: seq[HttpMethod],
allowedOrigins: seq[string],
allowedHeaders: Option[seq[string]],
reqOrigin = none[string]()): HttpHeaders =
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] =
if req.headers.contains("Origin"): some(req.headers["Origin"])
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*(
req: Request,
body: JsonNode,
code = Http200,
allowedOrigins = newSeq[string](),
headersToSend: HttpHeaders = @{:}) =
## Immediately send a JSON response and stop processing the request.
var headers =
headersToSend &
makeCorsHeaders(@[req.httpMethod], allowedOrigins, req.origin) &
@[("Content-Type", CONTENT_TYPE_JSON),
("Cache-Control", "no-cache")]
req.respond(code.ord, headers, $body)
proc respond*[T](
req: Request,
resp: ApiResponse[T], resp: ApiResponse[T],
code = Http200, code = Http200,
allowedOrigins = newSeq[string](), allowedOrigins = newSeq[string](),
headersToSend: HttpHeaders = @{:}) = headersToSend: RawHeaders = @{:}) =
req.respondWithRawJson(%resp, code, allowedOrigins, headersToSend) sendJsonResp(%resp, code, allowedOrigins, headersToSend)
template sendErrorResp*(err: ref ApiError, knownOrigins: seq[string]): void =
log().error err.respMsg & ( if err.msg.len > 0: ": " & err.msg else: "")
if not err.parent.isNil: log().error " original exception: " & err.parent.msg
sendJsonResp( %*{"details":err.respMsg}, err.respCode, knownOrigins)
proc respondWithData*[T]( ## CORS support
req: Request, template sendOptionsResp*(
data: T,
code = Http200,
allowedOrigins = newSeq[string](),
headersToSend: HttpHeaders = @{:}) =
req.respond(initApiResponse[T](data = some(data)),
code, allowedOrigins, headersToSend)
proc respondToOptions*(
req: Request,
allowedMethods: seq[HttpMethod], allowedMethods: seq[HttpMethod],
allowedOrigins: seq[string]) = knownOrigins: seq[string]) =
req.respond( let reqOrigin =
Http200.ord, if headers(request).hasKey("Origin"): $(headers(request)["Origin"])
makeCorsHeaders(allowedMethods, allowedOrigins, req.origin), else: ""
"")
let corsHeaders =
if knownOrigins.contains(reqOrigin):
@{
"Access-Control-Allow-Origin": reqOrigin,
"Access-Control-Allow-Credentials": "true",
"Access-Control-Allow-Methods": allowedMethods.mapIt($it).join(", "),
"Access-Control-Allow-Headers": "DNT,User-Agent,X-Requested-With,If-Modified-Since,Cache-Control,Content-Type,Range,Authorization,X-CSRF-TOKEN"
}
else:
log().debug "Unrecognized Origin '" & reqOrigin & "', excluding CORS headers."
log().debug "Valid origins: " & knownOrigins.join(", ")
@{:}
halt(
Http200,
corsHeaders,
""
)
+50 -131
View File
@@ -1,23 +1,17 @@
import std/[cookies, json, options, sequtils, strtabs, strutils, tables, times] import std/httpclient, std/json, std/logging, std/options, std/sequtils,
import mummy, uuids, webby std/strutils, std/tables, std/times
import std/httpclient except HttpHeaders import jester, namespaced_logging
import jwt_full, jwt_full/encoding import jwt_full, jwt_full/encoding
import ./[apiutils,jsonutils] import ./jsonutils
const SUPPORTED_SIGNATURE_ALGORITHMS = @[ HS256, RS256 ] const SUPPORTED_SIGNATURE_ALGORITHMS = @[ HS256, RS256 ]
type type
AuthError* = object of CatchableError AuthError* = object of CatchableError
additionalInfo*: Option[TableRef[string, JsonNode]]
ApiAuthContext* = object ApiAuthContext* = ref 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
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
issuer*: string ## The JWT issuer for tokens created by this API issuer*: string ## The JWT issuer for tokens created by this API
@@ -29,35 +23,22 @@ 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: Table[string, JwkSet] issuerKeys: TableRef[string, JwkSet]
var logNs {.threadvar.}: LoggingNamespace
proc failAuth*[T]( template log(): untyped =
reason: string, if logNs.isNil: logNs = getLoggerForNamespace("buffoonery/auth", lvlDebug)
additionalInfo: TableRef[string, T], logNs
parentException: ref Exception = nil) =
## Syntactic sugar to raise an AuthError. Reason will be the exception
## message and should be considered an internal message.
let err = newException(AuthError, reason, parentException)
err.additionalInfo = some(newTable[string, JsonNode]())
for key, val in additionalInfo.pairs:
err.additionalInfo.get[key] = %val
raise err
proc failAuth*(reason: string, parentException: ref Exception = nil) = proc failAuth*(reason: string, parentException: ref Exception = nil) =
## Syntactic sugar to raise an AuthError. Reason will be the exception
## message and should be considered an internal message.
raise newException(AuthError, reason, parentException) raise newException(AuthError, reason, parentException)
proc validateSigningKey(k: JWK): void = proc validateSigningKey(k: JWK): void =
if k.alg.isNone: failAuth "JWK is missing 'alg'" if k.alg.isNone: failAuth "JWK is missing 'alg'"
if k.kid.isNone: failAuth "JWK is missing 'kid'" if k.kid.isNone: failAuth "JWK is missing 'kid'"
proc initApiAuthContext*( proc initApiAuthContext*(
appDomain: string,
cookiePrefix: string, cookiePrefix: string,
validAudiences: seq[string], validAudiences: seq[string],
issuer: string, issuer: string,
@@ -68,52 +49,46 @@ proc initApiAuthContext*(
for k in signingKeys: validateSigningKey(k) for k in signingKeys: validateSigningKey(k)
result = ApiAuthContext( result = ApiAuthContext(
appDomain: appDomain,
cookiePrefix: cookiePrefix, cookiePrefix: cookiePrefix,
validAudiences: validAudiences, validAudiences: validAudiences,
issuer: issuer, issuer: issuer,
trustedIssuers: trustedIssuers, trustedIssuers: trustedIssuers,
signingKid: signingKid, signingKid: signingKid,
issuerKeys: toTable([(issuer, signingKeys)])) issuerKeys: newTable[string, JwkSet]([(issuer, signingKeys)]))
proc fetchJWKs(openIdConfigUrl: string): JwkSet {.gcsafe.} = proc fetchJWKs(openIdConfigUrl: string): JwkSet {.gcsafe.} =
## Fetch signing keys for an OAuth issuer. `openIdConfigUrl` is expected to ## Fetch signing keys for an OAuth issuer. `openIdConfigUrl` is expected to
## be a well-known URL (ISSUER_BASE/.well-known/openid-configuration) ## be a well-known URL (ISSUER_BASE/.well-known/openid-configuration)
var jwksKeysURI: string
try: try:
let http = newHttpClient() let http = newHttpClient()
# Inspect the OAuth metadata via the well-known address. # Inspect the OAuth metadata via the well-known address.
log.debug "fetchJwks: Fetching metadata from " & openIdConfigUrl
let metadata = parseJson(http.getContent(openIdConfigUrl)) let metadata = parseJson(http.getContent(openIdConfigUrl))
# Fetch the keys from the jwk_keys URI. # Fetch the keys from the jwk_keys URI.
jwksKeysURI = metadata.getOrFail("jwks_uri").getStr let jwksKeysURI = metadata.getOrFail("jwks_uri").getStr
debug "fetchJwks: Fetching JWKs from " & jwksKeysURI
let jwksKeys = parseJson(http.getContent(jwksKeysURI)) let jwksKeys = parseJson(http.getContent(jwksKeysURI))
# Parse and load the keys provided. # Parse and load the keys provided.
return initJwkSet(jwksKeys) return initJwkSet(jwksKeys)
except Exception: except:
#failAuth "unable to fetch isser signing keys" log.error "unable to fetch issuer signing keys: " & getCurrentExceptionMsg()
failAuth( failAuth "unable to fetch isser signing keys"
reason = "unable to fetch isser signing keys",
additionalInfo = newTable[string, string]({
"openIdConfigUrl": openIdConfigUrl,
"jwksKeysURI": jwksKeysURI }),
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:
log.error "unable to add a set of signing keys: " & getCurrentExceptionMsg()
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.
## ##
@@ -123,8 +98,10 @@ proc findSigningKey*(ctx: var ApiAuthContext, jwt: JWT, allowFetch = true): JWK
## [OpenID Connect standard discovery mechanism](https://openid.net/specs/openid-connect-discovery-1_0.html) ## [OpenID Connect standard discovery mechanism](https://openid.net/specs/openid-connect-discovery-1_0.html)
try: try:
if jwt.claims.iss.isNone: failAuth "JWT is missing 'iss' claim." if jwt.claims.iss.isNone: failAuth "Missing 'iss' claim."
if jwt.header.kid.isNone: failAuth "JWT is missing 'kid' header." if jwt.header.kid.isNone: failAuth "Missing 'kid' header."
if ctx.issuerKeys.isNil: ctx.issuerKeys = newTable[string, JwkSet]()
let jwtIssuer = jwt.claims.iss.get let jwtIssuer = jwt.claims.iss.get
@@ -141,35 +118,27 @@ proc findSigningKey*(ctx: var ApiAuthContext, jwt: JWT, allowFetch = true): JWK
fetchJWKs(jwtIssuer & "/.well-known/openid-configuration") fetchJWKs(jwtIssuer & "/.well-known/openid-configuration")
return ctx.findSigningKey(jwt, false) return ctx.findSigningKey(jwt, false)
failAuth( failAuth "unable to find JWT signing key"
reason = "unable to find JWT signing key",
additionalInfo = newTable[string, string]({
"jwtIssuer": jwtIssuer,
"jwtKid": jwt.header.kid.get}))
except: except:
log.error "unable to find JWT signing key: " & getCurrentExceptionMsg()
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:
if jwt.claims.iss.isNone: failAuth "JWT is missing 'iss' claim." log.debug "Validating JWT: " & $jwt
if jwt.claims.iss.isNone: failAuth "Missing 'iss' claim."
let jwtIssuer = jwt.claims.iss.get let jwtIssuer = jwt.claims.iss.get
if not ctx.trustedIssuers.contains(jwtIssuer): if not ctx.trustedIssuers.contains(jwtIssuer):
failAuth( failAuth "JWT is issued by $# but we only trust $#" %
reason = "We don't trust the JWT's issuer.", [jwtIssuer, $ctx.trustedIssuers]
additionalInfo = newTable[string, JsonNode]({
"issuer": %jwtIssuer,
"trustedIssuers": %ctx.trustedIssuers}))
if jwt.header.alg.isNone: failAuth "JWT is missing 'alg' header property." if jwt.header.alg.isNone: failAuth "Missing 'alg' header property."
if jwt.claims.aud.isNone: failAuth "JWT is missing 'aud' claim."
if jwt.claims.iss.isNone: failAuth "JWT is missing 'iss' claim." if jwt.claims.aud.isNone: failAuth "Missing 'aud' claim."
if jwt.claims.sub.isNone: failAuth "JWT is missing 'sub' claim."
if jwt.claims.exp.isNone: failAuth "JWT is missing or invalid 'exp' claim."
if jwt.claims["aud"].get.kind == JString: if jwt.claims["aud"].get.kind == JString:
# If the token is for a single audience, check that it is for us. # If the token is for a single audience, check that it is for us.
@@ -194,8 +163,7 @@ proc validateJWT*(ctx: var ApiAuthContext, jwt: JWT) =
except: except:
failAuth(getCurrentExceptionMsg(), getCurrentException()) failAuth(getCurrentExceptionMsg(), getCurrentException())
proc extractValidJwt*(ctx: ApiAuthContext, req: Request): 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.
@@ -209,7 +177,7 @@ proc extractValidJwt*(ctx: var ApiAuthContext, req: Request, validateCsrf = true
## - Split JWT via two cookies: ## - Split JWT via two cookies:
## ##
## - `${cookiePrefix}-user`: Contains the JWT header and payload, but not the ## - `${cookiePrefix}-user`: Contains the JWT header and payload, but not the
## signature. This cookie should be set Secure. The JWT payload should ## signature. This cookie should be set Secure. The JWT payload should
## have a defined expiration date (matching the Max-Age of the cookie) ## have a defined expiration date (matching the Max-Age of the cookie)
## and a CSRF token. This cookie is accessible by the web application. ## and a CSRF token. This cookie is accessible by the web application.
## ##
@@ -219,86 +187,39 @@ proc extractValidJwt*(ctx: var ApiAuthContext, req: Request, validateCsrf = true
## ##
## In the split-cookie mode we also check that the `csrfToken` claim in the ## In the split-cookie mode we also check that the `csrfToken` claim in the
## JWT payload matches the CSRF value passed via the `X-CSRF-TOKEN` header. ## JWT payload matches the CSRF value passed via the `X-CSRF-TOKEN` header.
## This CSRF check can be disabled by setting `validateCsrf` to `false`.
## This option is proivded to support occasional use-cases where you want to
## be able to serve a request using cookie auth when the client can't set
## custom headers (e.g. a simple link from an <a> tag). Obviously, this is a
## security risk and should only be used with caution with a full
## understanding of the risk.
try: try:
if req.headers.contains("Authorization"): if headers(req).hasKey("Authorization"):
# Using a Bearer token. # Using a Bearer token.
result = toJWT(req.headers["Authorization"][7..^1]) result = toJWT(headers(req)["Authorization"][7..^1])
elif req.headers.contains("Cookie"): else:
# Using a user/session cookie pair # Using a user/session cookie pair
let userCookieName = ctx.cookiePrefix & "-user" let userCookieName = ctx.cookiePrefix & "-user"
let sessionCookieName = ctx.cookiePrefix & "-session" let sessionCookieName = ctx.cookiePrefix & "-session"
let cookies = parseCookies(req.headers["Cookie"]) if not cookies(req).hasKey(userCookieName):
if not cookies.contains(userCookieName):
failAuth "missing cookie '$#'" % userCookieName failAuth "missing cookie '$#'" % userCookieName
if not cookies.contains(sessionCookieName): if not cookies(req).hasKey(sessionCookieName):
failAuth "missing cookie '$#'" % sessionCookieName failAuth "missing cookie '$#'" % sessionCookieName
let userVal = cookies[userCookieName] let userVal = cookies(req)[userCookieName]
let sessionVal = cookies[sessionCookieName] let sessionVal = cookies(req)[sessionCookieName]
result = toJWT(userVal & "." & sessionVal) result = toJWT(userVal & "." & sessionVal)
# Because this is a web session, check that the CSRF is present and # Because this is a web session, check that the CSRF is present and
# matches. # matches.
if validateCsrf: if not headers(req).hasKey("X-CSRF-TOKEN") or
if not req.headers.contains("X-CSRF-TOKEN") or not result.claims["csrfToken"].isSome:
not result.claims["csrfToken"].isSome: failAuth "missing CSRF token"
failAuth "missing CSRF token"
if req.headers["X-CSRF-TOKEN"] != result.claims["csrfToken"].get.getStr(""): if headers(req)["X-CSRF-TOKEN"] != result.claims["csrfToken"].get.getStr(""):
failAuth( failAuth "invalid CSRF token"
reason = "invalid CSRF token",
additionalInfo = newTable[string, string]({
"header": req.headers["X-CSRF-TOKEN"],
"jwt": result.claims["csrfToken"].get.getStr("")}))
else: failAuth "no auth token, no Authorization or Cookie headers" ctx.validateJwt(result)
ctx.validateJWT(result)
except: except:
failAuth(getCurrentExceptionMsg(), getCurrentException()) failAuth(getCurrentExceptionMsg(), getCurrentException())
proc createSessionCookies*(ctx: ApiAuthContext, jwt: JWT): HttpHeaders =
# Split the token to get the user and session cookie values.
let strToken = $jwt
let splitToken = strToken.rsplit('.', 1)
# User cookie (accessible by the application)
let userCookie = setCookie(
key = ctx.cookiePrefix & "-user",
value = splitToken[0],
domain = ctx.appDomain,
expires = jwt.claims.exp.get.utc,
httpOnly = false,
path = "/",
sameSite = SameSite.Strict,
secure = true)
# Session cookie (used by the API)
let sessionCookie = setCookie(
key = ctx.cookiePrefix & "-session",
value = splitToken[1],
domain = ctx.appDomain,
httpOnly = true,
path = "/",
sameSite = SameSite.Strict,
secure = true)
for c in [userCookie, sessionCookie]:
let parts = c.split(": ")
result &= [(parts[0], parts[1])]
proc createSignedJWT*(ctx: ApiAuthContext, claims: JsonNode, kid: string): JWT = proc createSignedJWT*(ctx: ApiAuthContext, claims: JsonNode, kid: string): JWT =
## Given a set of claims, create a JWT using the given key for our issuer ## Given a set of claims, create a JWT using the given key for our issuer
## (as defined in the ApiAuthContext). This is an opinionated method that ## (as defined in the ApiAuthContext). This is an opinionated method that
@@ -320,12 +241,10 @@ proc createSignedJWT*(ctx: ApiAuthContext, claims: JsonNode, kid: string): JWT =
initJwtClaims(claims), initJwtClaims(claims),
sigKey) sigKey)
proc newApiAccessToken*(ctx: ApiAuthContext, sub: string, duration = 1.hours): JWT = proc newApiAccessToken*(ctx: ApiAuthContext, sub: string, duration = 1.hours): JWT =
## Create a new JWT for API access. ## Create a new JWT for API access.
result = ctx.createSignedJWT( result = ctx.createSignedJWT(
%*{ %*{
"sid": $genUUID(),
"sub": sub, "sub": sub,
"iss": ctx.issuer, "iss": ctx.issuer,
"iat": now().utc.toTime.toUnix.int, "iat": now().utc.toTime.toUnix.int,
-5
View File
@@ -3,7 +3,6 @@ import json, times, timeutils, uuids
const MONTH_FORMAT* = "YYYY-MM" const MONTH_FORMAT* = "YYYY-MM"
func getOrFail*(n: JsonNode, key: string): JsonNode = func getOrFail*(n: JsonNode, key: string): JsonNode =
## convenience method to get a key from a JObject or raise an exception ## convenience method to get a key from a JObject or raise an exception
if not n.hasKey(key): if not n.hasKey(key):
@@ -11,18 +10,14 @@ func getOrFail*(n: JsonNode, key: string): JsonNode =
return n[key] return n[key]
func parseUUID*(n: JsonNode, key: string): UUID = func parseUUID*(n: JsonNode, key: string): UUID =
return parseUUID(n.getOrFail(key).getStr) return parseUUID(n.getOrFail(key).getStr)
proc parseIso8601*(n: JsonNode, key: string): DateTime = proc parseIso8601*(n: JsonNode, key: string): DateTime =
return parseIso8601(n.getOrFail(key).getStr) return parseIso8601(n.getOrFail(key).getStr)
proc parseMonth*(n: JsonNode, key: string): DateTime = proc parseMonth*(n: JsonNode, key: string): DateTime =
return parse(n.getOrFail(key).getStr, MONTH_FORMAT) return parse(n.getOrFail(key).getStr, MONTH_FORMAT)
func formatMonth*(dt: DateTime): string = func formatMonth*(dt: DateTime): string =
return dt.format(MONTH_FORMAT) return dt.format(MONTH_FORMAT)
-1
View File
@@ -1,3 +1,2 @@
switch("path", "../src") switch("path", "../src")
switch("verbosity", "0") switch("verbosity", "0")
switch("threads", "on")
-131
View File
@@ -1,131 +0,0 @@
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"