// Copyright 2025 International Digital Economy Academy
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
///|
// A single TLS session: handshake state, encrypted/plain buffers, and wants flags.
priv struct TlsConnectionHandle(UInt64)
///|
priv enum TlsTrustMode {
TrustNoVerification = 0
TrustSystemRoot = 1
TrustCustomRoot = 2
}
///|
priv enum TlsState {
Completed = 0
WantRead = 1
WantWrite = 2
Error = 3
Eof = 4
ReNegotiation = 5
}
///|
// Host TLS calls return non-negative byte counts and these negative status
// sentinels through the same Int. Values must match moonrun's Rust TLS status
// constants.
const TLS_ERROR = -1
///|
const TLS_CLOSED = -2
///|
const TLS_WOULD_BLOCK = -3
///|
const TLS_RENEGOTIATION = -4
///|
fn TlsConnectionHandle::is_null(self : TlsConnectionHandle) -> Bool {
self.0 == 0UL
}
///|
fn TrustedRoot::mode(self : TrustedRoot) -> TlsTrustMode {
match self {
NoVerification => TrustNoVerification
SystemRoot => TrustSystemRoot
CustomPemFile(_) => TrustCustomRoot
}
}
///|
fn tls_take_global_error_message() -> String {
let buf = tls_take_global_error()
defer buf.free()
@os_string.decode(buf)
}
///|
#unsafe_skip_stub_check
#borrow(buf)
fn tls_buffer_length(buf : @c_buffer.Buffer) -> Int = "moonbitlang/async" "c_buffer/length"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::new() -> TlsConnectionHandle = "moonbitlang/async" "tls/connection/new"
///|
#unsafe_skip_stub_check
#borrow(host)
fn TlsConnectionHandle::set_client(
self : TlsConnectionHandle,
host : String,
host_len? : Int = host.length(),
sni~ : Bool,
trust : TlsTrustMode,
) -> Int = "moonbitlang/async" "tls/connection/set_client"
///|
#unsafe_skip_stub_check
#borrow(root)
fn TlsConnectionHandle::add_root_certificate(
self : TlsConnectionHandle,
root : Bytes,
root_len? : Int = root.length(),
) -> Int = "moonbitlang/async" "tls/connection/add_root_certificate"
///|
#unsafe_skip_stub_check
#borrow(private_key_file, certificate_file)
fn TlsConnectionHandle::set_server_files(
self : TlsConnectionHandle,
private_key_file : @os_string.OsString,
private_key_file_len? : Int = private_key_file.to_string().length(),
private_key_type : X509FileType,
certificate_file : @os_string.OsString,
certificate_file_len? : Int = certificate_file.to_string().length(),
certificate_type : X509FileType,
) -> Int = "moonbitlang/async" "tls/connection/set_server_files"
///|
#unsafe_skip_stub_check
#borrow(pfx_content)
fn TlsConnectionHandle::set_server_pfx(
self : TlsConnectionHandle,
pfx_content : Bytes,
pfx_content_len? : Int = pfx_content.length(),
) -> Int = "moonbitlang/async" "tls/connection/set_server_pfx"
///|
#unsafe_skip_stub_check
fn tls_take_global_error() -> @c_buffer.Buffer = "moonbitlang/async" "tls/error/take_global"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::free_ffi(self : TlsConnectionHandle) -> Unit = "moonbitlang/async" "tls/connection/free"
///|
fn TlsConnectionHandle::free(self : TlsConnectionHandle) -> Unit {
if !self.is_null() {
self.free_ffi()
}
}
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::take_error(
self : TlsConnectionHandle,
) -> @c_buffer.Buffer = "moonbitlang/async" "tls/connection/take_error"
///|
fn TlsConnectionHandle::take_error_message(
self : TlsConnectionHandle,
) -> String {
let buf = self.take_error()
defer buf.free()
@os_string.decode(buf)
}
///|
#unsafe_skip_stub_check
#borrow(in_buffer, out_buffer, plain_buffer)
fn TlsConnectionHandle::read_plain(
self : TlsConnectionHandle,
in_buffer~ : FixedArray[Byte],
in_buffer_offset~ : Int,
in_buffer_len~ : Int,
out_buffer~ : FixedArray[Byte],
out_buffer_offset~ : Int,
out_buffer_len~ : Int,
plain_buffer~ : FixedArray[Byte],
plain_buffer_offset~ : Int,
plain_buffer_len~ : Int,
) -> Int = "moonbitlang/async" "tls/connection/read_plain"
///|
#unsafe_skip_stub_check
#borrow(in_buffer, out_buffer, plain_buffer)
fn TlsConnectionHandle::write_plain(
self : TlsConnectionHandle,
in_buffer~ : FixedArray[Byte],
in_buffer_offset~ : Int,
in_buffer_len~ : Int,
out_buffer~ : FixedArray[Byte],
out_buffer_offset~ : Int,
out_buffer_len~ : Int,
plain_buffer~ : Bytes,
plain_buffer_offset~ : Int,
plain_buffer_len~ : Int,
) -> Int = "moonbitlang/async" "tls/connection/write_plain"
///|
#unsafe_skip_stub_check
#borrow(in_buffer, out_buffer)
fn TlsConnectionHandle::connect(
self : TlsConnectionHandle,
in_buffer~ : FixedArray[Byte],
in_buffer_offset~ : Int,
in_buffer_len~ : Int,
out_buffer~ : FixedArray[Byte],
out_buffer_offset~ : Int,
out_buffer_len~ : Int,
) -> TlsState = "moonbitlang/async" "tls/connection/connect"
///|
#unsafe_skip_stub_check
#borrow(in_buffer, out_buffer)
fn TlsConnectionHandle::accept(
self : TlsConnectionHandle,
in_buffer~ : FixedArray[Byte],
in_buffer_offset~ : Int,
in_buffer_len~ : Int,
out_buffer~ : FixedArray[Byte],
out_buffer_offset~ : Int,
out_buffer_len~ : Int,
) -> TlsState = "moonbitlang/async" "tls/connection/accept"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::bytes_read(self : TlsConnectionHandle) -> Int = "moonbitlang/async" "tls/connection/bytes_read"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::bytes_to_write(self : TlsConnectionHandle) -> Int = "moonbitlang/async" "tls/connection/bytes_to_write"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::wants_read(self : TlsConnectionHandle) -> Bool = "moonbitlang/async" "tls/connection/wants_read"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::wants_write(self : TlsConnectionHandle) -> Bool = "moonbitlang/async" "tls/connection/wants_write"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::shutdown(self : TlsConnectionHandle) -> Int = "moonbitlang/async" "tls/connection/shutdown"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::peer_certificate(
self : TlsConnectionHandle,
) -> @c_buffer.Buffer = "moonbitlang/async" "tls/connection/peer_certificate"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::unique_channel_binding(
self : TlsConnectionHandle,
) -> @c_buffer.Buffer = "moonbitlang/async" "tls/connection/unique_channel_binding"
///|
#unsafe_skip_stub_check
fn TlsConnectionHandle::server_endpoint_channel_binding(
self : TlsConnectionHandle,
) -> @c_buffer.Buffer = "moonbitlang/async" "tls/connection/server_endpoint_channel_binding"
///|
fn TlsConnectionHandle::check_new(
self : TlsConnectionHandle,
) -> TlsConnectionHandle raise TlsError {
if self.is_null() {
raise TlsError(tls_take_global_error_message())
}
self
}
///|
struct Tls {
conn : TlsConnectionHandle
host : String?
is_client : Bool
transport : Transport
read_buf : @io.ReaderBuffer
mut shutdown : Bool
mut closed : Bool
}
///|
fn[R : @io.Reader, W : @io.Writer] Tls::from_pair(
conn : TlsConnectionHandle,
r : R,
w : W,
is_client~ : Bool,
host~ : String?,
) -> Tls {
{
conn,
host,
is_client,
transport: Transport::new(r, w),
read_buf: @io.ReaderBuffer::new(),
shutdown: false,
closed: false,
}
}
///|
async fn Tls::connect(self : Tls) -> Unit {
let read_buf = self.transport.reader._get_internal_buffer().repr()
let write_buf = self.transport.write_buf
for ;; {
let ret = self.conn.connect(
in_buffer=read_buf.buf,
in_buffer_offset=read_buf.start,
in_buffer_len=read_buf.len,
out_buffer=write_buf.buf,
out_buffer_offset=write_buf.start + write_buf.len,
out_buffer_len=write_buf.buf.length() - write_buf.start - write_buf.len,
)
read_buf.drop(self.conn.bytes_read())
write_buf.len += self.conn.bytes_to_write()
match ret {
Completed => {
self.transport.flush_write()
break
}
WantRead => {
self.transport.flush_write()
self.transport.read_more()
}
WantWrite => self.transport.flush_write()
Eof => raise ConnectionClosed
Error => raise TlsError(self.conn.take_error_message())
ReNegotiation => panic()
}
}
}
///|
async fn Tls::accept(self : Tls) -> Unit {
let read_buf = self.transport.reader._get_internal_buffer().repr()
let write_buf = self.transport.write_buf
for ;; {
let ret = self.conn.accept(
in_buffer=read_buf.buf,
in_buffer_offset=read_buf.start,
in_buffer_len=read_buf.len,
out_buffer=write_buf.buf,
out_buffer_offset=write_buf.start + write_buf.len,
out_buffer_len=write_buf.buf.length() - write_buf.start - write_buf.len,
)
read_buf.drop(self.conn.bytes_read())
write_buf.len += self.conn.bytes_to_write()
match ret {
Completed => {
self.transport.flush_write()
break
}
WantRead => {
self.transport.flush_write()
self.transport.read_more()
}
WantWrite => self.transport.flush_write()
Eof => raise ConnectionClosed
Error => raise TlsError(self.conn.take_error_message())
ReNegotiation => panic()
}
}
}
///|
#label_migration(verify, fill=false, msg="use `trust` instead")
pub async fn[R : @io.Reader, W : @io.Writer] Tls::client_from_pair(
r : R,
w : W,
verify? : Bool = true,
host? : String,
sni? : Bool = true,
trust? : TrustedRoot,
) -> Tls {
let trust = match trust {
Some(trust) => trust
None => if verify { SystemRoot } else { NoVerification }
}
let host_name = match host {
Some(host) => host
None => ""
}
let conn = TlsConnectionHandle::new().check_new()
try {
if trust is CustomPemFile(path) {
let pem = @fs.read_file(path).text()
for cert in decode_pem_certificates(pem) {
let status = conn.add_root_certificate(cert)
if status != 0 {
raise TlsError(conn.take_error_message())
}
}
}
let status = conn.set_client(host_name, sni~, trust.mode())
if status != 0 {
raise TlsError(conn.take_error_message())
}
} catch {
err => {
conn.free()
raise err
}
}
let self = Tls::from_pair(conn, r, w, is_client=true, host~)
try {
self.connect()
self
} catch {
err => {
self.close()
raise err
}
}
}
///|
#internal(internal, "do not use, for internal testing only")
pub async fn[R : @io.Reader, W : @io.Writer] Tls::server_from_pair(
r : R,
w : W,
private_key_file? : String,
private_key_type? : X509FileType,
certificate_file? : String,
certificate_type? : X509FileType,
pfx_file? : String,
) -> Tls {
let conn = TlsConnectionHandle::new().check_new()
try {
let status = if @event_loop.platform is Windows {
let pfx_file = match pfx_file {
Some(pfx_file) => pfx_file
None =>
raise TlsError(
"pfx_file is required for TLS servers on Windows hosts",
)
}
let pfx_content = @fs.read_file(pfx_file).binary()
conn.set_server_pfx(pfx_content)
} else {
let private_key_file = match private_key_file {
Some(private_key_file) => private_key_file
None =>
raise TlsError(
"private_key_file is required for TLS servers on non-Windows hosts",
)
}
let certificate_file = match certificate_file {
Some(certificate_file) => certificate_file
None =>
raise TlsError(
"certificate_file is required for TLS servers on non-Windows hosts",
)
}
let private_key_type = match private_key_type {
Some(private_key_type) => private_key_type
None =>
raise TlsError(
"private_key_type is required for TLS servers on non-Windows hosts",
)
}
let certificate_type = match certificate_type {
Some(certificate_type) => certificate_type
None =>
raise TlsError(
"certificate_type is required for TLS servers on non-Windows hosts",
)
}
conn.set_server_files(
@os_string.encode(private_key_file),
private_key_type,
@os_string.encode(certificate_file),
certificate_type,
)
}
if status != 0 {
raise TlsError(conn.take_error_message())
}
} catch {
err => {
conn.free()
raise err
}
}
let self = Tls::from_pair(conn, r, w, is_client=false, host=None)
try {
self.accept()
self
} catch {
err => {
self.close()
raise err
}
}
}
///|
#internal(internal, "do not use, for internal testing only")
pub async fn[Inner : @io.Reader + @io.Writer] Tls::server(
inner : Inner,
private_key_file? : String,
private_key_type? : X509FileType,
certificate_file? : String,
certificate_type? : X509FileType,
pfx_file? : String,
) -> Tls {
Tls::server_from_pair(
inner,
inner,
private_key_file?,
private_key_type?,
certificate_file?,
certificate_type?,
pfx_file?,
)
}
///|
pub impl @io.Reader for Tls with fn _direct_read(self, buf, offset~, max_len~) {
let read_buf = self.transport.reader._get_internal_buffer().repr()
let write_buf = self.transport.write_buf
while self.conn.read_plain(
in_buffer=read_buf.buf,
in_buffer_offset=read_buf.start,
in_buffer_len=read_buf.len,
out_buffer=write_buf.buf,
out_buffer_offset=write_buf.start + write_buf.len,
out_buffer_len=write_buf.buf.length() -
write_buf.start -
write_buf.len,
plain_buffer=buf,
plain_buffer_offset=offset,
plain_buffer_len=max_len,
)
is n {
let bytes_read = self.conn.bytes_read()
read_buf.drop(bytes_read)
let bytes_to_write = self.conn.bytes_to_write()
write_buf.len += bytes_to_write
if bytes_to_write > 0 {
self.transport.flush_write()
}
match n {
n if n > 0 => return n
TLS_CLOSED => {
self.shutdown()
return 0
}
TLS_RENEGOTIATION =>
if self.is_client {
self.connect()
} else {
self.accept()
}
TLS_WOULD_BLOCK =>
if bytes_read > 0 || bytes_to_write > 0 {
continue
} else if self.transport.state is Closed &&
self.transport.reader._get_internal_buffer().repr().len == 0 {
return 0
} else if self.conn.wants_read() {
self.transport.read_more()
} else if self.conn.wants_write() {
self.transport.flush_write()
} else {
return 0
}
TLS_ERROR => raise TlsError(self.conn.take_error_message())
_ => abort("unexpected status code \{n} from TLS read")
}
} nobreak {
0
}
}
///|
pub impl @io.Reader for Tls with fn _get_internal_buffer(self) {
self.read_buf
}
///|
pub impl @io.Writer for Tls with fn write_once(self, buf, offset~, len~) {
self.transport.flush_write()
let read_buf = self.transport.reader._get_internal_buffer().repr()
let write_buf = self.transport.write_buf
while self.conn.write_plain(
in_buffer=read_buf.buf,
in_buffer_offset=read_buf.start,
in_buffer_len=read_buf.len,
out_buffer=write_buf.buf,
out_buffer_offset=write_buf.start + write_buf.len,
out_buffer_len=write_buf.buf.length() -
write_buf.start -
write_buf.len,
plain_buffer=buf,
plain_buffer_offset=offset,
plain_buffer_len=len,
)
is n {
let bytes_read = self.conn.bytes_read()
read_buf.drop(bytes_read)
let bytes_to_write = self.conn.bytes_to_write()
write_buf.len += bytes_to_write
if bytes_to_write > 0 {
self.transport.flush_write()
}
match n {
n if n >= 0 => return n
TLS_RENEGOTIATION =>
if self.is_client {
self.connect()
} else {
self.accept()
}
TLS_WOULD_BLOCK =>
if bytes_read > 0 || bytes_to_write > 0 {
continue
} else if self.conn.wants_read() {
self.transport.read_more()
} else if self.conn.wants_write() {
return 0
} else {
return 0
}
TLS_CLOSED => raise ConnectionClosed
TLS_ERROR => raise TlsError(self.conn.take_error_message())
_ => abort("unexpected status code \{n} from TLS write")
}
} nobreak {
0
}
}
///|
pub fn Tls::close(self : Tls) -> Unit {
guard !self.closed else { }
self.closed = true
if self.transport.state is Normal {
self.transport.state = Closed
}
self.conn.free()
ignore(self.host)
}
///|
pub async fn Tls::shutdown(self : Tls) -> Unit {
guard !self.shutdown else { }
self.shutdown = true
match self.transport.state {
Normal => ()
Closed => return
Error(err) => raise err
}
self.transport.flush_write()
match self.conn.shutdown() {
0 => ()
TLS_CLOSED => return
TLS_ERROR => raise TlsError(self.conn.take_error_message())
status => abort("unexpected status code \{status} from TLS shutdown")
}
if self.is_client {
self.connect()
} else {
self.accept()
}
}
///|
pub fn Tls::get_peer_certificate(self : Tls) -> Bytes? raise {
let cert = self.conn.peer_certificate()
guard !cert.is_null() else { raise TlsError(self.conn.take_error_message()) }
defer cert.free()
let len = tls_buffer_length(cert)
guard! len > 0
let result = FixedArray::make(len, b'\x00')
cert.blit_to_bytes(dst=result, len~)
Some(result.unsafe_reinterpret_as_bytes())
}
///|
pub fn Tls::unique_channel_binding(self : Tls) -> Bytes raise {
let binding = self.conn.unique_channel_binding()
guard !binding.is_null() else {
raise TlsError(self.conn.take_error_message())
}
defer binding.free()
let len = tls_buffer_length(binding)
guard len > 0 else {
raise TlsError("tls-unique channel binding unavailable")
}
let result = FixedArray::make(len, b'\x00')
binding.blit_to_bytes(dst=result, len~)
result.unsafe_reinterpret_as_bytes()
}
///|
pub fn Tls::server_endpoint_channel_binding(self : Tls) -> Bytes raise {
let binding = self.conn.server_endpoint_channel_binding()
guard !binding.is_null() else {
raise TlsError(self.conn.take_error_message())
}
defer binding.free()
let len = tls_buffer_length(binding)
guard len > 0 else {
raise TlsError("tls-server-endpoint channel binding unavailable")
}
let result = FixedArray::make(len, b'\x00')
binding.blit_to_bytes(dst=result, len~)
result.unsafe_reinterpret_as_bytes()
}
///|
let _unused : Unit = {
ignore(@bytes_util.ascii_to_string)
ignore(@os_error.check_errno)
ignore(@os_string.encode)
ignore((tls : Tls) => tls.is_client)
ignore((transport : Transport) => transport.flush_write())
ignore((_ : @fs.File) => ())
}