summaryrefslogtreecommitdiff
path: root/src/quic.zig
diff options
context:
space:
mode:
Diffstat (limited to 'src/quic.zig')
-rw-r--r--src/quic.zig302
1 files changed, 302 insertions, 0 deletions
diff --git a/src/quic.zig b/src/quic.zig
new file mode 100644
index 0000000..c78a79d
--- /dev/null
+++ b/src/quic.zig
@@ -0,0 +1,302 @@
+const std = @import("std");
+const libc = std.c;
+
+/// QUIC transport with caller-selected ALPN. OpenSSL bindings are injected so
+/// applications control linkage. Certificates are ephemeral and peers unauthenticated.
+pub fn Quic(comptime ssl: type, comptime protocol: []const u8) type {
+ if (protocol.len == 0 or protocol.len > 255) @compileError("invalid ALPN length");
+ return struct {
+ comptime {
+ if (ssl.OPENSSL_VERSION_NUMBER < 0x30600000)
+ @compileError("9P over QUIC requires OpenSSL 3.6 or newer");
+ }
+
+ pub const alpn = protocol;
+ pub const Error = error{ Tls, Socket, SocketFlags, SocketOption, Bind, Address, Closed, InvalidWrite };
+
+ pub const Listener = struct {
+ fd: c_int,
+ handle: *ssl.SSL,
+ address: std.Io.net.IpAddress,
+
+ pub fn init(address: std.Io.net.IpAddress) Error!Listener {
+ ssl.ERR_clear_error();
+ const ctx = ssl.SSL_CTX_new(ssl.OSSL_QUIC_server_method()) orelse return error.Tls;
+ defer ssl.SSL_CTX_free(ctx);
+ const key = ssl.EVP_PKEY_Q_keygen(null, null, "EC", @as([*:0]const u8, "prime256v1")) orelse return error.Tls;
+ defer ssl.EVP_PKEY_free(key);
+ const cert = ssl.X509_new() orelse return error.Tls;
+ defer ssl.X509_free(cert);
+ if (ssl.X509_set_version(cert, 2) != 1 or
+ ssl.ASN1_INTEGER_set(ssl.X509_get_serialNumber(cert), 1) != 1 or
+ ssl.X509_gmtime_adj(ssl.X509_getm_notBefore(cert), -60) == null or
+ ssl.X509_gmtime_adj(ssl.X509_getm_notAfter(cert), 365 * 24 * 60 * 60) == null or
+ ssl.X509_set_pubkey(cert, key) != 1) return error.Tls;
+ const name = ssl.X509_get_subject_name(cert) orelse return error.Tls;
+ if (ssl.X509_NAME_add_entry_by_txt(name, "CN", ssl.MBSTRING_ASC, protocol.ptr, @intCast(protocol.len), -1, 0) != 1 or
+ ssl.X509_set_issuer_name(cert, name) != 1 or
+ ssl.X509_sign(cert, key, ssl.EVP_sha256()) <= 0 or
+ ssl.SSL_CTX_use_certificate(ctx, cert) != 1 or
+ ssl.SSL_CTX_use_PrivateKey(ctx, key) != 1) return error.Tls;
+ ssl.SSL_CTX_set_verify(ctx, ssl.SSL_VERIFY_NONE, null);
+ ssl.SSL_CTX_set_alpn_select_cb(ctx, selectAlpn, null);
+ var addr: libc.sockaddr.storage = undefined;
+ const addr_len = sockaddr(address, &addr);
+ const fd = try udp(addr.family);
+ errdefer _ = libc.close(fd);
+ if (libc.bind(fd, @ptrCast(&addr), addr_len) != 0) return error.Bind;
+ var actual_len: libc.socklen_t = @sizeOf(@TypeOf(addr));
+ if (libc.getsockname(fd, @ptrCast(&addr), &actual_len) != 0) return error.Address;
+ var actual = address;
+ actual.setPort(switch (address) {
+ .ip4 => std.mem.bigToNative(u16, @as(*const libc.sockaddr.in, @ptrCast(&addr)).port),
+ .ip6 => std.mem.bigToNative(u16, @as(*const libc.sockaddr.in6, @ptrCast(&addr)).port),
+ });
+ const handle = ssl.SSL_new_listener(ctx, 0) orelse return error.Tls;
+ errdefer ssl.SSL_free(handle);
+ if (ssl.SSL_set_fd(handle, fd) != 1 or ssl.SSL_set_blocking_mode(handle, 0) != 1 or
+ ssl.SSL_listen(handle) != 1) return error.Tls;
+ return .{ .fd = fd, .handle = handle, .address = actual };
+ }
+
+ pub fn accept(l: *Listener) Error!?Connection {
+ ssl.ERR_clear_error();
+ const handle = ssl.SSL_accept_connection(l.handle, ssl.SSL_ACCEPT_CONNECTION_NO_BLOCK) orelse {
+ if (ssl.ERR_peek_error() != 0) return error.Tls;
+ return null;
+ };
+ errdefer ssl.SSL_free(handle);
+ if (ssl.SSL_set_default_stream_mode(handle, ssl.SSL_DEFAULT_STREAM_MODE_NONE) != 1 or
+ ssl.SSL_set_blocking_mode(handle, 0) != 1) return error.Tls;
+ return .{ .handle = handle };
+ }
+
+ pub fn events(l: *Listener) Error!void {
+ ssl.ERR_clear_error();
+ if (ssl.SSL_handle_events(l.handle) != 1) return error.Tls;
+ }
+
+ pub fn poll(l: *const Listener) libc.pollfd {
+ return pollFd(l.handle, l.fd);
+ }
+
+ pub fn nextDue(l: *const Listener) ?i32 {
+ return due(l.handle);
+ }
+
+ // Accepted connections must be released before the shared UDP socket.
+ pub fn deinit(l: *Listener) void {
+ ssl.SSL_free(l.handle);
+ _ = libc.close(l.fd);
+ l.* = undefined;
+ }
+ };
+
+ pub const Connection = struct {
+ handle: *ssl.SSL,
+ stream: ?*ssl.SSL = null,
+ fd: c_int = -1,
+ pending_write_len: usize = 0,
+
+ pub fn dial(address: std.Io.net.IpAddress) Error!Connection {
+ ssl.ERR_clear_error();
+ const ctx = ssl.SSL_CTX_new(ssl.OSSL_QUIC_client_method()) orelse return error.Tls;
+ defer ssl.SSL_CTX_free(ctx);
+ ssl.SSL_CTX_set_verify(ctx, ssl.SSL_VERIFY_NONE, null);
+ const fd = try udp(if (address == .ip4) libc.AF.INET else libc.AF.INET6);
+ errdefer _ = libc.close(fd);
+ const handle = ssl.SSL_new(ctx) orelse return error.Tls;
+ errdefer ssl.SSL_free(handle);
+ if (ssl.SSL_set_fd(handle, fd) != 1 or ssl.SSL_set_blocking_mode(handle, 0) != 1 or
+ ssl.SSL_set_default_stream_mode(handle, ssl.SSL_DEFAULT_STREAM_MODE_NONE) != 1) return error.Tls;
+ const protocols = [_]u8{alpn.len} ++ protocol[0..protocol.len].*;
+ if (ssl.SSL_set_alpn_protos(handle, &protocols, protocols.len) != 0) return error.Tls;
+ const peer = ssl.BIO_ADDR_new() orelse return error.Tls;
+ defer ssl.BIO_ADDR_free(peer);
+ const made = switch (address) {
+ .ip4 => |ip| ssl.BIO_ADDR_rawmake(peer, libc.AF.INET, &ip.bytes, ip.bytes.len, std.mem.nativeToBig(u16, ip.port)),
+ .ip6 => |ip| ssl.BIO_ADDR_rawmake(peer, libc.AF.INET6, &ip.bytes, ip.bytes.len, std.mem.nativeToBig(u16, ip.port)),
+ };
+ if (made != 1 or ssl.SSL_set1_initial_peer_addr(handle, peer) != 1) return error.Tls;
+ return .{ .handle = handle, .fd = fd };
+ }
+
+ pub fn handshake(c: *Connection) Error!bool {
+ ssl.ERR_clear_error();
+ var close_info: ssl.SSL_CONN_CLOSE_INFO = undefined;
+ if (ssl.SSL_get_conn_close_info(c.handle, &close_info, @sizeOf(@TypeOf(close_info))) == 1)
+ return error.Closed;
+ if (ssl.SSL_is_init_finished(c.handle) == 1) return true;
+ const rc = if (c.fd >= 0) ssl.SSL_connect(c.handle) else ssl.SSL_accept(c.handle);
+ if (rc == 1) return true;
+ try retry(c.handle, rc);
+ return false;
+ }
+
+ fn ready(c: *Connection) Error!bool {
+ if (!try c.handshake()) return false;
+ if (c.stream != null) return true;
+ ssl.ERR_clear_error();
+ const stream = if (c.fd >= 0)
+ ssl.SSL_new_stream(c.handle, ssl.SSL_STREAM_FLAG_NO_BLOCK)
+ else
+ ssl.SSL_accept_stream(c.handle, ssl.SSL_ACCEPT_STREAM_NO_BLOCK);
+ if (stream == null) {
+ if (ssl.ERR_peek_error() != 0) return error.Tls;
+ return false;
+ }
+ errdefer ssl.SSL_free(stream);
+ if (ssl.SSL_set_blocking_mode(stream, 0) != 1 or
+ ssl.SSL_get_stream_id(stream) != 0 or
+ ssl.SSL_set_incoming_stream_policy(c.handle, ssl.SSL_INCOMING_STREAM_POLICY_REJECT, 0) != 1)
+ return error.Tls;
+ _ = ssl.SSL_set_mode(stream, ssl.SSL_MODE_ACCEPT_MOVING_WRITE_BUFFER);
+ c.stream = stream;
+ return true;
+ }
+
+ pub fn read(c: *Connection, bytes: []u8) Error!?usize {
+ if (!try c.ready()) return null;
+ ssl.ERR_clear_error();
+ var len: usize = 0;
+ const rc = ssl.SSL_read_ex(c.stream, bytes.ptr, bytes.len, &len);
+ if (rc == 1) return len;
+ if (ssl.SSL_get_error(c.stream, rc) == ssl.SSL_ERROR_ZERO_RETURN) return 0;
+ try retry(c.stream.?, rc);
+ return null;
+ }
+
+ pub fn pending(c: *const Connection) bool {
+ var close_info: ssl.SSL_CONN_CLOSE_INFO = undefined;
+ if (ssl.SSL_get_conn_close_info(c.handle, &close_info, @sizeOf(@TypeOf(close_info))) == 1) return true;
+ if (c.stream) |stream| {
+ var item: ssl.SSL_POLL_ITEM = .{ .desc = ssl.SSL_as_poll_descriptor(stream), .events = ssl.SSL_POLL_EVENT_RE, .revents = 0 };
+ const timeout: ssl.struct_timeval = .{ .tv_sec = 0, .tv_usec = 0 };
+ if (ssl.SSL_poll(&item, 1, @sizeOf(@TypeOf(item)), &timeout, ssl.SSL_POLL_FLAG_NO_HANDLE_EVENTS, null) != 1) return true;
+ return item.revents != 0;
+ }
+ return ssl.SSL_get_accept_stream_queue_len(c.handle) != 0;
+ }
+
+ pub fn write(c: *Connection, bytes: []const u8) Error!usize {
+ if (!try c.ready()) return 0;
+ if (bytes.len < c.pending_write_len) return error.InvalidWrite;
+ const requested = if (c.pending_write_len != 0) c.pending_write_len else bytes.len;
+ if (requested == 0) return 0;
+ ssl.ERR_clear_error();
+ var len: usize = 0;
+ const rc = ssl.SSL_write_ex(c.stream, bytes.ptr, requested, &len);
+ if (rc == 1) {
+ c.pending_write_len = 0;
+ return len;
+ }
+ try retry(c.stream.?, rc);
+ c.pending_write_len = requested;
+ return 0;
+ }
+
+ pub fn conclude(c: *Connection) Error!void {
+ if (c.pending_write_len != 0) return error.InvalidWrite;
+ if (!try c.ready()) return error.Closed;
+ ssl.ERR_clear_error();
+ if (ssl.SSL_stream_conclude(c.stream, 0) != 1) return error.Tls;
+ }
+
+ pub fn events(c: *Connection) Error!void {
+ if (c.fd < 0) return;
+ ssl.ERR_clear_error();
+ if (ssl.SSL_handle_events(c.handle) != 1) return error.Tls;
+ }
+
+ pub fn poll(c: *const Connection) ?libc.pollfd {
+ return if (c.fd >= 0) pollFd(c.handle, c.fd) else null;
+ }
+
+ pub fn nextDue(c: *const Connection) ?i32 {
+ return if (c.fd >= 0) due(c.handle) else null;
+ }
+
+ pub fn deinit(c: *Connection) void {
+ ssl.ERR_clear_error();
+ _ = ssl.SSL_shutdown_ex(c.handle, ssl.SSL_SHUTDOWN_FLAG_RAPID | ssl.SSL_SHUTDOWN_FLAG_NO_STREAM_FLUSH | ssl.SSL_SHUTDOWN_FLAG_NO_BLOCK, null, 0);
+ ssl.SSL_free(c.stream);
+ ssl.SSL_free(c.handle);
+ if (c.fd >= 0) _ = libc.close(c.fd);
+ c.* = undefined;
+ }
+ };
+
+ fn udp(family: u16) Error!c_int {
+ const fd = libc.socket(family, libc.SOCK.DGRAM, 0);
+ if (fd < 0) return error.Socket;
+ errdefer _ = libc.close(fd);
+ const flags = libc.fcntl(fd, libc.F.GETFL, @as(c_int, 0));
+ if (flags < 0) return error.SocketFlags;
+ var options: libc.O = @bitCast(@as(u32, @bitCast(flags)));
+ options.NONBLOCK = true;
+ if (libc.fcntl(fd, libc.F.SETFL, @as(c_int, @bitCast(@as(u32, @bitCast(options))))) != 0 or
+ libc.fcntl(fd, libc.F.SETFD, @as(c_int, libc.FD_CLOEXEC)) != 0) return error.SocketFlags;
+ if (family == libc.AF.INET6) {
+ const enabled: c_int = 1;
+ const v6only = if (@import("builtin").os.tag.isDarwin()) 27 else libc.IPV6.V6ONLY;
+ if (libc.setsockopt(fd, libc.IPPROTO.IPV6, v6only, &enabled, @sizeOf(c_int)) != 0) return error.SocketOption;
+ }
+ return fd;
+ }
+
+ fn sockaddr(address: std.Io.net.IpAddress, out: *libc.sockaddr.storage) libc.socklen_t {
+ switch (address) {
+ .ip4 => |ip| {
+ const addr: *libc.sockaddr.in = @ptrCast(out);
+ addr.* = .{ .port = std.mem.nativeToBig(u16, ip.port), .addr = @bitCast(ip.bytes) };
+ return @sizeOf(libc.sockaddr.in);
+ },
+ .ip6 => |ip| {
+ const addr: *libc.sockaddr.in6 = @ptrCast(out);
+ addr.* = .{ .port = std.mem.nativeToBig(u16, ip.port), .addr = ip.bytes, .flowinfo = 0, .scope_id = 0 };
+ return @sizeOf(libc.sockaddr.in6);
+ },
+ }
+ }
+
+ fn pollFd(handle: *ssl.SSL, fd: c_int) libc.pollfd {
+ var result: libc.pollfd = .{ .fd = fd, .events = 0, .revents = 0 };
+ if (ssl.SSL_net_read_desired(handle) == 1) result.events |= libc.POLL.IN;
+ if (ssl.SSL_net_write_desired(handle) == 1) result.events |= libc.POLL.OUT;
+ return result;
+ }
+
+ fn due(handle: *ssl.SSL) ?i32 {
+ var tv: ssl.struct_timeval = undefined;
+ var infinite: c_int = undefined;
+ if (ssl.SSL_get_event_timeout(handle, &tv, &infinite) != 1) return 0;
+ if (infinite != 0) return null;
+ const ms = @as(i128, tv.tv_sec) * 1000 + @divFloor(@as(i128, tv.tv_usec) + 999, 1000);
+ return @intCast(std.math.clamp(ms, 0, std.math.maxInt(i32)));
+ }
+
+ fn retry(handle: *ssl.SSL, rc: c_int) Error!void {
+ switch (ssl.SSL_get_error(handle, rc)) {
+ ssl.SSL_ERROR_WANT_READ, ssl.SSL_ERROR_WANT_WRITE => {},
+ ssl.SSL_ERROR_ZERO_RETURN => return error.Closed,
+ else => return error.Tls,
+ }
+ }
+
+ fn selectAlpn(_: ?*ssl.SSL, out: [*c][*c]const u8, outlen: [*c]u8, input: [*c]const u8, len: c_uint, _: ?*anyopaque) callconv(.c) c_int {
+ var offset: usize = 0;
+ while (offset < len) {
+ const size = input[offset];
+ offset += 1;
+ if (size > len - offset) return ssl.SSL_TLSEXT_ERR_ALERT_FATAL;
+ if (std.mem.eql(u8, input[offset..][0..size], alpn)) {
+ out.* = input + offset;
+ outlen.* = size;
+ return ssl.SSL_TLSEXT_ERR_OK;
+ }
+ offset += size;
+ }
+ return ssl.SSL_TLSEXT_ERR_ALERT_FATAL;
+ }
+ };
+}