diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 4a961e8..62c3a88 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -25,4 +25,13 @@ jobs: if: ${{ matrix.atomic_arc }} - run: nim c -r -d:useMalloc tests/test_http2.nim - run: nim c -r -d:useMalloc tests/test_websockets.nim + - run: nim c -r -d:useMalloc -d:ssl tests/test_http.nim + if: ${{ runner.os == 'Linux' }} + - run: nim c -r -d:useMalloc -d:ssl tests/test_tls.nim + if: ${{ runner.os == 'Linux' }} + # The fuzzer creates a few thousand servers in seconds. On Windows every + # SelectEvent is a loopback socket pair, so the dynamic port range runs + # out part-way through and newServer raises (upstream's Windows job flakes + # the same way); Linux is where the fuzzer is meaningful. - run: nim c -r -d:useMalloc -d:mummyNoWorkers tests/fuzz_recv.nim + if: ${{ runner.os == 'Linux' }} diff --git a/README.md b/README.md index 271e1df..676f6ff 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,40 @@ -# Mummy +# Mummy (FrameOS fork) + +This is [FrameOS](https://github.com/FrameOS/frameos)'s fork of +[guzba/mummy](https://github.com/guzba/mummy). It adds what FrameOS needs to +serve its on-device API over HTTPS without a separate proxy process: + +- **TLS listeners.** Compile with `-d:ssl` and pass a `TlsConfig` to + `addListener`; the handshake, reads and writes run inside the same epoll + loop as plain sockets (OpenSSL via Nim's `std/openssl`, no extra bindings). + The certificate chain and private key are loaded from memory with + `newTlsConfig(certificateChainPem, privateKeyPem)`, so the key never has to + be written to disk. TLS 1.2 is the minimum version. +- **Several listeners per server**, plain and TLS side by side, added with + `addListener(port, address, tls)` and removed with `removeListener` — from + any thread, before or while serving. `serve()` with no arguments serves on + all of them; `serve(port, address)` still works as before. +- **`Request.secure`**, true for requests that arrived over a TLS listener. + +Everything else is upstream mummy. Pin it by commit from a nimble file: +`requires "https://github.com/FrameOS/mummy#"`. + +```nim +import mummy + +proc handler(request: Request) = + request.respond(200, emptyHttpHeaders(), "secure: " & $request.secure) + +let server = newServer(handler) +discard server.addListener(Port(8080), "0.0.0.0") +let tls = newTlsConfig(readFile("cert.pem"), readFile("key.pem")) +discard server.addListener(Port(8443), "0.0.0.0", tls) +server.serve() +``` + +`nim c --threads:on --mm:orc -d:ssl -r tls_server.nim` + +--- `nimble install mummy` @@ -225,4 +261,6 @@ Requests/sec: 9,171.55 ## Testing +The TLS listeners are covered by `nim c -r -d:ssl tests/test_tls.nim` (plain and TLS side by side, `Request.secure`, a multi-megabyte response, WebSocket over TLS, a client stalled mid-handshake, listeners added and removed while serving). + A fuzzer has been run against Mummy's socket reading and parsing code to ensure Mummy does not crash or otherwise misbehave on bad data from sockets. You can run the fuzzer any time by running `nim c -r tests/fuzz_recv.nim`. diff --git a/src/mummy.nim b/src/mummy.nim index 9f48078..9d69c92 100644 --- a/src/mummy.nim +++ b/src/mummy.nim @@ -30,6 +30,32 @@ elif defined(posix): import std/locks +when defined(ssl): + # TLS listeners terminate connections with OpenSSL inside the same epoll + # loop that serves plain sockets. Nim's stdlib wrapper already binds the + # whole server side (TLS_server_method, SSL_accept/read/write/pending, + # SSL_get_error, SSL_CTX_set_mode); the handful of symbols it lacks are + # declared below, bound the way the wrapper binds everything else: through + # its DLLSSLName / DLLUtilName library patterns (stock Nim loads OpenSSL + # at run time; a Nim built to link -lssl still resolves these). + import std/openssl + + proc PEM_read_bio_X509( + bp: BIO, x: ptr PX509, cb: pointer, u: pointer + ): PX509 {.cdecl, dynlib: DLLUtilName, importc.} + proc SSL_CTX_use_certificate( + ctx: SslCtx, x: PX509 + ): cint {.cdecl, dynlib: DLLSSLName, importc.} + proc SSL_CTX_use_PrivateKey( + ctx: SslCtx, pkey: EVP_PKEY + ): cint {.cdecl, dynlib: DLLSSLName, importc.} + proc SSL_get_version(ssl: SslPtr): cstring {.cdecl, dynlib: DLLSSLName, importc.} + proc mummyX509Free(cert: PX509) {.cdecl, dynlib: DLLUtilName, importc: "X509_free".} + + const + SSL_CTRL_SET_MIN_PROTO_VERSION = 123 + TLS1_2_VERSION = 0x0303 + export Port, common, httpheaders, queryparams const @@ -53,6 +79,7 @@ type headers*: HttpHeaders ## HTTP headers key-value pairs. body*: string ## Request body. remoteAddress*: string ## Network address of the request sender. + secure*: bool ## True when the request arrived over a TLS listener. server: Server clientSocket: SocketHandle clientId: uint64 @@ -83,6 +110,31 @@ type message: Message ) {.gcsafe.} + TlsConfig* {.acyclic.} = ref object + ## Server certificate and key, loaded once and shared by any number of + ## TLS listeners. Create with `newTlsConfig`. Requires `-d:ssl`. + when defined(ssl): + ctx: SslCtx + + Listener* {.acyclic.} = ref object + ## A bound and listening socket the server accepts connections from. + ## Returned by `addListener`, handed back to `removeListener`. + ## Acyclic, like the other refs that cross threads here (DataEntry, + ## OutgoingBuffer): a listener is created on the caller's thread and + ## released on the serving thread, and ORC's cycle-candidate roots are + ## per thread — a ref registered as a root on one thread and freed on + ## another crashes in unregisterCycle. + id: int + socket: SocketHandle + address: string + port: Port + tls: TlsConfig + + ListenerOp = object + remove: bool + listener: Listener # add + id: int # remove + ServerObj = object handler: RequestHandler websocketHandler: WebSocketHandler @@ -93,7 +145,10 @@ type workerThreads: seq[Thread[Server]] serving: Atomic[bool] destroyCalled: bool - socket: SocketHandle + listeners: seq[Listener] # Only touched by the serving thread + listenerOps: Deque[ListenerOp] # Applied by the serving thread on wake-up + listenerOpsLock: Lock + nextListenerId: int # Under listenerOpsLock selector: Selector[DataEntry] responseQueued, sendQueued, shutdown: SelectEvent clientSockets: HashSet[SocketHandle] @@ -120,12 +175,20 @@ type DataEntry {.acyclic.} = ref object case kind: DataEntryKind: of ServerSocketEntry: - discard + listener: Listener of EventEntry: event: SelectEvent of ClientSocketEntry: clientId: uint64 remoteAddress: string + secure: bool # Accepted from a TLS listener + acceptedAt: float64 + when defined(ssl): + ssl: SslPtr # nil on plain connections + tlsHandshaken: bool + # SSL_accept or SSL_read asked for the socket to become writable + # before it can make progress; keep Write armed until it has. + tlsWantWrite: bool recvBuf: string bytesReceived: int requestState: IncomingRequestState @@ -259,6 +322,201 @@ proc setNoDelay( "Error setting TCP_NODELAY: ", e.msg ) +when defined(ssl): + proc tlsErrorText(): string = + let code = ERR_get_error() + if code == 0: + return "unknown OpenSSL error" + var buf: array[256, char] + discard ERR_error_string(code, cast[cstring](buf[0].addr)) + result = $cast[cstring](buf[0].addr) + # Drain what is left so the next error is not misattributed + while ERR_get_error() != 0: + discard + + proc newTlsConfig*( + certificateChainPem: string, + privateKeyPem: string + ): TlsConfig {.raises: [MummyError].} = + ## Loads a PEM certificate chain (leaf first) and its PEM private key from + ## memory, so the key never has to touch the file system. TLS 1.2 is the + ## minimum protocol version; ciphers are OpenSSL's defaults. + ## The returned config can back any number of TLS listeners and lives for + ## the rest of the process. + if certificateChainPem.len == 0 or privateKeyPem.len == 0: + raise newException(MummyError, "TLS certificate and key must not be empty") + + let serverMethod = + try: + TLS_server_method() + except LibraryError as e: + raise newException(MummyError, "OpenSSL is not available: " & e.msg) + let ctx = SSL_CTX_new(serverMethod) + if ctx == nil: + raise newException(MummyError, "SSL_CTX_new failed: " & tlsErrorText()) + + proc fail(ctx: SslCtx, msg: string) {.raises: [MummyError].} = + SSL_CTX_free(ctx) + raise newException(MummyError, msg) + + if SSL_CTX_ctrl( + ctx, SSL_CTRL_SET_MIN_PROTO_VERSION.cint, TLS1_2_VERSION.clong, nil + ) != 1: + fail(ctx, "Setting the minimum TLS version failed: " & tlsErrorText()) + + # Partial writes let the loop keep its plain-socket bookkeeping (bytesSent + # advances by whatever went out); the moving-buffer mode is needed because + # outgoing buffers are strings whose address may change between retries. + discard SSLCTXSetMode( + ctx, SSL_MODE_ENABLE_PARTIAL_WRITE or SSL_MODE_ACCEPT_MOVING_WRITE_BUFFER + ) + + # Certificate chain: the first PEM block is the leaf, the rest are + # intermediates handed to the context as extra chain certificates. + let certBio = BIO_new_mem_buf(certificateChainPem[0].unsafeAddr, certificateChainPem.len.cint) + if certBio == nil: + fail(ctx, "BIO_new_mem_buf failed") + var certCount = 0 + while true: + let cert = PEM_read_bio_X509(certBio, nil, nil, nil) + if cert == nil: + # The end of the PEM input also reports as an error; only a missing + # leaf is a real failure. + while ERR_get_error() != 0: + discard + break + if certCount == 0: + if SSL_CTX_use_certificate(ctx, cert) != 1: + mummyX509Free(cert) + discard BIO_free(certBio) + fail(ctx, "SSL_CTX_use_certificate failed: " & tlsErrorText()) + mummyX509Free(cert) # The context holds its own reference + else: + # SSL_CTRL_EXTRA_CHAIN_CERT takes ownership of the certificate + if SSL_CTX_ctrl(ctx, SSL_CTRL_EXTRA_CHAIN_CERT.cint, 0, cert) != 1: + mummyX509Free(cert) + discard BIO_free(certBio) + fail(ctx, "Adding an intermediate certificate failed: " & tlsErrorText()) + inc certCount + discard BIO_free(certBio) + if certCount == 0: + fail(ctx, "No certificate found in the PEM certificate chain") + + let keyBio = BIO_new_mem_buf(privateKeyPem[0].unsafeAddr, privateKeyPem.len.cint) + if keyBio == nil: + fail(ctx, "BIO_new_mem_buf failed") + let key = PEM_read_bio_PrivateKey(keyBio, nil, nil, nil) + discard BIO_free(keyBio) + if key == nil: + fail(ctx, "Reading the PEM private key failed: " & tlsErrorText()) + if SSL_CTX_use_PrivateKey(ctx, key) != 1: + EVP_PKEY_free(key) + fail(ctx, "SSL_CTX_use_PrivateKey failed: " & tlsErrorText()) + EVP_PKEY_free(key) # The context holds its own reference + + if SSL_CTX_check_private_key(ctx) != 1: + fail(ctx, "The private key does not match the certificate: " & tlsErrorText()) + + result = TlsConfig() + result.ctx = ctx + +proc port*(listener: Listener): Port = + ## The port the listener is bound to. Useful after `addListener` with + ## port 0, where the operating system picked the port. + listener.port + +proc address*(listener: Listener): string = + ## The address the listener is bound to. + listener.address + +proc secure*(listener: Listener): bool = + ## Whether connections accepted by this listener are TLS. + listener.tls != nil + +proc closeListenerSocket(listener: Listener) = + if listener.socket.int != 0: + listener.socket.close() + listener.socket = SocketHandle(0) + +proc addListener*( + server: Server, + port: Port, + address = "localhost", + tls: TlsConfig = nil +): Listener {.raises: [MummyError].} = + ## Binds a listening socket and hands it to the server. Can be called + ## before `serve()` and, from any thread, while the server is serving — + ## the serving thread starts accepting from it on its next loop iteration. + ## Pass a `TlsConfig` to terminate TLS on this listener (requires `-d:ssl`). + ## Port 0 lets the operating system choose; read it back with `listener.port`. + ## Raises if the socket cannot be bound, without touching the server. + when not defined(ssl): + if tls != nil: + raise newException(MummyError, "TLS listeners require compiling with -d:ssl") + + let listener = Listener() + listener.address = address + listener.port = port + listener.tls = tls + try: + listener.socket = createNativeSocket( + Domain.AF_INET, + SockType.SOCK_STREAM, + Protocol.IPPROTO_TCP, + false + ) + if listener.socket == osInvalidSocket: + raiseOSError(osLastError()) + + listener.socket.setBlocking(false) + listener.socket.setSockOptInt(SOL_SOCKET, SO_REUSEADDR, 1) + + let ai = getAddrInfo( + address, + port, + Domain.AF_INET, + SockType.SOCK_STREAM, + Protocol.IPPROTO_TCP, + ) + try: + if bindAddr(listener.socket, ai.ai_addr, ai.ai_addrlen.SockLen) < 0: + raiseOSError(osLastError()) + finally: + freeAddrInfo(ai) + + if nativesockets.listen(listener.socket, listenBacklogLen) < 0: + raiseOSError(osLastError()) + + if port == Port(0): + let (_, boundPort) = getLocalAddr(listener.socket, Domain.AF_INET) + listener.port = boundPort + except Exception as e: + listener.closeListenerSocket() + raise currentExceptionAsMummyError() + + withLock server.listenerOpsLock: + inc server.nextListenerId + listener.id = server.nextListenerId + server.listenerOps.addLast(ListenerOp(listener: listener)) + + # Listener changes ride the responseQueued event: the loop drains the ops + # queue whenever it wakes up for a queued response, and a trigger with + # nothing queued costs one empty pass. + if server.serving.load(moRelaxed): + server.trigger(server.responseQueued) + + listener + +proc removeListener*(server: Server, listener: Listener) {.raises: [].} = + ## Stops accepting on the listener and closes its socket. Connections it + ## already accepted are unaffected. Safe to call from any thread. + if listener == nil: + return + withLock server.listenerOpsLock: + server.listenerOps.addLast(ListenerOp(remove: true, id: listener.id)) + if server.serving.load(moRelaxed): + server.trigger(server.responseQueued) + proc send*( websocket: WebSocket, data: sink string, @@ -762,6 +1020,7 @@ proc popRequest( result.clientSocket = clientSocket result.clientId = dataEntry.clientId result.remoteAddress = dataEntry.remoteAddress + result.secure = dataEntry.secure result.httpVersion = dataEntry.requestState.httpVersion result.httpMethod = move dataEntry.requestState.httpMethod result.uri = move dataEntry.requestState.uri @@ -1106,18 +1365,42 @@ proc afterSend( return true # If we don't have any more outgoing buffers, update the selector if dataEntry.outgoingBuffers.len == 0: - server.selector.updateHandle2(clientSocket, {Read}) + var keepWrite = false + when defined(ssl): + # A TLS read or handshake step may still be waiting for writability + keepWrite = dataEntry.tlsWantWrite + if not keepWrite: + server.selector.updateHandle2(clientSocket, {Read}) proc destroy(server: Server, joinThreads: bool) {.raises: [].} = withLock server.taskQueueLock: server.destroyCalled = true + when defined(ssl): + # Free the per-connection SSL objects while the selector can still map + # a socket to its entry. + if server.selector != nil: + for clientSocket in server.clientSockets: + try: + let dataEntry = server.selector.getData(clientSocket) + if dataEntry != nil and dataEntry.kind == ClientSocketEntry and + dataEntry.ssl != nil: + SSL_free(dataEntry.ssl) + dataEntry.ssl = nil + except Exception as e: + discard # Ignore if server.selector != nil: try: server.selector.close() except Exception as e: discard # Ignore - if server.socket.int != 0: - server.socket.close() + for listener in server.listeners: + listener.closeListenerSocket() + server.listeners.setLen(0) + withLock server.listenerOpsLock: + while server.listenerOps.len > 0: + let op = server.listenerOps.popFirst() + if not op.remove: + op.listener.closeListenerSocket() for clientSocket in server.clientSockets: clientSocket.close() broadcast(server.taskQueueCond) @@ -1128,18 +1411,16 @@ proc destroy(server: Server, joinThreads: bool) {.raises: [].} = deinitLock(server.responseQueueLock) deinitLock(server.sendQueueLock) deinitLock(server.websocketQueuesLock) - try: - server.responseQueued.close() - except Exception as e: - discard # Ignore - try: - server.sendQueued.close() - except Exception as e: - discard # Ignore - try: - server.shutdown.close() - except Exception as e: - discard # Ignore + deinitLock(server.listenerOpsLock) + # newServer can fail between creating these (on Windows each is a + # loopback socket pair, and a busy process can run out of ephemeral + # ports); closing an event that was never created is a nil dereference. + for event in [server.responseQueued, server.sendQueued, server.shutdown]: + if event != nil: + try: + event.close() + except Exception as e: + discard # Ignore `=destroy`(server[]) deallocShared(server) else: @@ -1147,6 +1428,110 @@ proc destroy(server: Server, joinThreads: bool) {.raises: [].} = # The process is likely going to be exiting anyway discard +when defined(ssl): + proc tlsStep( + server: Server, + clientSocket: SocketHandle, + dataEntry: DataEntry + ): tuple[received, sent, close: bool] {.raises: [IOSelectorsException].} = + ## One step of a TLS connection, run on any Read or Write readiness: + ## finishes the handshake, drains every record OpenSSL has buffered into + ## recvBuf (one epoll wake-up can carry several), and pushes as much of + ## the head outgoing buffer as SSL_write accepts. Ends by arming the + ## selector with what OpenSSL says it needs next. + dataEntry.tlsWantWrite = false + var writeWantsRead = false + + if not dataEntry.tlsHandshaken: + let ret = SSL_accept(dataEntry.ssl) + if ret == 1: + dataEntry.tlsHandshaken = true + server.log( + DebugLevel, + "TLS handshake ", $SSL_get_version(dataEntry.ssl), " ", + $int((epochTime() - dataEntry.acceptedAt) * 1000), " ms ", + dataEntry.remoteAddress + ) + else: + case SSL_get_error(dataEntry.ssl, ret): + of SSL_ERROR_WANT_READ: + server.selector.updateHandle2(clientSocket, {Read}) + return + of SSL_ERROR_WANT_WRITE: + dataEntry.tlsWantWrite = true + server.selector.updateHandle2(clientSocket, {Read, Write}) + return + else: + server.log(DebugLevel, "TLS handshake failed: ", tlsErrorText()) + return (false, false, true) + + # Read everything OpenSSL can give us + while true: + # Expand the buffer if it is full + if dataEntry.bytesReceived == dataEntry.recvBuf.len: + dataEntry.recvBuf.setLen(dataEntry.recvBuf.len * 2) + let ret = SSL_read( + dataEntry.ssl, + dataEntry.recvBuf[dataEntry.bytesReceived].addr, + dataEntry.recvBuf.len - dataEntry.bytesReceived + ) + if ret > 0: + dataEntry.bytesReceived += ret + result.received = true + continue + case SSL_get_error(dataEntry.ssl, ret): + of SSL_ERROR_WANT_READ: + discard + of SSL_ERROR_WANT_WRITE: + dataEntry.tlsWantWrite = true + else: + # close_notify (ZERO_RETURN), the peer going away (SYSCALL) or a + # protocol error (SSL). Data read just before it is still handed + # to the parser; the close follows on the next wake-up. + if not result.received: + return (false, result.sent, true) + while ERR_get_error() != 0: + discard + break + + # Write the head of the outgoing queue + if dataEntry.outgoingBuffers.len > 0: + let + outgoingBuffer = dataEntry.outgoingBuffers.peekFirst() + totalBytes = outgoingBuffer.buffer1.len + outgoingBuffer.buffer2.len + if outgoingBuffer.bytesSent < totalBytes: + let ret = + if outgoingBuffer.bytesSent < outgoingBuffer.buffer1.len: + SSL_write( + dataEntry.ssl, + cast[cstring](outgoingBuffer.buffer1[outgoingBuffer.bytesSent].addr), + outgoingBuffer.buffer1.len - outgoingBuffer.bytesSent + ) + else: + let buffer2Pos = outgoingBuffer.bytesSent - outgoingBuffer.buffer1.len + SSL_write( + dataEntry.ssl, + cast[cstring](outgoingBuffer.buffer2[buffer2Pos].addr), + outgoingBuffer.buffer2.len - buffer2Pos + ) + if ret > 0: + outgoingBuffer.bytesSent += ret + result.sent = true + else: + case SSL_get_error(dataEntry.ssl, ret): + of SSL_ERROR_WANT_WRITE: + discard # Socket buffer full, Write stays armed below + of SSL_ERROR_WANT_READ: + writeWantsRead = true + else: + return (result.received, false, true) + + var events = {Read} + if dataEntry.tlsWantWrite or + (dataEntry.outgoingBuffers.len > 0 and not writeWantsRead): + events.incl(Write) + server.selector.updateHandle2(clientSocket, events) + proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = var readyKeys: array[maxEventsPerSelectLoop, ReadyKey] @@ -1154,6 +1539,7 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = needClosing: HashSet[SocketHandle] encodedResponses: seq[OutgoingBuffer] encodedFrames: seq[OutgoingBuffer] + listenerOps: seq[ListenerOp] while true: receivedFrom.setLen(0) sentTo.setLen(0) @@ -1178,7 +1564,21 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = else: discard + if shutdownTriggered: + server.destroy(true) + return + if responseQueuedTriggered: + # Listeners added or removed from other threads (or before serving) + # share this wake-up; they are applied at the end of the iteration so + # the selector is only ever touched by this thread, and so a removed + # listener's socket stays registered until the ready keys of this + # iteration have been handled (closing it earlier would let the kernel + # hand its descriptor number to a client accepted below). + withLock server.listenerOpsLock: + while server.listenerOps.len > 0: + listenerOps.add(server.listenerOps.popFirst()) + # If we have responses queued move them to the outgoing buffer queue of # the appropriate socket and update the socket selector to include Write @@ -1261,19 +1661,30 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = else: server.log(DebugLevel, "Dropped message to disconnected client") - if shutdownTriggered: - server.destroy(true) - return - # This is the main client socket select loop for i in 0 ..< readyCount: let readyKey = readyKeys[i] # echo "Socket ready: ", readyKey.fd, " ", readyKey.events - if readyKey.fd == server.socket.int: + if User in readyKey.events: + continue # Handled above + + let dataEntry = + try: + server.selector.getData(readyKey.fd) + except Exception as e: + nil # Unregistered earlier in this iteration (a removed listener) + if dataEntry == nil: + continue + + case dataEntry.kind: + of EventEntry: + discard + of ServerSocketEntry: # We should have a new client socket to accept if Read in readyKey.events: + let listener = dataEntry.listener let (clientSocket, remoteAddress) = when defined(linux) and not defined(nimdoc): var @@ -1282,7 +1693,7 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = let socket = accept4( - server.socket, + listener.socket, sockAddr.addr, addrLen.addr, SOCK_CLOEXEC or SOCK_NONBLOCK @@ -1294,7 +1705,7 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = "" (socket, sockAddrStr) else: - server.socket.accept() + listener.socket.accept() if clientSocket == osInvalidSocket: continue @@ -1306,19 +1717,50 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = if server.tcpNoDelay: server.setNoDelay(clientSocket) - server.clientSockets.incl(clientSocket) + let clientDataEntry = DataEntry(kind: ClientSocketEntry) + clientDataEntry.clientId = server.rand.next() + clientDataEntry.remoteAddress = remoteAddress + clientDataEntry.acceptedAt = epochTime() + clientDataEntry.recvBuf.setLen(initialRecvBufLen) + + if listener.tls != nil: + when defined(ssl): + let ssl = SSL_new(listener.tls.ctx) + if ssl == nil or SSL_set_fd(ssl, clientSocket) != 1: + server.log(ErrorLevel, "SSL_new failed: ", tlsErrorText()) + if ssl != nil: + SSL_free(ssl) + clientSocket.close() + continue + clientDataEntry.ssl = ssl + clientDataEntry.secure = true + else: + clientSocket.close() + continue - let dataEntry = DataEntry(kind: ClientSocketEntry) - dataEntry.clientId = server.rand.next() - dataEntry.remoteAddress = remoteAddress - dataEntry.recvBuf.setLen(initialRecvBufLen) - server.selector.registerHandle2(clientSocket, {Read}, dataEntry) - else: # Client socket + server.clientSockets.incl(clientSocket) + server.selector.registerHandle2(clientSocket, {Read}, clientDataEntry) + of ClientSocketEntry: if Error in readyKey.events: needClosing.incl(readyKey.fd.SocketHandle) continue - let dataEntry = server.selector.getData(readyKey.fd) + var isTls = false + when defined(ssl): + isTls = dataEntry.ssl != nil + + if isTls: + when defined(ssl): + let (received, sent, close) = + server.tlsStep(readyKey.fd.SocketHandle, dataEntry) + if close: + needClosing.incl(readyKey.fd.SocketHandle) + continue + if received: + receivedFrom.add(readyKey.fd.SocketHandle) + if sent: + sentTo.add(readyKey.fd.SocketHandle) + continue if Read in readyKey.events: # Expand the buffer if it is full @@ -1388,6 +1830,14 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = # Leaks DataEntry for this socket server.log(DebugLevel, "Error unregistering client socket") finally: + when defined(ssl): + if dataEntry.ssl != nil: + # Best effort close_notify; the socket is nonblocking so one + # call is all the peer gets. + if dataEntry.tlsHandshaken: + discard SSL_shutdown(dataEntry.ssl) + SSL_free(dataEntry.ssl) + dataEntry.ssl = nil clientSocket.close() server.clientSockets.excl(clientSocket) if dataEntry.upgradedToWebSocket: @@ -1402,65 +1852,56 @@ proc loopForever(server: Server) {.raises: [OSError, IOSelectorsException].} = var close = WebSocketUpdate(event: CloseEvent) websocket.postWebSocketUpdate(close) + # Apply listener changes last, see the note above + for op in listenerOps: + if op.remove: + for i in 0 ..< server.listeners.len: + let listener = server.listeners[i] + if listener.id == op.id: + try: + server.selector.unregister(listener.socket) + except Exception as e: + server.log(DebugLevel, "Error unregistering listener socket") + listener.closeListenerSocket() + server.listeners.delete(i) + break + else: + let dataEntry = DataEntry(kind: ServerSocketEntry) + dataEntry.listener = op.listener + server.selector.registerHandle2(op.listener.socket, {Read}, dataEntry) + server.listeners.add(op.listener) + listenerOps.setLen(0) + proc close*(server: Server) {.raises: [], gcsafe.} = ## Cleanly stops and deallocates the server. ## In-flight request handler calls will be allowed to finish. ## No additional handler calls will be dispatched even if they are queued. - if server.socket.int != 0: + if server.serving.load(moRelaxed): server.trigger(server.shutdown) else: server.destroy(true) -proc serve*( - server: Server, - port: Port, - address = "localhost" -) {.raises: [MummyError].} = - ## The server will serve on the address and port. The default address is - ## localhost. Use "0.0.0.0" to make the server externally accessible (with - ## caution). +proc serve*(server: Server) {.raises: [MummyError].} = + ## Serves on every listener added with `addListener`, and on any added + ## later while serving. At least one listener must have been added. ## This call does not return unless server.close() is called from another ## thread. - - if server.socket.int != 0: - raise newException(MummyError, "Server already has a socket") - - try: - server.socket = createNativeSocket( - Domain.AF_INET, - SockType.SOCK_STREAM, - Protocol.IPPROTO_TCP, - false - ) - if server.socket == osInvalidSocket: - raiseOSError(osLastError()) - - server.socket.setBlocking(false) - server.socket.setSockOptInt(SOL_SOCKET, SO_REUSEADDR, 1) - - let ai = getAddrInfo( - address, - port, - Domain.AF_INET, - SockType.SOCK_STREAM, - Protocol.IPPROTO_TCP, - ) - try: - if bindAddr(server.socket, ai.ai_addr, ai.ai_addrlen.SockLen) < 0: - raiseOSError(osLastError()) - finally: - freeAddrInfo(ai) - - if nativesockets.listen(server.socket, listenBacklogLen) < 0: - raiseOSError(osLastError()) - - let dataEntry = DataEntry(kind: ServerSocketEntry) - server.selector.registerHandle2(server.socket, {Read}, dataEntry) - except Exception as e: + if server.serving.load(moRelaxed): + raise newException(MummyError, "Server is already serving") + + var hasListener: bool + withLock server.listenerOpsLock: + for op in server.listenerOps: + if not op.remove: + hasListener = true + break + if not hasListener: server.destroy(true) - raise currentExceptionAsMummyError() + raise newException(MummyError, "Server has no listeners, call addListener first") server.serving.store(true, moRelaxed) + # Pending listeners are registered by the first loop iteration + server.trigger(server.responseQueued) try: server.loopForever() @@ -1469,6 +1910,23 @@ proc serve*( server.destroy(false) raise currentExceptionAsMummyError() +proc serve*( + server: Server, + port: Port, + address = "localhost" +) {.raises: [MummyError].} = + ## The server will serve on the address and port. The default address is + ## localhost. Use "0.0.0.0" to make the server externally accessible (with + ## caution). + ## This call does not return unless server.close() is called from another + ## thread. + try: + discard server.addListener(port, address) + except MummyError as e: + server.destroy(true) + raise e + server.serve() + proc newServer*( handler: RequestHandler, websocketHandler: WebSocketHandler = nil, @@ -1505,6 +1963,15 @@ proc newServer*( result.workerThreads.setLen(workerThreads) + # The locks first: destroy() acquires taskQueueLock, so a failure below + # must not find them uninitialised. + initLock(result.taskQueueLock) + initCond(result.taskQueueCond) + initLock(result.responseQueueLock) + initLock(result.sendQueueLock) + initLock(result.websocketQueuesLock) + initLock(result.listenerOpsLock) + # Stuff that can fail try: result.responseQueued = newSelectEvent() @@ -1525,12 +1992,6 @@ proc newServer*( shutdownData.event = result.shutdown result.selector.registerEvent(result.shutdown, shutdownData) - initLock(result.taskQueueLock) - initCond(result.taskQueueCond) - initLock(result.responseQueueLock) - initLock(result.sendQueueLock) - initLock(result.websocketQueuesLock) - for i in 0 ..< workerThreads: createThread(result.workerThreads[i], workerProc, result) except Exception as e: diff --git a/tests/test_tls.nim b/tests/test_tls.nim new file mode 100644 index 0000000..bee3278 --- /dev/null +++ b/tests/test_tls.nim @@ -0,0 +1,253 @@ +## TLS listeners: compile with -d:ssl. +## +## A self-signed P-256 certificate for localhost, valid until 2126, in the +## SEC 1 "EC PRIVATE KEY" form (what `openssl ec` writes) so the loader is +## exercised on the same shape FrameOS hands it. Regenerate with: +## openssl req -x509 -newkey ec -pkeyopt ec_paramgen_curve:prime256v1 -nodes \ +## -keyout key.pem -out cert.pem -days 36500 -subj "/CN=localhost" \ +## -addext "subjectAltName=DNS:localhost,IP:127.0.0.1" +## openssl ec -in key.pem -out key-ec.pem + +import std/[httpclient, net, os, strutils], mummy + +const + testCert = """-----BEGIN CERTIFICATE----- +MIIBmjCCAUGgAwIBAgIUZamAN7dinEGglG7utfMsGMaoY5swCgYIKoZIzj0EAwIw +FDESMBAGA1UEAwwJbG9jYWxob3N0MCAXDTI2MDkxMjEyMjAwMFoYDzIxMjYwODE5 +MTIyMDAwWjAUMRIwEAYDVQQDDAlsb2NhbGhvc3QwWTATBgcqhkjOPQIBBggqhkjO +PQMBBwNCAASMI/JAkpfQaY6MDcq29H17JILDjmU+ewB91LvQhsxdpaE5OaszaSUA +ZNK9QzwNkdAzyY1K0TAToIc5i46mQGm7o28wbTAdBgNVHQ4EFgQUJ4P/LWfu8ZcB +SLelGeRgm5ioawQwHwYDVR0jBBgwFoAUJ4P/LWfu8ZcBSLelGeRgm5ioawQwDwYD +VR0TAQH/BAUwAwEB/zAaBgNVHREEEzARgglsb2NhbGhvc3SHBH8AAAEwCgYIKoZI +zj0EAwIDRwAwRAIgId13ZagrbcVFwPpJKQoawNrBB0m0zXb9UAKYErVrt+gCIBvt +ed4IFxd0P+pnRIL9P6rkTp35FXAiA3YX696h4QuJ +-----END CERTIFICATE----- +""" + testKey = """-----BEGIN EC PRIVATE KEY----- +MHcCAQEEIJyAmtKAVs7XPLTMJD7guygpiNJd3O9y2MtywgoTnARroAoGCCqGSM49 +AwEHoUQDQgAEjCPyQJKX0GmOjA3KtvR9eySCw45lPnsAfdS70IbMXaWhOTmrM2kl +AGTSvUM8DZHQM8mNStEwE6CHOYuOpkBpuw== +-----END EC PRIVATE KEY----- +""" + plainPort = 8091 + tlsPort = 8092 + bigBodyLen = 4 * 1024 * 1024 # Well past any socket buffer + +proc handler(request: Request) = + case request.uri: + of "/": + var headers: mummy.HttpHeaders + headers["Content-Type"] = "text/plain" + request.respond(200, headers, "Hello, World!") + of "/secure": + request.respond(200, emptyHttpHeaders(), $request.secure) + of "/echo": + request.respond(200, emptyHttpHeaders(), request.body) + of "/big": + var body = newString(bigBodyLen) + for i in 0 ..< body.len: + body[i] = char(ord('a') + (i mod 26)) + var headers: mummy.HttpHeaders + headers["Content-Type"] = "application/octet-stream" + request.respond(200, headers, body) + of "/ws": + let websocket = request.upgradeToWebSocket() + websocket.send("hello from " & (if request.secure: "wss" else: "ws")) + else: + request.respond(404) + +proc websocketHandler( + websocket: WebSocket, + event: WebSocketEvent, + message: Message +) = + case event: + of MessageEvent: + websocket.send("echo " & message.data, message.kind) + else: + discard + +proc noVerify(): SslContext = + newContext(verifyMode = CVerifyNone) + +proc tlsSocket(port: int): Socket = + result = newSocket() + noVerify().wrapSocket(result) + result.connect("localhost", Port(port)) + +proc readHttpResponse(socket: Socket): tuple[status: string, headers: string, body: string] = + var headers = "" + while true: + let line = socket.recvLine(timeout = 5000) + if line.len == 0 or line == "\r\n": + break + headers &= line & "\n" + result.headers = headers + result.status = headers.splitLines()[0] + var contentLength = 0 + for line in headers.splitLines(): + if line.toLowerAscii().startsWith("content-length:"): + contentLength = parseInt(line.split(':', 1)[1].strip()) + if contentLength > 0: + result.body = socket.recv(contentLength, timeout = 5000) + +proc websocketRoundTrip(socket: Socket, secure: bool) = + socket.send( + "GET /ws HTTP/1.1\r\nHost: localhost\r\nConnection: Upgrade\r\n" & + "Upgrade: websocket\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" & + "Sec-WebSocket-Version: 13\r\n\r\n" + ) + let upgrade = socket.readHttpResponse() + doAssert upgrade.status.startsWith("HTTP/1.1 101"), upgrade.status + doAssert upgrade.headers.contains("Sec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=") + + proc readFrame(): string = + let header = socket.recv(2, timeout = 5000) + doAssert header.len == 2 + doAssert (header[0].uint8 and 0x0f) == 0x1 # Text + let payloadLen = (header[1].uint8 and 0x7f).int + doAssert payloadLen <= 125 + socket.recv(payloadLen, timeout = 5000) + + doAssert readFrame() == "hello from " & (if secure: "wss" else: "ws") + + # A masked client frame, as the RFC requires + let payload = "ping over tls" + var frame = "" + frame.add(char(0x81)) + frame.add(char(0x80 or payload.len.uint8)) + let mask = [0x12'u8, 0x34, 0x56, 0x78] + for b in mask: + frame.add(char(b)) + for i, c in payload: + frame.add(char(c.uint8 xor mask[i mod 4])) + socket.send(frame) + doAssert readFrame() == "echo " & payload + +let + tls = newTlsConfig(testCert, testKey) + server = newServer(handler, websocketHandler) + +discard server.addListener(Port(plainPort), "localhost") +let tlsListener = server.addListener(Port(tlsPort), "localhost", tls) +doAssert tlsListener.secure +doAssert tlsListener.port == Port(tlsPort) + +var requesterThread: Thread[void] + +proc requesterBody() = + server.waitUntilReady() + + block: # Bad certificate material is refused up front + doAssertRaises(MummyError): + discard newTlsConfig("not a certificate", testKey) + doAssertRaises(MummyError): + discard newTlsConfig(testCert, "not a key") + doAssertRaises(MummyError): + discard newTlsConfig("", "") + + block: # GET over TLS + let client = newHttpClient(sslContext = noVerify()) + doAssert client.getContent("https://localhost:" & $tlsPort & "/") == "Hello, World!" + + block: # POST over TLS with a body + let client = newHttpClient(sslContext = noVerify()) + let body = "x".repeat(100_000) + let response = client.post("https://localhost:" & $tlsPort & "/echo", body) + doAssert response.status.startsWith("200") + doAssert response.body == body + + block: # Plain and TLS listeners at once, Request.secure tells them apart + let plain = newHttpClient() + doAssert plain.getContent("http://localhost:" & $plainPort & "/secure") == "false" + let secure = newHttpClient(sslContext = noVerify()) + doAssert secure.getContent("https://localhost:" & $tlsPort & "/secure") == "true" + + block: # A response far larger than the socket buffer (partial SSL_write) + let client = newHttpClient(sslContext = noVerify()) + let body = client.getContent("https://localhost:" & $tlsPort & "/big") + doAssert body.len == bigBodyLen + doAssert body[0] == 'a' and body[25] == 'z' and body[26] == 'a' + doAssert body[^1] == char(ord('a') + ((bigBodyLen - 1) mod 26)) + + block: # Several requests on one keep-alive TLS connection + let client = newHttpClient(sslContext = noVerify()) + for i in 0 ..< 20: + doAssert client.getContent("https://localhost:" & $tlsPort & "/secure") == "true" + + block: # WebSocket over TLS and over plain + let secure = tlsSocket(tlsPort) + secure.websocketRoundTrip(secure = true) + secure.close() + let plain = newSocket() + plain.connect("localhost", Port(plainPort)) + plain.websocketRoundTrip(secure = false) + plain.close() + + block: # A client that stalls mid-handshake must not block the loop + let stalled = newSocket() + stalled.connect("localhost", Port(tlsPort)) + # The first five bytes of a TLS record header, then silence + stalled.send("\x16\x03\x01\x02\x00") + let silent = newSocket() + silent.connect("localhost", Port(tlsPort)) + let client = newHttpClient(sslContext = noVerify()) + doAssert client.getContent("https://localhost:" & $tlsPort & "/") == "Hello, World!" + let plain = newHttpClient() + doAssert plain.getContent("http://localhost:" & $plainPort & "/") == "Hello, World!" + stalled.close() + silent.close() + + block: # Plain text sent to the TLS port is rejected without harm + let junk = newSocket() + junk.connect("localhost", Port(tlsPort)) + junk.send("GET / HTTP/1.1\r\nHost: localhost\r\n\r\n") + var got = "" + try: + got = junk.recv(1024, timeout = 2000) + except CatchableError: + discard + doAssert not got.startsWith("HTTP/1.1 200") + junk.close() + let client = newHttpClient(sslContext = noVerify()) + doAssert client.getContent("https://localhost:" & $tlsPort & "/") == "Hello, World!" + + block: # Listeners added and removed while serving + let extra = server.addListener(Port(0), "localhost") + doAssert extra.port != Port(0) + doAssert not extra.secure + let client = newHttpClient() + doAssert client.getContent("http://localhost:" & $extra.port.int & "/secure") == "false" + let extraTls = server.addListener(Port(0), "localhost", tls) + let secure = newHttpClient(sslContext = noVerify()) + doAssert secure.getContent("https://localhost:" & $extraTls.port.int & "/secure") == "true" + + server.removeListener(extra) + server.removeListener(extraTls) + # The loop applies the removal on its next wake-up; a request to a + # still-open listener gets served, so poll until the port refuses. + var refused = false + for attempt in 0 ..< 100: + let probe = newSocket() + try: + probe.connect("localhost", Port(extra.port.int), timeout = 1000) + probe.close() + sleep(20) + except OSError: + refused = true + break + doAssert refused + # The original listeners are unaffected + doAssert client.getContent("http://localhost:" & $plainPort & "/secure") == "false" + doAssert secure.getContent("https://localhost:" & $tlsPort & "/secure") == "true" + + echo "Done, shut down the server" + server.close() + +proc requesterProc() = + {.cast(gcsafe).}: # tls and server are read-only globals here + requesterBody() + +createThread(requesterThread, requesterProc) + +server.serve()