Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 16 additions & 1 deletion src/Funogram.Tests/UnionDiscriminators.fs
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
module Funogram.Tests.UnionDiscriminators

open System.Text
open System.Text.Json
open Funogram
open Funogram.Telegram.Types
open Xunit
open Helpers
Expand Down Expand Up @@ -81,4 +84,16 @@ let ``MessageOrigin deserializes to the case matching the type value`` (json: st
| MessageOrigin.Chat _ -> "Chat"
| MessageOrigin.Channel _ -> "Channel"
Assert.Equal(expectedCase, actual)
| Error e -> failwith e.Description
| Error e -> failwith e.Description

[<Fact>]
let ``Strict RichBlock throws on a gibberish discriminator instead of guessing a case`` () =
let json = """{"type":"foo","text":"x"}"""
let bytes = Encoding.UTF8.GetBytes json
Assert.Throws<JsonException>(fun () -> JsonSerializer.Deserialize<RichBlock>(bytes, Tools.strictOptions) |> ignore) |> ignore
Assert.NotNull(box (JsonSerializer.Deserialize<RichBlock>(bytes, Tools.options)))

[<Fact>]
let ``Strict RichBlock throws when a known discriminator is missing a required field`` () =
let bytes = Encoding.UTF8.GetBytes """{"type":"paragraph"}"""
Assert.Throws<JsonException>(fun () -> JsonSerializer.Deserialize<RichBlock>(bytes, Tools.strictOptions) |> ignore) |> ignore
33 changes: 29 additions & 4 deletions src/Funogram/Converters.fs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ open System.Runtime.CompilerServices
open System.Text.Json
open System.Text.Json.Serialization
open Funogram.Types
open Microsoft.FSharp.Reflection
open TypeShape.Core
open TypeShape.Core.SubtypeExtensions

Expand Down Expand Up @@ -98,7 +99,7 @@ module internal Converters =
t.IsArray
|| (t.IsGenericType && t.GetGenericTypeDefinition() = fsharpListTypeDef)

type DiscriminatedUnionConverter<'a>() =
type DiscriminatedUnionConverter<'a>(strict: bool) =
inherit JsonConverter<'a>()

let shape = shapeof<'a>
Expand Down Expand Up @@ -127,6 +128,19 @@ module internal Converters =
|> CaseDescriptor.Object)
|> Seq.toArray

// "required" = reference-typed, non-option property (Nullable<_> is a value type, excluded by IsValueType)
let requiredProps =
union.UnionCases
|> Seq.map (fun c ->
let tp = if c.Fields.Length = 0 then typeof<unit> else c.Fields[0].Member.Type
if c.Fields.Length = 0 || tp.IsPrimitive || tp = typeof<string> || isArrayLike tp then [||]
else
tp.GetProperties()
|> Array.filter (fun p ->
not p.PropertyType.IsValueType
&& not (p.PropertyType.IsGenericType && p.PropertyType.GetGenericTypeDefinition() = typedefof<_ option>)))
|> Seq.toArray

let alwaysFields =
cases
|> Seq.collect (fun x ->
Expand Down Expand Up @@ -181,6 +195,9 @@ module internal Converters =

match exact with
| Some _ -> exact
| None when strict && not requiredFields.IsEmpty ->
let discriminator = requiredFields |> Set.toList |> List.map (fun (n, v) -> $"{n}={v}") |> String.concat ", "
raise (JsonException($"Unable to match JSON to any case of union {typeof<'a>.Name}: no case declares discriminator {discriminator}"))
| None ->
let scored =
cases
Expand Down Expand Up @@ -259,7 +276,13 @@ module internal Converters =

match idx with
| Some i ->
deserializers[i].Deserialize(&reader, options)
let result = deserializers[i].Deserialize(&reader, options)
if strict && requiredProps[i].Length > 0 then
let caseInfo, fields = FSharpValue.GetUnionFields(box result, typeof<'a>)
for p in requiredProps[i] do
if isNull (p.GetValue(fields[0])) then
raise (JsonException($"Missing required field '{p.Name}' for {typeof<'a>.Name} case {caseName caseInfo}"))
result
| None ->
raise (JsonException($"Unable to match JSON to any case of union {typeof<'a>.Name}"))

Expand All @@ -269,12 +292,14 @@ module internal Converters =
| Shape.FSharpUnion _ -> true
| _ -> false

type DiscriminatedUnionConverterFactory() =
type DiscriminatedUnionConverterFactory(?strict: bool) =
inherit JsonConverterFactory()

let strict = defaultArg strict false

override x.CreateConverter(typeToConvert, _) =
let g = typedefof<DiscriminatedUnionConverter<_>>.MakeGenericType(typeToConvert)
Activator.CreateInstance(g) :?> JsonConverter
Activator.CreateInstance(g, [| box strict |]) :?> JsonConverter

override x.CanConvert(typeToConvert) =
match TypeShape.Create(typeToConvert) with
Expand Down
19 changes: 12 additions & 7 deletions src/Funogram/Tools.fs
Original file line number Diff line number Diff line change
Expand Up @@ -76,23 +76,28 @@ module internal RequestLogger =
logger.Text.Append("Res: ").Append(e.ToString()) |> ignore
logger.Logger.Log(logger.Text.ToString())

/// Shared JSON serializer settings used by Funogram for request and response payloads.
///
/// Reuse this instance when serializing or deserializing Funogram types so external code
/// stays aligned with the library's snake_case wire format, union handling, Unix timestamps,
/// and null-skipping behavior.
let options =
let private mkOptions (strict: bool) =
let o =
JsonSerializerOptions(
WriteIndented = false,
PropertyNamingPolicy = JsonNamingPolicy.SnakeCaseLower,
DefaultIgnoreCondition = JsonIgnoreCondition.WhenWritingNull
)
o.Converters.Add(DiscriminatedUnionConverterFactory())
o.Converters.Add(DiscriminatedUnionConverterFactory(strict))
o.Converters.Add(UnixTimestampDateTimeConverter())
o.Converters.Add(OptionConverterFactory())
o

/// Shared JSON serializer settings used by Funogram for request and response payloads.
///
/// Reuse this instance when serializing or deserializing Funogram types so external code
/// stays aligned with the library's snake_case wire format, union handling, Unix timestamps,
/// and null-skipping behavior.
let options = mkOptions false

/// Like `options`, but an unmatched union discriminator or a missing required field raises JsonException instead of guessing a case.
let strictOptions = mkOptions true

let private getUrl (config: BotConfig) methodName =
let botToken = sprintf "%s%s" (config.ApiEndpointUrl |> string) config.Token

Expand Down