Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 10 additions & 2 deletions generate_lsquic_ffi.nim
Original file line number Diff line number Diff line change
Expand Up @@ -7,21 +7,29 @@ from os import parentDir, `/`

import boringssl

proc dropGeneratedCEnumsImpl(opirOutput: JsonNode): JsonNode =
const preludeTypes = ["struct_lsquic_cid", "lsquic_cid_t"]

proc dropPreludeTypesAndGeneratedCEnumsImpl(opirOutput: JsonNode): JsonNode =
var resp = newJArray()
for node in opirOutput:
# enums are generated manually to avoid issue described in
# https://github.com/PMunch/futhark/issues/152
if node{"kind"}.getStr("") == "enum":
continue

# Futhark incorrectly maps the struct alignment on lsquic_cid_t to field alignment,
# so the prelude supplies the correct native layout instead.
if node{"name"}.getStr("") in preludeTypes:
continue

resp.add node
resp

importc:
outputPath currentSourcePath.parentDir / "tmp_lsquic_ffi.nim"
path currentSourcePath.parentDir / "libs/lsquic/include"
addopircallback proc(opirOutput: JsonNode): JsonNode {.closure.} =
dropGeneratedCEnumsImpl(opirOutput)
dropPreludeTypesAndGeneratedCEnumsImpl(opirOutput)
rename FILE, CFile # Rename `FILE` that STB uses to `CFile` which is the Nim equivalent
rename struct_sockaddr, SockAddr # Rename `struct_sockaddr` for chronos SockAddr
"lsquic.h"
6 changes: 6 additions & 0 deletions lsquic/context/client.nim
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,7 @@ method dial*(
return err("could not dial: " & $remote)

quicClientConn.lsquicConn = conn
ctx.trackConnectionCid(conn)
ctx.processWhenReady()

ok(quicClientConn)
Expand All @@ -121,6 +122,7 @@ proc new*(T: typedesc[ClientContext], tlsConfig: TLSConfig): Result[T, string] =
ctx.tlsConfig = tlsConfig
ctx.running = true
ctx.setupSSLContext()
ctx.initCidTracking()

lsquic_engine_init_settings(addr ctx.settings, 0)
ctx.settings.es_versions = 1.cuint shl LSQVER_I001.cuint #IETF QUIC v1
Expand Down Expand Up @@ -155,6 +157,10 @@ proc new*(T: typedesc[ClientContext], tlsConfig: TLSConfig): Result[T, string] =
ea_get_ssl_ctx: getSSLCtx,
ea_packets_out: sendPacketsOut,
)
ctx.api.ea_new_scids = addCids
ctx.api.ea_live_scids = addCids
ctx.api.ea_old_scids = removeCids
ctx.api.ea_cids_update_ctx = cast[pointer](ctx)

ctx.engine = lsquic_engine_new(0, addr ctx.api)
if ctx.engine.isNil:
Expand Down
103 changes: 91 additions & 12 deletions lsquic/context/context.nim
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 OR MIT
# Copyright (c) Status Research & Development GmbH

import std/deques
import std/[deques, hashes, sets, strutils]
import boringssl
import chronos
import chronos/osdefs
Expand All @@ -12,17 +12,96 @@ import
let SSL_CTX_ID = SSL_CTX_get_ex_new_index(0, nil, nil, nil, nil) # Yes, this is global
doAssert SSL_CTX_ID >= 0, "could not generate global ssl_ctx id"

type QuicContext* = ref object of RootObj
settings*: struct_lsquic_engine_settings
api*: struct_lsquic_engine_api
engine*: ptr struct_lsquic_engine
stream_if*: struct_lsquic_stream_if
tlsConfig*: TLSConfig
tickTimeout*: Timeout
sslCtx*: ptr SSL_CTX
fd*: cint
processing: bool
running*: bool
type
CidKey* = object
len*: uint8
bytes*: array[MAX_CID_LEN, uint8]

LsquicCidArray = UncheckedArray[lsquic_cid_t]

QuicContext* = ref object of RootObj
settings*: struct_lsquic_engine_settings
api*: struct_lsquic_engine_api
engine*: ptr struct_lsquic_engine
stream_if*: struct_lsquic_stream_if
tlsConfig*: TLSConfig
tickTimeout*: Timeout
sslCtx*: ptr SSL_CTX
fd*: cint
processing: bool
running*: bool
ownedCids*: HashSet[CidKey]
Comment thread
richard-ramos marked this conversation as resolved.
Outdated

