(* Copyright (C) 2017--2025 Petter A. Urkedal * * This library is free software; you can redistribute it and/or modify it * under the terms of the GNU Lesser General Public License as published by * the Free Software Foundation, either version 3 of the License, or (at your * option) any later version, with the LGPL-3.0 Linking Exception. * * This library is distributed in the hope that it will be useful, but WITHOUT * ANY WARRANTY; without even the implied warranty of MERCHANTABILITY or * FITNESS FOR A PARTICULAR PURPOSE. See the GNU Lesser General Public * License for more details. * * You should have received a copy of the GNU Lesser General Public License * and the LGPL-3.0 Linking Exception along with this library. If not, see * and , respectively. *) [@@@alert "-caqti_private"] open Caqti_template open Caqti_platform open Printf let dialect = let server_version = Version.of_string_unsafe (Sqlite3.sqlite_version_info ()) in Dialect.create_sqlite ~server_version () let driver_info = Caqti_driver_info.of_dialect dialect type Caqti_connection_sig.driver_connection += Driver_connection of Sqlite3.db let get_uri_bool uri name = (match Uri.get_query_param uri name with | Some ("true" | "yes") -> Some true | Some ("false" | "no") -> Some false | Some _ -> ksprintf invalid_arg "Boolean expected for URI parameter %s." name | None -> None) let get_uri_int uri name = (match Uri.get_query_param uri name with | Some s -> (try Some (int_of_string s) with | Failure _ -> ksprintf invalid_arg "Integer expected for URI parameter %s." name) | None -> None) type Caqti_error.msg += Error_msg of { errcode: Sqlite3.Rc.t; extended_errcode: int option; errmsg: string option; } let cause_of_rc : Sqlite3.Rc.t * int option -> _ = function | CONSTRAINT, Some 1299 -> `Not_null_violation | CONSTRAINT, Some 787 -> `Foreign_key_violation | CONSTRAINT, Some 2067 -> `Unique_violation | CONSTRAINT, Some 275 -> `Check_violation | CONSTRAINT, _ -> `Integrity_constraint_violation__don't_match | NOMEM, _ -> `Out_of_memory | FULL, _ -> `Disk_full | _, _ -> `Unspecified__don't_match let () = let pp ppf = function | Error_msg {errcode; errmsg; extended_errcode; _} -> Format.pp_print_string ppf (match errmsg with | None -> Sqlite3.Rc.to_string errcode | Some errmsg -> errmsg); (match extended_errcode with | None -> Format.pp_print_char ppf '.' | Some erc -> Format.fprintf ppf " (ERC#%d)." erc) | _ -> assert false in let cause = function | Error_msg {errcode; extended_errcode; _} -> cause_of_rc (errcode, extended_errcode) | _ -> assert false in Caqti_error.define_msg ~pp ~cause [%extension_constructor Error_msg] let wrap_rc = (match Sqlite3_shim.extended_errcode_int with | None -> fun ?db errcode -> let errmsg = Option.map Sqlite3.errmsg db in Error_msg {errcode; errmsg; extended_errcode = None} | Some extended_errcode_int -> fun ?db errcode -> let errmsg = Option.map Sqlite3.errmsg db in let extended_errcode = Option.map extended_errcode_int db in Error_msg {errcode; errmsg; extended_errcode}) let data_of_value : type a. a Field_type.t -> a -> Sqlite3.Data.t = fun field_type x -> (match field_type with | Bool -> Sqlite3.Data.INT (if x then 1L else 0L) | Int -> Sqlite3.Data.INT (Int64.of_int x) | Int16 -> Sqlite3.Data.INT (Int64.of_int x) | Int32 -> Sqlite3.Data.INT (Int64.of_int32 x) | Int64 -> Sqlite3.Data.INT x | Float -> Sqlite3.Data.FLOAT x | String -> Sqlite3.Data.TEXT x | Enum _ -> Sqlite3.Data.TEXT x | Octets -> Sqlite3.Data.BLOB x | Pdate -> Sqlite3.Data.TEXT (Conv.iso8601_of_pdate x) | Ptime -> (* This is the suggested time representation according to https://sqlite.org/lang_datefunc.html, and is consistent with current_timestamp. Three subsecond digits are significant. *) let s = Ptime.to_rfc3339 ~space:true ~frac_s:3 ~tz_offset_s:0 x in Sqlite3.Data.TEXT (String.sub s 0 23) | Ptime_span -> Sqlite3.Data.FLOAT (Ptime.Span.to_float_s x)) (* TODO: Check integer ranges? The Int64.to_* functions don't raise. *) let value_of_data : type a. uri: Uri.t -> a Field_type.t -> Sqlite3.Data.t -> a = fun ~uri field_type data -> let to_ptime_span x = (match Ptime.Span.of_float_s x with | Some t -> t | None -> let msg = Caqti_error.Msg "Interval out of range for Ptime.span." in let typ = Row_type.field field_type in Request_utils.raise_decode_rejected ~uri ~typ msg) in let cannot_convert_to ft = let msg = Printf.sprintf "Cannot convert %s to %s." (Sqlite3.Data.to_string_debug data) ft in let typ = Row_type.field field_type in Request_utils.raise_decode_rejected ~uri ~typ (Caqti_error.Msg msg) in (match field_type, data with | Bool, Sqlite3.Data.INT y -> y <> 0L | Bool, _ -> cannot_convert_to "bool" | Int, Sqlite3.Data.INT y -> Int64.to_int y | Int, _ -> cannot_convert_to "int" | Int16, Sqlite3.Data.INT y -> Int64.to_int y | Int16, _ -> cannot_convert_to "int" | Int32, Sqlite3.Data.INT y -> Int64.to_int32 y | Int32, _ -> cannot_convert_to "int32" | Int64, Sqlite3.Data.INT y -> y | Int64, _ -> cannot_convert_to "int64" | Float, Sqlite3.Data.FLOAT y -> y | Float, Sqlite3.Data.INT y -> Int64.to_float y | Float, _ -> cannot_convert_to "float" | String, Sqlite3.Data.TEXT y -> y | String, _ -> cannot_convert_to "string" | Enum _, Sqlite3.Data.TEXT y -> y | Enum _, _ -> cannot_convert_to "enum" | Octets, Sqlite3.Data.BLOB y -> y | Octets, _ -> cannot_convert_to "octets" | Pdate as field_type, Sqlite3.Data.TEXT y -> (match Conv.pdate_of_iso8601 y with | Ok y -> y | Error msg -> let msg = Caqti_error.Msg msg in let typ = Row_type.field field_type in Request_utils.raise_decode_rejected ~uri ~typ msg) | Pdate, _ -> cannot_convert_to "date" | Ptime as field_type, Sqlite3.Data.TEXT y -> (* TODO: Improve parsing. *) (match Conv.ptime_of_rfc3339_utc y with | Ok y -> y | Error msg -> let msg = Caqti_error.Msg msg in let typ = Row_type.field field_type in Request_utils.raise_decode_rejected ~uri ~typ msg) | Ptime, _ -> cannot_convert_to "time" | Ptime_span, Sqlite3.Data.FLOAT x -> to_ptime_span x | Ptime_span, Sqlite3.Data.INT x -> to_ptime_span (Int64.to_float x) | Ptime_span, _ -> cannot_convert_to "time span") let query_quotes q = let rec loop : Query.t -> _ = function | L _ | P _ | E _ -> Fun.id | V (t, v) -> List.cons (data_of_value t v) | Q s -> List.cons (Sqlite3.Data.TEXT s) | S qs -> List_ext.fold loop qs in List.rev (loop q []) let query_string q = let quotes = query_quotes q in let buf = Buffer.create 64 in let iQ = ref 1 in let iP0 = List.length quotes + 1 in let rec loop : Query.t -> _ = function | L s -> Buffer.add_string buf s | V _ -> bprintf buf "?%d" !iQ; incr iQ | Q _ -> bprintf buf "?%d" !iQ; incr iQ | P i -> bprintf buf "?%d" (iP0 + i) | E _ -> assert false | S qs -> List.iter loop qs in loop q; (quotes, Buffer.contents buf) let bind_quotes ~uri ~db stmt oq = let aux j x = (match Sqlite3.bind stmt (j + 1) x with | Sqlite3.Rc.OK -> Ok () | rc -> let typ = Row_type.string in Error (Caqti_error.encode_failed ~uri ~typ (wrap_rc ~db rc))) in List_ext.iteri_r aux oq let encode_null_field ~uri ~db stmt field_type iP = (match Sqlite3.bind stmt (iP + 1) Sqlite3.Data.NULL with | Sqlite3.Rc.OK -> () | rc -> let typ = Row_type.field field_type in Request_utils.raise_encode_failed ~uri ~typ (wrap_rc ~db rc)) let encode_field ~uri ~db stmt field_type field_value iP = let d = data_of_value field_type field_value in (match Sqlite3.bind stmt (iP + 1) d with | Sqlite3.Rc.OK -> () | rc -> let typ = Row_type.field field_type in Request_utils.raise_encode_failed ~uri ~typ (wrap_rc ~db rc)) let encode_param ~uri ~db stmt t = let write_value ~uri ft fv iP = encode_field ~uri ~db stmt ft fv iP; iP + 1 in let write_null ~uri ft iP = encode_null_field ~uri ~db stmt ft iP; iP + 1 in let encode = Request_utils.encode_param ~uri {write_value; write_null} t in fun x iP -> try Ok (encode x iP) with | Caqti_error.Exn (#Caqti_error.call as err) -> Error err let decode_row ~uri ~query row_type = let read_value ~uri ft (stmt, j) = let fv = value_of_data ~uri ft (Sqlite3.column stmt j) in (fv, (stmt, j + 1)) in let skip_null n (stmt, j) = let j' = j + n in let rec check k = k = j' || Sqlite3.column stmt k = Sqlite3.Data.NULL && check (k + 1) in if check j then Some (stmt, j') else None in let field_decoder = {Request_utils.read_value; skip_null} in let decode = Request_utils.decode_row ~uri field_decoder row_type in fun stmt -> let (y, (_, n)) = decode (stmt, 0) in let n' = Sqlite3.data_count stmt in if n = n' then Some y else let msg = sprintf "Decoded only %d of %d fields." n n' in let msg = Caqti_error.Msg msg in Request_utils.raise_response_rejected ~uri ~query msg module Q = struct open Caqti_template.Create let start = static T.(unit -->. unit) "BEGIN" let commit = static T.(unit -->. unit) "COMMIT" let rollback = static T.(unit -->. unit) "ROLLBACK" end module Connect_functor (System : Caqti_platform.System_sig.S) (System_unix : Caqti_platform_unix.System_sig.S with type 'a fiber := 'a System.Fiber.t and type stdenv := System.stdenv) = struct open System open System_unix open System.Fiber.Infix open System_utils.Monad_syntax (System.Fiber) module H = Connection_utils.Make_helpers (System) let driver_info = driver_info module type CONNECTION = Caqti_connection_sig.S with type 'a fiber := 'a Fiber.t and type ('a, 'err) stream := ('a, 'err) Stream.t module Make_connection_base (Connection_arg : sig val subst : Query.subst val uri : Uri.t val db : Sqlite3.db val dynamic_capacity : int end) = struct open Connection_arg let using_db_ref = ref false let using_db f = H.assert_single_use ~what:"SQLite connection" using_db_ref f module Response = struct type ('b, 'm) t = { stmt: Sqlite3.stmt; row_type: 'b Row_type.t; query: string; mutable affected_count: int; mutable has_been_executed: bool; } let returned_count _ = Fiber.return (Error `Unsupported) let affected_count {affected_count; has_been_executed; query; _} = if has_been_executed then Fiber.return (Ok affected_count) else let msg = Caqti_error.Msg "Statement not executed yet, affected_count unavailable." in Fiber.return (Error (Caqti_error.response_rejected ~uri ~query msg)) let run_step response = let ret = Sqlite3.step response.stmt in if not response.has_been_executed then begin response.has_been_executed <- true; response.affected_count <- Sqlite3.changes db end; ret let fetch_row ({stmt; row_type; query; _} as response) = let decode = decode_row ~uri ~query row_type in fun () -> (match run_step response with | Sqlite3.Rc.DONE -> None | Sqlite3.Rc.ROW -> decode stmt | rc -> Request_utils.raise_response_failed ~uri ~query (wrap_rc ~db rc)) let exec ({row_type; query; _} as response) = assert (Row_type.unify row_type Row_type.unit <> None); let retrieve () = (match run_step response with | Sqlite3.Rc.DONE -> Ok () | Sqlite3.Rc.ROW -> let msg = Caqti_error.Msg "Received unexpected row for exec." in Error (Caqti_error.response_rejected ~uri ~query msg) | rc -> Error (Caqti_error.response_failed ~uri ~query (wrap_rc ~db rc))) in Preemptive.detach retrieve () let find resp = let retrieve () = try (match fetch_row resp () with | None -> let msg = Caqti_error.Msg "Received no rows for find." in Error (Caqti_error.response_rejected ~uri ~query:resp.query msg) | Some y -> (match fetch_row resp () with | None -> Ok y | Some _ -> let msg = "Received multiple rows for find." in let msg = Caqti_error.Msg msg in let query = resp.query in Error (Caqti_error.response_rejected ~uri ~query msg))) with Caqti_error.Exn (#Caqti_error.retrieve as err) -> Error err in Preemptive.detach retrieve () let find_opt resp = let retrieve () = try (match fetch_row resp () with | None -> Ok None | Some y -> (match fetch_row resp () with | None -> Ok (Some y) | Some _ -> let msg = "Received multiple rows for find_opt." in let msg = Caqti_error.Msg msg in let query = resp.query in Error (Caqti_error.response_rejected ~uri ~query msg))) with Caqti_error.Exn (#Caqti_error.retrieve as err) -> Error err in Preemptive.detach retrieve () let fold f resp acc = let fetch = fetch_row resp in let retrieve acc = let rec loop acc = (match fetch () with | None -> Ok acc | Some y -> loop (f y acc)) in try loop acc with | Caqti_error.Exn (#Caqti_error.retrieve as err) -> Error err in Preemptive.detach retrieve acc let fold_s f resp acc = let fetch = fetch_row resp in let retrieve acc = let rec loop acc = (match fetch () with | None -> Ok acc | Some y -> (match Preemptive.run_in_main (fun () -> f y acc) with | Ok acc -> loop acc | Error _ as r -> r)) in try loop acc with | Caqti_error.Exn (#Caqti_error.retrieve as err) -> Error err in Preemptive.detach retrieve acc let iter_s f resp = let fetch = fetch_row resp in let retrieve () = let rec loop () = (match fetch () with | None -> Ok () | Some y -> (match Preemptive.run_in_main (fun () -> f y) with | Ok () -> loop () | Error _ as r -> r)) in try loop () with | Caqti_error.Exn (#Caqti_error.retrieve as err) -> Error err in Preemptive.detach retrieve () let to_stream resp = let fetch = fetch_row resp in let rec seq () = (match fetch () with | None -> Stream.Nil | Some y -> Stream.Cons (y, Preemptive.detach seq) | exception Caqti_error.Exn (#Caqti_error.retrieve as err) -> Stream.Error err) in Preemptive.detach seq end module Pcache = Request_cache.Make (struct type t = Sqlite3.stmt * Sqlite3.Data.t list * string let weight _ = 1 end) let pcache = Pcache.create ~dynamic_capacity dialect let prepare req = let prepare_helper query = try let stmt = Sqlite3.prepare db query in (match Sqlite3.prepare_tail stmt with | None -> Ok stmt | Some stmt -> Ok stmt) with Sqlite3.Error msg -> let msg = Caqti_error.Msg msg in Error (Caqti_error.request_failed ~uri ~query msg) in let templ = Request.query req dialect in let templ = Query.expand subst templ in let quotes, query = query_string templ in Preemptive.detach prepare_helper query >|=? fun stmt -> (stmt, quotes, query) let pp_request_with_param ppf = Request.make_pp_with_param ~subst ~dialect () ppf let free_prepared (stmt, _, query) = (match Sqlite3.finalize stmt with | Sqlite3.Rc.OK -> Ok () | rc -> let query = sprintf "DEALLOCATE (%s)" query in Error (Caqti_error.request_failed ~uri ~query (wrap_rc ~db rc)) | exception Sqlite3.Error msg -> let query = sprintf "DEALLOCATE (%s)" query in let msg = Caqti_error.Msg msg in Error (Caqti_error.request_failed ~uri ~query msg)) let deallocate req = using_db @@ fun () -> (match Request.prepare_policy req with | Dynamic | Static -> (match Pcache.deallocate pcache req with | None -> Fiber.return (Ok ()) | Some (entry, commit) -> Preemptive.detach free_prepared entry >|=? commit) | Direct -> failwith "deallocate called on oneshot request") let deallocate_some () = let rec loop = function | [] -> Ok () | pcache_entry :: orphans -> (match free_prepared pcache_entry with | Ok () -> loop orphans | r -> r) in let orphans, commit = Pcache.trim pcache in Preemptive.detach loop orphans >|= Result.map commit let call ~f req param = using_db @@ fun () -> deallocate_some () >>=? fun () -> Log.debug ~src:Logging.request_log_src (fun f -> f "Sending %a" pp_request_with_param (req, param)) >>= fun () -> let param_type = Request.param_type req in let row_type = Request.row_type req in (match Request.prepare_policy req with | Direct -> prepare req | Dynamic | Static -> (match Pcache.find_and_promote pcache req with | Some pcache_entry -> Fiber.return (Ok pcache_entry) | None -> prepare req >|=? fun pcache_entry -> Pcache.add pcache req pcache_entry; pcache_entry)) >>=? fun (stmt, quotes, query) -> (* CHECKME: Does binding involve IO? *) Fiber.return (bind_quotes ~uri ~db stmt quotes) >>=? fun () -> let nQ = List.length quotes in (match encode_param ~uri ~db stmt param_type param nQ with | Ok nQP -> let nP = Row_type.length param_type in if nQP > nQ + nP then failwith "Too many arguments passed to query; \ check that the parameter type is correct." else if nQP < nQ + nP then failwith "Too few arguments passed to query; \ check that the parameter type is correct." else let resp = Response.{ stmt; query; row_type; has_been_executed=false; affected_count = -1; } in Fiber.return (Ok resp) | Error _ as r -> Fiber.return r) >>=? fun resp -> (* CHECKME: Does finalize or reset involve IO? *) let cleanup () = (match Request.prepare_policy req with | Direct -> (match Sqlite3.finalize stmt with | Sqlite3.Rc.OK -> Fiber.return () | rc -> Log.warn (fun p -> p "Ignoring error %s when finalizing statement." (Sqlite3.Rc.to_string rc))) | Dynamic | Static -> (match Sqlite3.reset stmt with | Sqlite3.Rc.OK -> Fiber.return () | _ -> Log.warn (fun p -> p "Dropping cache statement due to error.") >|= fun () -> Pcache.remove_and_discard pcache req)) in Fiber.finally (fun () -> f resp) cleanup let disconnect () = using_db @@ fun () -> let finalize_error_count = ref 0 in let not_busy = ref false in Preemptive.detach begin fun () -> let uncache (stmt, _, _) = (match Sqlite3.finalize stmt with | Sqlite3.Rc.OK -> () | _ -> finalize_error_count := !finalize_error_count + 1) in Pcache.iter uncache pcache; not_busy := Sqlite3.db_close db (* If this reports busy, it means we missed an Sqlite3.finalize or other * cleanup action, so this should not happen. *) end () >>= fun () -> (if !finalize_error_count = 0 then Fiber.return () else Log.warn (fun p -> p "Finalization of %d during disconnect return error." !finalize_error_count)) >>= fun () -> (if !not_busy then Fiber.return () else Log.warn (fun p -> p "Sqlite reported still busy when closing handle.")) let validate () = Fiber.return true let check f = f true let exec q p = call ~f:Response.exec q p let start () = exec Q.start () let commit () = exec Q.commit () let rollback () = exec Q.rollback () let set_statement_timeout _ = Fiber.return (Ok ()) end let setup ~config db = let tweaks_version = Caqti_connect_config.(get tweaks_version) config in if tweaks_version < (1, 8) then Fiber.return () else Preemptive.detach (Sqlite3.exec db) "PRAGMA foreign_keys = ON" >>= fun rc -> if rc = Sqlite3.Rc.OK then Fiber.return () else Log.warn (fun f -> f "Could not turn on foreign key support: %s" (Sqlite3.Rc.to_string rc)) let connect ~sw:_ ~stdenv:_ ~subst ~config uri = try (* Check URI and extract parameters. *) assert (Uri.scheme uri = Some "sqlite3"); (match Uri.userinfo uri, Uri.host uri with | None, (None | Some "") -> () | _ -> invalid_arg "Sqlite URI cannot contain user or host components."); let mode = (match get_uri_bool uri "write", get_uri_bool uri "create" with | Some false, Some true -> invalid_arg "Create mode presumes write." | (Some false), (Some false | None) -> Some `READONLY | (Some true | None), (Some true | None) -> None | (Some true | None), (Some false) -> Some `NO_CREATE) in let busy_timeout = get_uri_int uri "busy_timeout" in (* Connect, configure, wrap. *) Preemptive.detach (fun () -> Sqlite3.db_open ~mutex:`FULL ?mode (Uri.path uri |> Uri.pct_decode)) () >>= fun db -> setup ~config db >|= fun () -> (match busy_timeout with | None -> () | Some timeout -> Sqlite3.busy_timeout db timeout); let module Arg = struct let subst = subst dialect let uri = uri let db = db let dynamic_capacity = Caqti_connect_config.(get dynamic_prepare_capacity) config end in let module Connection_base = Make_connection_base (Arg) in let module Connection = struct let driver_info = driver_info let dialect = dialect let driver_connection = Some (Driver_connection db) include Connection_base include Connection_utils.Make_convenience (System) (Connection_base) include Connection_utils.Make_populate (System) (Connection_base) end in Ok (module Connection : CONNECTION) with | Invalid_argument msg -> Fiber.return (Error (Caqti_error.connect_rejected ~uri (Caqti_error.Msg msg))) | Sqlite3.Error msg -> Fiber.return (Error (Caqti_error.connect_failed ~uri (Caqti_error.Msg msg))) end let () = Caqti_platform_unix.Driver_loader.register "sqlite3" (module Connect_functor)