Something went wrong. Try again.
Native PostgreSQL driver / client for Zig
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231const std = @import("std");const lib = @import("lib.zig");
const openssl = @cImport({ @cInclude("openssl/ssl.h"); @cInclude("openssl/err.h");});
const posix = std.posix;
const Conn = lib.Conn;const Allocator = std.mem.Allocator;
const DEFAULT_HOST = "127.0.0.1";
pub const Stream = if (lib.has_openssl) TLSStream else PlainStream;
const TLSStream = struct { valid: bool, ssl: ?*openssl.SSL, socket: posix.socket_t,
pub fn connect(allocator: Allocator, opts: Conn.Opts, ctx_: ?*openssl.SSL_CTX) !Stream { const plain = try PlainStream.connect(allocator, opts, null); errdefer plain.close();
const socket = plain.socket;
var ssl: ?*openssl.SSL = null; if (ctx_) |ctx| { // PostgreSQL TLS starts off as a plain connection which we upgrade try writeSocket(socket, &.{ 0, 0, 0, 8, 4, 210, 22, 47 }); var buf = [1]u8{0}; _ = try readSocket(socket, &buf); if (buf[0] != 'S') { return error.SSLNotSupportedByServer; }
ssl = openssl.SSL_new(ctx) orelse return error.SSLNewFailed; errdefer openssl.SSL_free(ssl);
if (opts.host) |host| { if (isHostName(host)) { // don't send this for an ip address var owned = false; const h = opts._hostz orelse blk: { owned = true; break :blk try allocator.dupeZ(u8, host); };
defer if (owned) { allocator.free(h); };
if (openssl.SSL_set_tlsext_host_name(ssl, h.ptr) != 1) { return error.SSLHostNameFailed; } } switch (opts.tls) { .verify_full => openssl.SSL_set_verify(ssl, openssl.SSL_VERIFY_PEER, null), else => {}, } }
if (openssl.SSL_set_fd(ssl, socket) != 1) { return error.SSLSetFdFailed; }
{ const ret = openssl.SSL_connect(ssl); if (ret != 1) { const verification_code = openssl.SSL_get_verify_result(ssl); if (comptime lib._stderr_tls) { lib.printSSLError(); } if (verification_code != openssl.X509_V_OK) { if (comptime lib._stderr_tls) { std.debug.print("ssl verification error: {s}\n", .{openssl.X509_verify_cert_error_string(verification_code)}); } return error.SSLCertificationVerificationError; } return error.SSLConnectFailed; } } }
return .{ .ssl = ssl, .valid = true, .socket = socket, }; }
pub fn close(self: *Stream) void { if (self.ssl) |ssl| { if (self.valid) { _ = openssl.SSL_shutdown(ssl); self.valid = false; } openssl.SSL_free(ssl); } posix.close(self.socket); }
pub fn writeAll(self: *Stream, data: []const u8) !void { if (self.ssl) |ssl| { const result = openssl.SSL_write(ssl, data.ptr, @intCast(data.len)); if (result <= 0) { self.valid = false; return error.SSLWriteFailed; } return; } return writeSocket(self.socket, data); }
pub fn read(self: *Stream, buf: []u8) !usize { if (self.ssl) |ssl| { var read_len: usize = undefined; const result = openssl.SSL_read_ex(ssl, buf.ptr, @intCast(buf.len), &read_len); if (result <= 0) { self.valid = false; return error.SSLReadFailed; } return read_len; }
return readSocket(self.socket, buf); }};
const PlainStream = struct { socket: posix.socket_t,
pub fn connect(allocator: Allocator, opts: Conn.Opts, _: anytype) !PlainStream { var is_tcp = false; const socket = blk: { if (opts.unix_socket) |path| { if (comptime std.net.has_unix_sockets == false or std.posix.AF == void) { return error.UnixPathNotSupported; } break :blk (try std.net.connectUnixSocket(path)).handle; } else { const host = opts.host orelse DEFAULT_HOST; const port = opts.port orelse 5432; is_tcp = true; break :blk (try std.net.tcpConnectToHost(allocator, host, port)).handle; } }; errdefer posix.close(socket);
if (is_tcp) { try setKeepalive(socket, opts); }
return .{ .socket = socket, }; }
pub fn close(self: *const PlainStream) void { posix.close(self.socket); }
pub fn writeAll(self: *const PlainStream, data: []const u8) !void { return writeSocket(self.socket, data); }
pub fn read(self: *const PlainStream, buf: []u8) !usize { return readSocket(self.socket, buf); }};
fn setKeepalive(handle: posix.socket_t, opts: Conn.Opts) !void { if (opts.keepalive == false) { return; }
const on: c_int = 1; try posix.setsockopt(handle, posix.SOL.SOCKET, posix.SO.KEEPALIVE, std.mem.asBytes(&on));
const TCP = posix.TCP; const level = posix.IPPROTO.TCP;
if (opts.keepalive_idle) |idle| { const optname: ?u32 = comptime if (@hasDecl(TCP, "KEEPIDLE")) TCP.KEEPIDLE else if (@hasDecl(TCP, "KEEPALIVE")) TCP.KEEPALIVE else null; if (optname) |name| { const v: c_int = @intCast(idle); posix.setsockopt(handle, level, name, std.mem.asBytes(&v)) catch {}; } }
if (opts.keepalive_interval) |intvl| { if (comptime @hasDecl(TCP, "KEEPINTVL")) { const v: c_int = @intCast(intvl); posix.setsockopt(handle, level, TCP.KEEPINTVL, std.mem.asBytes(&v)) catch {}; } }
if (opts.keepalive_count) |cnt| { if (comptime @hasDecl(TCP, "KEEPCNT")) { const v: c_int = @intCast(cnt); posix.setsockopt(handle, level, TCP.KEEPCNT, std.mem.asBytes(&v)) catch {}; } }}
fn readSocket(socket: posix.socket_t, buf: []u8) !usize { return posix.read(socket, buf);}
fn writeSocket(socket: posix.socket_t, data: []const u8) !void { var i: usize = 0; while (i < data.len) { i += try posix.write(socket, data[i..]); }}
fn isHostName(host: []const u8) bool { if (std.mem.indexOfScalar(u8, host, ':') != null) { // IPv6 return false; } return std.mem.indexOfNone(u8, host, "0123456789.") != null;}