func hash*(cid: CidKey): Hash =
var h = hash(cid.len)
for i in 0 ..< cid.len.int:
h = h !& hash(cid.bytes[i])
!$h

func shortLog*(cid: CidKey): string =
var ret = $cid.len & ":"
for i in 0 ..< min(cid.len.int, 8):
ret.add(toHex(cid.bytes[i], 2))
ret

chronicles.formatIt(CidKey):
shortLog(it)

func toCidKey(cid: lsquic_cid_t, key: var CidKey): bool =
if cid.len == 0 or cid.len.int > MAX_CID_LEN:
return false

key = CidKey(len: cid.len)
for i in 0 ..< cid.len.int:
key.bytes[i] = cid.buf[i]
true

proc initCidTracking*(ctx: QuicContext) {.raises: [].} =
ctx.ownedCids = initHashSet[CidKey]()

proc addCids*(
ctx: pointer, peerCtxs: ptr pointer, cids: ptr lsquic_cid_t, nCids: cuint
) {.cdecl, raises: [].} =
Comment thread
richard-ramos marked this conversation as resolved.
let quicCtx = cast[QuicContext](ctx)
if quicCtx.isNil or cids.isNil:
return

let cidsArr = cast[ptr LsquicCidArray](cids)
for i in 0 ..< nCids.int:
var key: CidKey
if toCidKey(cidsArr[i], key):
quicCtx.ownedCids.incl(key)
trace "Registered CID", cid = key, cidCount = quicCtx.ownedCids.len

proc removeCids*(
ctx: pointer, peerCtxs: ptr pointer, cids: ptr lsquic_cid_t, nCids: cuint
) {.cdecl, raises: [].} =
Comment thread
richard-ramos marked this conversation as resolved.
let quicCtx = cast[QuicContext](ctx)
if quicCtx.isNil or cids.isNil:
return

let cidsArr = cast[ptr LsquicCidArray](cids)
for i in 0 ..< nCids.int:
var key: CidKey
if toCidKey(cidsArr[i], key):
quicCtx.ownedCids.excl(key)
trace "Removed CID", cid = key, cidCount = quicCtx.ownedCids.len

proc trackConnectionCid*(ctx: QuicContext, conn: ptr lsquic_conn_t) {.raises: [].} =
if ctx.isNil or conn.isNil:
return

let cid = lsquic_conn_id(conn)
if cid.isNil:
return

var key: CidKey
if toCidKey(cid[], key):
ctx.ownedCids.incl(key)
trace "Tracked connection CID", cid = key, cidCount = ctx.ownedCids.len

proc ownsCid*(ctx: QuicContext, cid: CidKey): bool {.raises: [].} =
not ctx.isNil and cid in ctx.ownedCids

proc isRunning*(ctx: QuicContext): bool {.raises: [].} =
not ctx.isNil and ctx.running and not ctx.engine.isNil
Expand Down
2 changes: 2 additions & 0 deletions lsquic/context/io.nim
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

import chronos
import chronos/osdefs
import chronicles
import ./context
import ../[lsquic_ffi, datagram]
import ../helpers/[openarray, sequninit, transportaddr]
Expand Down Expand Up @@ -144,6 +145,7 @@ proc sendPacketsOut*(

let res = sendmsg(SocketHandle(quicCtx.fd), msg.addr, 0)
if res < 0:
trace "sendmsg failed", sent, nspecs
Comment thread
richard-ramos marked this conversation as resolved.
break

sent.inc
Expand Down
8 changes: 7 additions & 1 deletion lsquic/context/server.nim
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,9 @@ proc onNewConn(
onClose: proc() =
discard,
)
GC_ref(quicConn) # Keep it pinned until on_conn_closed is called
let serverCtx = cast[ServerContext](stream_if_ctx)
serverCtx.trackConnectionCid(conn)
GC_ref(quicConn) # Keep it pinned until on_conn_closed is called
serverCtx.incoming.putNoWait(quicConn)
cast[ptr lsquic_conn_ctx_t](quicConn)

Expand All @@ -58,6 +59,7 @@ proc new*(T: typedesc[ServerContext], tlsConfig: TLSConfig): Result[T, string] =
ctx.running = true
ctx.incoming = newAsyncQueue[QuicConnection]()
ctx.setupSSLContext()
ctx.initCidTracking()

lsquic_engine_init_settings(addr ctx.settings, LSENG_SERVER)
ctx.settings.es_versions = 1.cuint shl LSQVER_I001.cuint #IETF QUIC v1
Expand Down Expand Up @@ -92,6 +94,10 @@ proc new*(T: typedesc[ServerContext], tlsConfig: TLSConfig): Result[T, string] =
ea_get_ssl_ctx: getSSLCtx,
ea_packets_out: sendPacketsOut,
)
ctx.api.ea_new_scids = addCids
ctx.api.ea_live_scids = addCids
ctx.api.ea_old_scids = removeCids
ctx.api.ea_cids_update_ctx = cast[pointer](ctx)

