Packages

LiveView-style runtime for Gleam.

Current section

Files

Jump to
lightspeed src lightspeed transport wisp_websocket.gleam
Raw

src/lightspeed/transport/wisp_websocket.gleam

//// Wisp-style WebSocket adapter for Lightspeed session transport.
import lightspeed/agent/session
import lightspeed/diff
import lightspeed/protocol
import lightspeed/transport/contract
/// WebSocket connect request context.
pub type WebSocketRequest {
WebSocketRequest(
route: String,
csrf_token: String,
origin: String,
now_ms: Int,
)
}
/// Adapter state containing one bound session actor.
pub type AdapterState {
AdapterState(session: session.Session, owner: String, csrf_token: String)
}
/// Connect result for WebSocket upgrade.
pub type ConnectResult {
Connected(state: AdapterState, outbound_frames: List(String))
Rejected(error: contract.AdapterError)
}
/// Receive result after processing one client frame.
pub type ReceiveResult {
ReceiveResult(state: AdapterState, outbound_frames: List(String))
}
/// Authenticate and mount a WebSocket session.
pub fn connect(
session_state: session.Session,
request: WebSocketRequest,
auth_hook: contract.AuthHook,
) -> ConnectResult {
connect_with_hooks(
session_state,
request,
auth_hook,
contract.allow_protection(),
)
}
/// Authenticate, enforce protection hooks, and mount a WebSocket session.
pub fn connect_with_hooks(
session_state: session.Session,
request: WebSocketRequest,
auth_hook: contract.AuthHook,
protection_hook: contract.ProtectionHook,
) -> ConnectResult {
case authenticate_owner(session_state, request, auth_hook) {
Error(error) -> Rejected(error: error)
Ok(owner) ->
case ensure_protection(session_state, owner, request, protection_hook) {
Error(error) -> Rejected(error: error)
Ok(Nil) -> {
let connected =
session.handle(
session_state,
session.InboxMessage(
owner: owner,
event: session.Connect(
route: request.route,
csrf_token: request.csrf_token,
now_ms: request.now_ms,
),
),
)
let #(connected, outbox) = session.flush_outbox(connected)
Connected(
state: AdapterState(
session: connected,
owner: owner,
csrf_token: request.csrf_token,
),
outbound_frames: [
protocol.encode(protocol.hello()),
..encode_outbox(outbox)
],
)
}
}
}
}
/// Authenticate and reconnect a WebSocket session.
pub fn reconnect(
state: AdapterState,
request: WebSocketRequest,
auth_hook: contract.AuthHook,
) -> ConnectResult {
reconnect_with_hooks(state, request, auth_hook, contract.allow_protection())
}
/// Authenticate, enforce protection hooks, and reconnect a WebSocket session.
pub fn reconnect_with_hooks(
state: AdapterState,
request: WebSocketRequest,
auth_hook: contract.AuthHook,
protection_hook: contract.ProtectionHook,
) -> ConnectResult {
case authenticate_owner(state.session, request, auth_hook) {
Error(error) -> Rejected(error: error)
Ok(owner) ->
case ensure_protection(state.session, owner, request, protection_hook) {
Error(error) -> Rejected(error: error)
Ok(Nil) -> {
let reconnected =
session.handle(
state.session,
session.InboxMessage(
owner: owner,
event: session.Reconnect(
route: request.route,
now_ms: request.now_ms,
),
),
)
let #(reconnected, outbox) = session.flush_outbox(reconnected)
Connected(
state: AdapterState(
session: reconnected,
owner: owner,
csrf_token: request.csrf_token,
),
outbound_frames: [
protocol.encode(protocol.hello()),
..encode_outbox(outbox)
],
)
}
}
}
}
/// Receive one client frame payload and return outbound server frames.
pub fn receive(
state: AdapterState,
payload: String,
now_ms: Int,
) -> ReceiveResult {
receive_with_hooks(state, payload, now_ms, contract.allow_rate_limit())
}
/// Receive one client frame payload with a rate-limit hook.
pub fn receive_with_hooks(
state: AdapterState,
payload: String,
now_ms: Int,
rate_limit_hook: contract.RateLimitHook,
) -> ReceiveResult {
case protocol.decode(payload) {
Error(decode_error) ->
failure_only(
state,
protocol.Failure(
ref: "",
reason: contract.error_to_string(contract.ProtocolDecodeFailed(
decode_error,
)),
),
)
Ok(frame) ->
case frame {
protocol.Ack(ref) -> apply_and_flush(state, session.Ack(ref: ref))
protocol.Event(ref, name, payload) ->
apply_event_frame(state, ref, name, payload, now_ms, rate_limit_hook)
protocol.Hello(_, _) ->
failure_only(
state,
protocol.Failure(
ref: "",
reason: contract.error_to_string(contract.UnsupportedClientFrame(
"hello",
)),
),
)
protocol.Diff(ref, _) ->
failure_only(
state,
protocol.Failure(
ref: ref,
reason: contract.error_to_string(contract.UnsupportedClientFrame(
"diff",
)),
),
)
protocol.Failure(ref, _) ->
failure_only(
state,
protocol.Failure(
ref: ref,
reason: contract.error_to_string(contract.UnsupportedClientFrame(
"failure",
)),
),
)
}
}
}
/// Access session state.
pub fn session_state(state: AdapterState) -> session.Session {
state.session
}
fn authenticate_owner(
session_state: session.Session,
request: WebSocketRequest,
auth_hook: contract.AuthHook,
) -> Result(String, contract.AdapterError) {
let context =
contract.AuthContext(
session_id: session.id(session_state),
route: request.route,
csrf_token: request.csrf_token,
origin: request.origin,
)
case contract.authenticate(auth_hook, context) {
contract.Denied(reason) -> Error(contract.AuthenticationFailed(reason))
contract.Authorized(owner) ->
case owner == session.owner(session_state) {
True -> Ok(owner)
False -> Error(contract.InvalidAdapterState("owner_mismatch"))
}
}
}
fn apply_event_frame(
state: AdapterState,
ref: String,
name: String,
payload: String,
now_ms: Int,
rate_limit_hook: contract.RateLimitHook,
) -> ReceiveResult {
case enforce_rate_limit(state, name, now_ms, rate_limit_hook) {
Error(error) ->
failure_only(
state,
protocol.Failure(ref: ref, reason: contract.error_to_string(error)),
)
Ok(Nil) ->
case name {
"increment" -> apply_and_flush(state, session.Increment)
"decrement" -> apply_and_flush(state, session.Decrement)
"heartbeat" -> apply_and_flush(state, session.Heartbeat(now_ms: now_ms))
"shutdown" -> apply_and_flush(state, session.Shutdown(reason: payload))
_ ->
failure_only(
state,
protocol.Failure(
ref: ref,
reason: contract.error_to_string(contract.UnsupportedClientEvent(
name,
)),
),
)
}
}
}
fn apply_and_flush(
state: AdapterState,
event: session.InboxEvent,
) -> ReceiveResult {
let next =
session.handle(
state.session,
session.InboxMessage(owner: state.owner, event: event),
)
let #(next, outbox) = session.flush_outbox(next)
ReceiveResult(
state: AdapterState(
session: next,
owner: state.owner,
csrf_token: state.csrf_token,
),
outbound_frames: encode_outbox(outbox),
)
}
fn failure_only(state: AdapterState, frame: protocol.Frame) -> ReceiveResult {
ReceiveResult(state: state, outbound_frames: [protocol.encode(frame)])
}
fn encode_outbox(outbox: List(session.OutboxMessage)) -> List(String) {
case outbox {
[] -> []
[entry, ..rest] ->
case entry {
session.OutboxPatch(patch) -> [
protocol.encode(patch_to_frame(patch)),
..encode_outbox(rest)
]
session.OutboxTelemetry(_) -> encode_outbox(rest)
}
}
}
fn patch_to_frame(patch: session.PatchEnvelope) -> protocol.Frame {
let ref = session.patch_ref(patch)
let html = diff.encode(session.patch(patch))
protocol.Diff(ref: ref, html: html)
}
fn ensure_protection(
session_state: session.Session,
owner: String,
request: WebSocketRequest,
protection_hook: contract.ProtectionHook,
) -> Result(Nil, contract.AdapterError) {
let context =
contract.ProtectionContext(
session_id: session.id(session_state),
owner: owner,
route: request.route,
csrf_token: request.csrf_token,
origin: request.origin,
)
case contract.protect(protection_hook, context) {
contract.Protected -> Ok(Nil)
contract.Rejected(reason) -> Error(contract.ProtectionRejected(reason))
}
}
fn enforce_rate_limit(
state: AdapterState,
event_name: String,
now_ms: Int,
rate_limit_hook: contract.RateLimitHook,
) -> Result(Nil, contract.AdapterError) {
let context =
contract.RateLimitContext(
session_id: session.id(state.session),
owner: state.owner,
event_name: event_name,
now_ms: now_ms,
)
case contract.limit_rate(rate_limit_hook, context) {
contract.RateAllowed -> Ok(Nil)
contract.Limited(reason, retry_after_ms) ->
Error(contract.RateLimited(reason: reason, retry_after_ms: retry_after_ms))
}
}