ctx.engine = lsquic_engine_new(LSENG_SERVER, addr ctx.api)
if ctx.engine.isNil:
Expand Down
70 changes: 66 additions & 4 deletions lsquic/endpoint.nim
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
# Copyright (c) Status Research & Development GmbH

import chronos, chronicles, results
import ./[errors, connection, tlsconfig, datagram, connectionmanager]
import ./[errors, connection, tlsconfig, datagram, connectionmanager, lsquic_ffi]
import ./context/[server, client, context, io]

type
Expand Down Expand Up @@ -39,17 +39,79 @@ proc createClientContext(
context.fd = fd
context

proc scidLen(endpoint: QuicEndpoint): cuint {.raises: [].} =
if not endpoint.serverContext.isNil and
endpoint.serverContext.settings.es_scid_len != 0:
return endpoint.serverContext.settings.es_scid_len
if not endpoint.clientContext.isNil and
endpoint.clientContext.settings.es_scid_len != 0:
return endpoint.clientContext.settings.es_scid_len
LSQUIC_DF_SCID_LEN.cuint

proc packetDcid(
endpoint: QuicEndpoint, packet: seq[byte], cid: var CidKey
): bool {.raises: [].} =
if packet.len == 0:
return false

var cidLen: uint8
let offset = lsquic_dcid_from_packet(
unsafeAddr packet[0], packet.len.csize_t, endpoint.scidLen(), addr cidLen
)
if offset < 0:
return false

let start = offset.int
if cidLen == 0 or cidLen.int > MAX_CID_LEN or start + cidLen.int > packet.len:
return false

cid = CidKey(len: cidLen)
for i in 0 ..< cidLen.int:
cid.bytes[i] = packet[start + i]
true

func isIetfInitial(packet: seq[byte]): bool {.raises: [].} =
if packet.len == 0:
return false
(packet[0] and 0xC0'u8) == 0xC0'u8 and (packet[0] and 0x30'u8) == 0

proc receiveDatagram(
endpoint: QuicEndpoint, data: seq[byte], local, remote: TransportAddress
) {.raises: [].} =
if endpoint.isNil or endpoint.stopped:
return

# Endpoints can have both contexts; each engine needs to see the packet.
if not endpoint.clientContext.isNil:
let
hasClientContext = not endpoint.clientContext.isNil
hasServerContext = not endpoint.serverContext.isNil

var cid: CidKey
if endpoint.packetDcid(data, cid):
if hasClientContext and endpoint.clientContext.ownsCid(cid):
trace "Routing datagram to client context", cid
endpoint.clientContext.receive(Datagram(data: data), local, remote)
return

if hasServerContext and endpoint.serverContext.ownsCid(cid):
trace "Routing datagram to server context", cid
endpoint.serverContext.receive(Datagram(data: data), local, remote)
return

if hasClientContext and not hasServerContext:
endpoint.clientContext.receive(Datagram(data: data), local, remote)
if not endpoint.serverContext.isNil:
return

if hasServerContext and not hasClientContext:
endpoint.serverContext.receive(Datagram(data: data), local, remote)
return

if hasServerContext and data.isIetfInitial():
trace "Routing initial datagram with unknown CID to server context",
bytes = data.len, local, remote
endpoint.serverContext.receive(Datagram(data: data), local, remote)
return

trace "Dropping datagram with unknown CID", bytes = data.len, local, remote
Comment thread
richard-ramos marked this conversation as resolved.

proc receiveFromUdp(
endpoint: QuicEndpoint, udp: DatagramTransport, remote: TransportAddress
Expand Down
Loading
Loading