diff --git a/.dialyzer_ignore.exs b/.dialyzer_ignore.exs index 125aa827..a30b4699 100644 --- a/.dialyzer_ignore.exs +++ b/.dialyzer_ignore.exs @@ -55,7 +55,7 @@ # defensive guard: ReqLLM.Response types provider_meta as map() with a %{} # default, but the struct does not enforce it (a caller can build one with # nil), and ReqLLM's own OpenTelemetry attributes guard it with is_map/1. - {"lib/imp/clients/req_llm.ex", :guard_fail, 1902}, + {"lib/imp/clients/req_llm.ex", :guard_fail, 1917}, # defensive error clause on an always-ok internal call {"lib/imp/clients/training.ex", :pattern_match, {1215, 13}}, # defensive error clause on an always-ok internal call diff --git a/CHANGELOG.md b/CHANGELOG.md index a915fe99..6116804e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,7 +4,7 @@ User-visible changes to Imp are recorded here. ## Unreleased -### Fixed +### Changed - Breaking: `Imp.Core.LMResponse.cost`, and the `:cost` on a `:model_response` event, is what the provider reported charging, as its @@ -23,6 +23,34 @@ User-visible changes to Imp are recorded here. unknown charge, not a free one; a host that wants the old number for calls with no reported charge reads `estimated_cost` for them, knowing it is an estimate. + +### Fixed + +- A call streamed through `Imp.Clients.ReqLLM` records the usage + and the cost the provider reported at the end of the stream. ReqLLM's + stream is a `Stream.resource`, which reports its end as halted rather than + done, and Imp ended such a stream without its terminal event, so every + streamed call was recorded with no usage and no cost. The stream now ends + with exactly one terminal event, `done: true` with the provider's usage + (including its `"cost"`), model and finish reason, taking the finish + reason, and usage no chunk reported, from ReqLLM's metadata handle. +- A stream that did not complete ends in `{:error, %Imp.LMError{}}`, not in + a completion of whatever text arrived first: one that carries a provider + error, finishes with reason `:error` or `:cancelled`, or is incomplete, + its body ending with no finish and no `[DONE]` (reason + `{:stream_finished, :incomplete}`). The error event carries the usage and + other metadata that arrived before it. +- A streamed call that fails after the provider reported usage records that + usage and cost on its failed `:model_response` event and in + `Imp.Usage`, since the provider may have charged for it. The caller still + receives `{:error, reason}`. +- A streamed call to a client built from an inline spec map, with atom or + string keys, or a `{provider, opts}` tuple no longer fails with + `{:lm_stream_failed, "protocol String.Chars not implemented ..."}`. A + streamed call records the model the provider reported, as a non-streamed + call does, and otherwise the configured model id (`gpt-test`, where it + recorded the whole `openai:gpt-test`), so its `Imp.Usage` key changes the + same way. - `:reasoning_effort` accepts `max`, when an LM is built, on a call and in a saved program. Imp's accepted efforts are read from ReqLLM's own `reasoning_effort` option, so they are every effort ReqLLM accepts, on every diff --git a/lib/imp/clients/req_llm.ex b/lib/imp/clients/req_llm.ex index 68457f73..ca575744 100644 --- a/lib/imp/clients/req_llm.ex +++ b/lib/imp/clients/req_llm.ex @@ -25,6 +25,15 @@ defmodule Imp.Clients.ReqLLM do relays an upstream refusal, is returned as that error, never as an empty completion. + `stream/3` ends with exactly one terminal event. A provider stream that runs + to its end closes with `done: true` and the metadata the provider reported + along the way: usage (with any `"cost"` the provider charged), model and + finish reason. One that raises, carries a provider error, or finishes with + reason `:error` or `:cancelled` closes with `{:error, %Imp.LMError{}}` + instead, never a completion, and that event carries the metadata that + arrived before it stopped. A consumer that stops early receives no terminal + event, and the provider stream is cancelled. + `:reasoning_effort` is the one reasoning option, on the client or on a call. It takes any value of ReqLLM's own `reasoning_effort` option, such as `high`, `xhigh`, `max` or `default`, as an atom or a string, on every provider; @@ -392,26 +401,28 @@ defmodule Imp.Clients.ReqLLM do # says nothing has declined to answer (`Imp.Predict.ReActV2`), so a refused # request would be recorded as a choice. It is the failed request it reports, # in the shape ReqLLM gives an HTTP error. - defp relayed_error(%ReqLLM.Response{provider_meta: %{} = meta}) do - case Map.get(meta, "error") || Map.get(meta, :error) do - %{} = error -> - code = Map.get(error, "code") || Map.get(error, :code) - - %ReqLLM.Error.API.Request{ - reason: Map.get(error, "message") || Map.get(error, :message) || inspect(error), - status: if(is_integer(code), do: code), - response_body: %{"error" => error} - } + defp relayed_error(%ReqLLM.Response{provider_meta: %{} = meta}), + do: provider_error(Map.get(meta, "error") || Map.get(meta, :error)) - message when is_binary(message) and message != "" -> - %ReqLLM.Error.API.Request{reason: message, response_body: %{"error" => message}} + defp relayed_error(_response), do: nil - _none -> - nil - end + # A provider's error object or message, as ReqLLM carries it in a response's + # `provider_meta` or a stream's metadata, in the shape ReqLLM gives an HTTP + # error; `nil` when there is none. + defp provider_error(%{} = error) do + code = Map.get(error, "code") || Map.get(error, :code) + + %ReqLLM.Error.API.Request{ + reason: Map.get(error, "message") || Map.get(error, :message) || inspect(error), + status: if(is_integer(code), do: code), + response_body: %{"error" => error} + } end - defp relayed_error(_response), do: nil + defp provider_error(message) when is_binary(message) and message != "", + do: %ReqLLM.Error.API.Request{reason: message, response_body: %{"error" => message}} + + defp provider_error(_none), do: nil # Every failed request becomes one `Imp.LMError`, classified here, where the # provider library's error shapes are known, so no caller has to know them. @@ -1832,7 +1843,11 @@ defmodule Imp.Clients.ReqLLM do defp model_id(%{"id" => id}) when is_binary(id), do: id defp model_id(%{model: id}) when is_binary(id), do: id defp model_id(%{"model" => id}) when is_binary(id), do: id - defp model_id({_provider, id}) when is_binary(id), do: id + # ReqLLM's two-element tuple is `{provider, opts}`, naming the model in + # `:id` or `:model` (`ReqLLM.model/1`). + defp model_id({_provider, opts}) when is_list(opts), + do: to_string(opts[:id] || opts[:model] || "") + defp model_id({_provider, id, _opts}) when is_binary(id), do: id defp model_id(model) when is_binary(model) do @@ -1914,16 +1929,32 @@ defmodule Imp.Clients.ReqLLM do } end - defp provider_name(%{provider: provider}), do: to_string(provider) + defp provider_name(model_spec), do: model_spec |> model_identity() |> elem(0) - defp provider_name(model_spec) when is_binary(model_spec) do - case String.split(model_spec, ":", parts: 2) do + @doc false + # The provider, as a string, and the model id of any model shape + # `Imp.req_llm/2` accepts: a `"provider:model"` string, a + # `{provider, opts}` or `{provider, model, opts}` tuple, or a spec map + # with atom or string keys. The provider is `nil` when the shape names none. + @spec model_identity(term()) :: {String.t() | nil, String.t()} + def model_identity(model), do: {provider_label(model), model_id(model)} + + defp provider_label(%{provider: provider}) when not is_nil(provider), do: to_string(provider) + + defp provider_label(%{"provider" => provider}) when not is_nil(provider), + do: to_string(provider) + + defp provider_label({provider, _model}) when is_atom(provider), do: to_string(provider) + defp provider_label({provider, _model, _opts}) when is_atom(provider), do: to_string(provider) + + defp provider_label(model) when is_binary(model) do + case String.split(model, ":", parts: 2) do [provider, _model] -> provider _other -> nil end end - defp provider_name(_model_spec), do: nil + defp provider_label(_model), do: nil defp sanitize_usage(nil), do: nil defp sanitize_usage(usage) when is_map(usage), do: sanitize_usage_value(usage) @@ -2134,16 +2165,13 @@ defmodule Imp.Clients.ReqLLM do metadata: metadata }} - {:done, _acc} -> - {[ - %Imp.Streaming.Messages.StreamResponse{ - done: true, - metadata: state.metadata - } - ], %{state | continuation: nil, started?: true, completed?: true}} - - {:halted, _acc} -> - {:halt, %{state | continuation: nil, started?: true, completed?: true}} + # The reducer only ever suspends and is only ever resumed with + # `{:cont, _}`, so neither result means a consumer stopped early: a list + # that runs out reports `{:done, _}`, and a `Stream.resource` (ReqLLM's + # stream is one) that runs out after a suspension reports `{:halted, _}`. + # Both are the provider stream's end. + {finished, _acc} when finished in [:done, :halted] -> + finish_stream(%{state | continuation: nil, started?: true}) end rescue error -> stream_failure(state, error) @@ -2151,6 +2179,75 @@ defmodule Imp.Clients.ReqLLM do kind, reason -> stream_failure(state, {kind, reason}) end + # The end of the provider stream closes with one terminal event carrying + # the metadata accumulated along the way. A stream whose metadata reports an + # error, or a finish reason of `:error`, `:cancelled` or `:incomplete`, did + # not complete, and ends as a failure. ReqLLM's own event projection + # (`ReqLLM.StreamResponse.events/1`) ends the first three the same way; it + # reports `:incomplete` as the reason of a `:finish` event, which Imp does + # not count as a completion. + # + # The chunks carry what the provider sent; ReqLLM's metadata handle carries + # what ReqLLM concluded about the stream as a whole. Its finish reason is + # `:incomplete` when the body ended with no termination event, a stream + # cut short that the chunks alone cannot tell from a finished one, and its + # usage stands in when no chunk reported any. + defp finish_stream(state) do + state = %{state | metadata: merge_handle_metadata(state.metadata, state.response)} + + case stream_end_error(state.metadata) do + nil -> + {[%Imp.Streaming.Messages.StreamResponse{done: true, metadata: state.metadata}], + %{state | completed?: true}} + + error -> + stream_failure(state, error) + end + end + + # The provider stream has ended, so ReqLLM's collection is finishing too; + # the wait is bounded so a handle that never answers cannot hold the + # terminal event, and a handle that fails or has stopped adds nothing. + @metadata_handle_timeout 5_000 + + defp merge_handle_metadata(metadata, %ReqLLM.StreamResponse{metadata_handle: handle}) + when is_pid(handle) do + handle_metadata = + try do + ReqLLM.StreamResponse.MetadataHandle.await(handle, @metadata_handle_timeout) + rescue + _error -> %{} + catch + :exit, _reason -> %{} + end + + metadata + |> put_handle_value(:finish_reason, handle_metadata, :always) + |> put_handle_value(:usage, handle_metadata, :when_missing) + end + + defp merge_handle_metadata(metadata, _response), do: metadata + + defp put_handle_value(metadata, key, handle_metadata, rule) do + case {Map.get(handle_metadata, key), rule, Map.get(metadata, key)} do + {nil, _rule, _current} -> metadata + {value, :always, _current} -> Map.put(metadata, key, value) + {value, :when_missing, nil} -> Map.put(metadata, key, value) + {_value, :when_missing, _current} -> metadata + end + end + + defp stream_end_error(metadata) do + with nil <- provider_error(Map.get(metadata, :error) || Map.get(metadata, "error")) do + case Map.get(metadata, :finish_reason) || Map.get(metadata, "finish_reason") do + reason when reason in [:error, "error"] -> {:stream_finished, :error} + reason when reason in [:cancelled, "cancelled"] -> {:stream_finished, :cancelled} + reason when reason in [:incomplete, "incomplete"] -> {:stream_finished, :incomplete} + _other -> nil + end + end + end + defp suspend_stream(stream) do Enumerable.reduce(stream, {:cont, nil}, fn chunk, _acc -> {:suspend, chunk} end) end @@ -2166,11 +2263,18 @@ defmodule Imp.Clients.ReqLLM do # The stream had opened, so the request reached the provider: sending it # again may be billed again, and repeats chunks the caller already has. + # The terminal event is an error, never a completion, and carries whatever + # metadata (usage, cost, finish reason) arrived before the stream broke. defp stream_failure(state, error) do reason = lm_error(error, true, state.provider) - {[%Imp.Streaming.Messages.StreamResponse{chunk: {:error, reason}, done: true}], - %{state | completed?: true, failed?: true}} + {[ + %Imp.Streaming.Messages.StreamResponse{ + chunk: {:error, reason}, + done: true, + metadata: state.metadata + } + ], %{state | completed?: true, failed?: true}} end defp cleanup_stream(state, lm) do diff --git a/lib/imp/exceptions.ex b/lib/imp/exceptions.ex index d0ae8b79..60bf199a 100644 --- a/lib/imp/exceptions.ex +++ b/lib/imp/exceptions.ex @@ -42,6 +42,11 @@ defmodule Imp.LMError do request because its input is longer than the model accepts. Sending it again unchanged will fail again; a shorter input may not. + An error the provider sends inside a stream arrives without its code, + because ReqLLM's stream decoder keeps only its message, so it has `:status` + `nil` and is retryable as a failed stream, where the same error in a + non-streamed response may carry a status that says otherwise. + `:reason` keeps the provider library's error unchanged for diagnostics. `Imp.Errors.retryable?/1` and `Imp.Errors.context_window_exceeded?/1` read these fields through the wrappers Imp puts around an LM error. diff --git a/lib/imp/lm.ex b/lib/imp/lm.ex index 2b763a4c..ec80a060 100644 --- a/lib/imp/lm.ex +++ b/lib/imp/lm.ex @@ -15,7 +15,9 @@ defmodule Imp.LM do metadata as `:purpose`; it is never part of what is sent. A request streamed through the client's `stream/3` (`Imp.stream/3` with `provider_stream: true`) is recorded the same way, with the usage the provider reported at the end of - the stream. + the stream. A streamed call that fails after the provider reported usage or + cost records them on its failed `:model_response` and in `Imp.Usage`, since + the provider may have charged for it. The response event's metadata carries the money for that call as two numbers, each a non-negative float in USD or `nil`, as on @@ -170,13 +172,16 @@ defmodule Imp.LM do @doc false # Performs `request` with `dispatch`, a function from the request to - # `{:ok, %Imp.Core.LMResponse{}}` or `{:error, reason}`, and records it: + # `{:ok, %Imp.Core.LMResponse{}}`, `{:error, reason}`, or + # `{:error, reason, %Imp.Core.LMResponse{}}` for a failure that carries a + # partial response, and records it: # usage in `Imp.Usage`, and inside an `Imp.Run` the `:tools_sent`, # `:model_request` and `:model_response` events. `request/2` dispatches to # the client's `request/2` or `generate/3`; a provider stream dispatches to # its `stream/3`. How the answer arrives does not change what is recorded. - @spec record(t(), Imp.Core.LMRequest.t(), (Imp.Core.LMRequest.t() -> result)) :: result - when result: {:ok, Imp.Core.LMResponse.t()} | {:error, term()} + @spec record(t(), Imp.Core.LMRequest.t(), (Imp.Core.LMRequest.t() -> dispatched)) :: result + when result: {:ok, Imp.Core.LMResponse.t()} | {:error, term()}, + dispatched: result | {:error, term(), Imp.Core.LMResponse.t()} def record(lm, %Imp.Core.LMRequest{} = request, dispatch) when is_function(dispatch, 1) do if Imp.Run.context() do call_id = Imp.Run.new_event_id("model") @@ -230,16 +235,35 @@ defmodule Imp.LM do ) ) + {:error, error, partial} -> + Imp.Run.emit(:model_response, + error: error, + metadata: + maybe_put_billing( + %{ + model_call_id: call_id, + model: request.config.model, + usage: partial.usage, + cost: partial.cost, + estimated_cost: partial.estimated_cost + }, + partial.billing + ) + ) + {:error, error} -> Imp.Run.emit(:model_response, error: error, metadata: %{model_call_id: call_id}) end - result + caller_result(result) else - perform_request(lm, request, dispatch) + lm |> perform_request(request, dispatch) |> caller_result() end end + defp caller_result({:error, reason, %Imp.Core.LMResponse{}}), do: {:error, reason} + defp caller_result(result), do: result + # A stable name for one tool roster: the SHA-256 of its canonical JSON, with # object keys sorted, so two requests offering the same definitions hash the # same however the terms were built. A request offering no tools has no hash. @@ -276,10 +300,23 @@ defmodule Imp.LM do # A dispatch that raises or throws, a client's `stream/3` failing before it # returns a stream say, fails as `{:lm_failed, lm, reason}` like a raising # `generate/3`, and its `:model_response` is still recorded. + # + # A dispatch that failed after the provider reported usage, a stream broken + # after its usage arrived, returns `{:error, reason, partial}`: the call may + # have been billed, so the partial response's usage and cost are recorded + # like a completed call's, and the caller still receives `{:error, reason}`. defp perform_request(lm, request, dispatch) do - with {:ok, response} <- call_request(fn -> dispatch.(request) end, lm) do - Imp.Usage.maybe_record(Imp.Core.legacy_response(response)) - {:ok, response} + case call_request(fn -> dispatch.(request) end, lm) do + {:ok, response} -> + Imp.Usage.maybe_record(Imp.Core.legacy_response(response)) + {:ok, response} + + {:error, _reason, partial} = failed -> + Imp.Usage.maybe_record(Imp.Core.legacy_response(partial)) + failed + + {:error, _reason} = error -> + error end end @@ -332,6 +369,7 @@ defmodule Imp.LM do defp call_request(fun, lm) do case fun.() do {:ok, %Imp.Core.LMResponse{} = response} -> {:ok, response} + {:error, _reason, %Imp.Core.LMResponse{}} = failed -> failed {:error, _reason} = error -> error other -> {:error, {:invalid_lm_response, lm_name(lm), other}} end diff --git a/lib/imp/streaming/execution.ex b/lib/imp/streaming/execution.ex index 91659f9e..4490061d 100644 --- a/lib/imp/streaming/execution.ex +++ b/lib/imp/streaming/execution.ex @@ -51,11 +51,10 @@ defmodule Imp.Streaming.Execution do |> Imp.LM.record(request, fn request -> {messages, opts} = Imp.Core.request_parts(request) - with {:ok, raw} <- - lm - |> module.stream(messages, opts) - |> consume_stream(context, name, lm) do - Imp.Core.response(raw) + case lm |> module.stream(messages, opts) |> consume_stream(context, name, lm) do + {:ok, raw} -> Imp.Core.response(raw) + {:error, reason, raw} -> partial_failure(reason, raw) + {:error, _reason} = error -> error end end) |> Imp.LM.legacy_result() @@ -64,6 +63,17 @@ defmodule Imp.Streaming.Execution do end end + # A stream that failed after the provider reported usage or cost may still + # have been billed, so the failure carries what arrived as a partial + # response for `Imp.LM.record/3` to record; the caller still gets only the + # reason. + defp partial_failure(reason, raw) do + case Imp.Core.response(raw) do + {:ok, partial} -> {:error, reason, partial} + {:error, _invalid} -> {:error, reason} + end + end + defp lm_module(%module{}), do: module defp lm_module(module), do: module @@ -72,8 +82,8 @@ defmodule Imp.Streaming.Execution do result = Enum.reduce_while(stream, {:ok, [], %{}}, fn - %StreamResponse{chunk: {:error, reason}}, _acc -> - {:halt, {:error, reason}} + %StreamResponse{chunk: {:error, reason}} = event, {:ok, chunks, metadata} -> + {:halt, {:error, reason, chunks, collect_metadata(event, metadata)}} %StreamResponse{} = event, {:ok, chunks, metadata} -> maybe_emit_raw(context, name, event) @@ -92,8 +102,11 @@ defmodule Imp.Streaming.Execution do |> envelope(normalize_metadata(lm, metadata)) |> then(&{:ok, &1}) - {:error, _reason} = error -> - error + {:error, reason, _chunks, metadata} when map_size(metadata) == 0 -> + {:error, reason} + + {:error, reason, chunks, metadata} -> + {:error, reason, envelope(materialize_chunks(chunks), normalize_metadata(lm, metadata))} end rescue error -> {:error, {:lm_stream_failed, Exception.message(error)}} @@ -203,16 +216,14 @@ defmodule Imp.Streaming.Execution do defp envelope(output, metadata), do: %{__imp_lm_output__: output, __imp_lm_metadata__: metadata} + # The model is the one the provider reported, as a non-streamed response + # records it, and otherwise the id the client was configured with. defp normalize_metadata(%Imp.Clients.ReqLLM{model: model}, metadata) do - provider = - case to_string(model) |> String.split(":", parts: 2) do - [provider, _model] -> provider - _other -> nil - end + {provider, model_id} = Imp.Clients.ReqLLM.model_identity(model) req_llm = %{ provider: provider, - model: to_string(model), + model: metadata[:model] || metadata["model"] || model_id, usage: metadata[:usage] || metadata["usage"], finish_reason: metadata[:finish_reason] || metadata["finish_reason"] } diff --git a/lib/mix/tasks/imp.benchmark.trace.ex b/lib/mix/tasks/imp.benchmark.trace.ex index 382c094a..68ff196b 100644 --- a/lib/mix/tasks/imp.benchmark.trace.ex +++ b/lib/mix/tasks/imp.benchmark.trace.ex @@ -669,6 +669,8 @@ defmodule Mix.Tasks.Imp.Benchmark.Trace.StreamReqLLM do @moduledoc false def stream_text(model, messages, _opts) do + {:ok, metadata_handle} = ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end) + {:ok, %ReqLLM.StreamResponse{ stream: [ @@ -677,7 +679,7 @@ defmodule Mix.Tasks.Imp.Benchmark.Trace.StreamReqLLM do ReqLLM.StreamChunk.text("Paris\n\n[[ ## completed ## ]]"), ReqLLM.StreamChunk.meta(%{finish_reason: "stop"}) ], - metadata_handle: self(), + metadata_handle: metadata_handle, cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) diff --git a/test/req_llm_client_test.exs b/test/req_llm_client_test.exs index 82277678..37e7cf9b 100644 --- a/test/req_llm_client_test.exs +++ b/test/req_llm_client_test.exs @@ -60,7 +60,7 @@ defmodule ReqLLMClientTest do ReqLLM.StreamChunk.text(~s(ng","score":7})), ReqLLM.StreamChunk.meta(%{finish_reason: "stop"}) ], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) @@ -110,7 +110,7 @@ defmodule ReqLLMClientTest do finish_reason: "stop" }) ], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) @@ -220,7 +220,7 @@ defmodule ReqLLMClientTest do ReqLLM.StreamChunk.text(~s({"answer":"Paris"})), ReqLLM.StreamChunk.meta(%{finish_reason: "stop"}) ], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) @@ -277,7 +277,7 @@ defmodule ReqLLMClientTest do }, ReqLLM.StreamChunk.meta(%{finish_reason: "tool_calls"}) ], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) @@ -321,7 +321,7 @@ defmodule ReqLLMClientTest do {:ok, %ReqLLM.StreamResponse{ stream: stream, - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> send(test_pid, :provider_cancelled) end, model: model, context: ReqLLM.Context.new(messages) @@ -1593,8 +1593,9 @@ defmodule ReqLLMClientTest do prediction = List.last(events) + # The stream reported no model, so the key names the configured model id. assert Imp.Prediction.get_lm_usage(prediction) == %{ - "openai/openai:gpt-test" => %{ + "openai/gpt-test" => %{ input_tokens: 3, output_tokens: 2, total_tokens: 5 diff --git a/test/req_llm_stream_end_test.exs b/test/req_llm_stream_end_test.exs new file mode 100644 index 00000000..72b2f127 --- /dev/null +++ b/test/req_llm_stream_end_test.exs @@ -0,0 +1,277 @@ +defmodule Imp.ReqLLMStreamEndTest do + use ExUnit.Case, async: true + + alias Imp.Streaming.Messages.StreamResponse + + # These streams come from ReqLLM itself, decoding server-sent events from a + # local server, so they end the way a provider's stream does: as a + # `Stream.resource` that has run out, not as a list. + + @content %{ + "id" => "gen-1", + "model" => "local-model", + "choices" => [%{"index" => 0, "delta" => %{"role" => "assistant", "content" => "pong"}}] + } + + @finish %{ + "id" => "gen-1", + "model" => "local-model", + "choices" => [%{"index" => 0, "delta" => %{}, "finish_reason" => "stop"}] + } + + @usage %{ + "id" => "gen-1", + "model" => "local-model", + "choices" => [], + "usage" => %{ + "prompt_tokens" => 7, + "completion_tokens" => 1, + "total_tokens" => 8, + "cost" => 0.00042 + } + } + + defp sse(events, done? \\ true) do + body = Enum.map_join(events, "", &("data: " <> Jason.encode!(&1) <> "\n\n")) + if done?, do: body <> "data: [DONE]\n\n", else: body + end + + defp lm(body) do + url = Imp.Test.LocalHTTP.start(fn _request -> {200, body} end) + + Imp.req_llm( + %{ + provider: :openrouter, + id: "local-model", + model: "local-model", + base_url: url <> "/api/v1" + }, + api_key: "local-test-key", + cache: false + ) + end + + defp stream(lm), + do: lm |> Imp.Clients.ReqLLM.stream([%{role: :user, content: "ping"}], []) |> Enum.to_list() + + test "a provider stream that runs to its end closes with one done event carrying usage and cost" do + events = stream(lm(sse([@content, @finish, @usage]))) + + assert [%StreamResponse{chunk: "pong", done: false}, %StreamResponse{done: true} = last] = + events + + assert last.chunk == nil + assert last.metadata.finish_reason == :stop + assert %{input_tokens: 7, output_tokens: 1, total_tokens: 8} = last.metadata.usage + assert last.metadata.usage["cost"] == 0.00042 + end + + test "an error the provider sends inside the stream ends it as a failure, with what arrived" do + error = %{"error" => %{"message" => "upstream died", "code" => 502}} + events = stream(lm(sse([@content, @usage, error], false))) + + assert [%StreamResponse{chunk: "pong"}, %StreamResponse{done: true} = last] = events + assert {:error, %Imp.LMError{retryable: true} = reason} = last.chunk + assert reason.message =~ "upstream died" + assert %{input_tokens: 7} = last.metadata.usage + assert last.metadata.usage["cost"] == 0.00042 + refute Enum.any?(events, &match?(%StreamResponse{done: true, chunk: nil}, &1)) + end + + test "a stream whose body stops with no finish and no [DONE] is not a completion" do + events = stream(lm(sse([@content], false))) + + assert [ + %StreamResponse{chunk: "pong"}, + %StreamResponse{ + chunk: {:error, %Imp.LMError{reason: {:stream_finished, :incomplete}}}, + done: true, + metadata: %{finish_reason: :incomplete} + } + ] = events + end + + # A ReqLLM stream built from chunks, with a real metadata handle answering + # what ReqLLM concluded about the stream. + defmodule FinishStub do + def stream_text(model, messages, opts) do + chunks = + Keyword.get_lazy(opts, :chunks, fn -> + [ + ReqLLM.StreamChunk.text("partial"), + ReqLLM.StreamChunk.meta(%{finish_reason: Keyword.fetch!(opts, :finish_reason)}) + ] + end) + + handle_metadata = Keyword.get(opts, :handle_metadata, %{}) + + {:ok, handle} = + ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> handle_metadata end) + + stream = + Stream.resource( + fn -> chunks end, + fn + [] -> {:halt, []} + [chunk | rest] -> {[chunk], rest} + end, + fn _state -> :ok end + ) + + {:ok, + %ReqLLM.StreamResponse{ + stream: stream, + metadata_handle: handle, + cancel: fn -> :ok end, + model: model, + context: ReqLLM.Context.new(messages) + }} + end + end + + test "a stream that finishes cancelled or in error is not a completion" do + for finish_reason <- [:cancelled, :error] do + events = + %{provider: :openai, id: "local-model", model: "local-model"} + |> Imp.req_llm(req_module: FinishStub, finish_reason: finish_reason, cache: false) + |> stream() + + assert [ + %StreamResponse{chunk: "partial"}, + %StreamResponse{ + chunk: {:error, %Imp.LMError{reason: {:stream_finished, ^finish_reason}}}, + done: true, + metadata: %{finish_reason: ^finish_reason} + } + ] = events + end + end + + test "usage only ReqLLM's metadata handle reported is on the done event" do + usage = %{input_tokens: 4, output_tokens: 2, total_tokens: 6} + + events = + %{provider: :openai, id: "local-model", model: "local-model"} + |> Imp.req_llm( + req_module: FinishStub, + chunks: [ReqLLM.StreamChunk.text("pong")], + handle_metadata: %{usage: usage, finish_reason: :stop}, + cache: false + ) + |> stream() + + assert [ + %StreamResponse{chunk: "pong"}, + %StreamResponse{done: true, metadata: %{usage: ^usage, finish_reason: :stop}} + ] = events + end + + test "the client names the provider and model id of every model shape it accepts" do + for model <- [ + "openai:gpt-x", + {:openai, id: "gpt-x"}, + {:openai, model: "gpt-x"}, + {:openai, "gpt-x", []}, + %{provider: :openai, id: "gpt-x"}, + %{"provider" => "openai", "id" => "gpt-x"} + ] do + assert Imp.Clients.ReqLLM.model_identity(model) == {"openai", "gpt-x"}, inspect(model) + end + end + + defmodule Collecting do + @behaviour Imp.Module + defstruct [:program] + + @impl true + def call(%__MODULE__{program: program}, inputs), + do: Imp.collect(program, inputs, provider_stream: true) + end + + test "a streamed call records the usage and cost the provider reported at the end" do + answer = "[[ ## answer ## ]]\nParis\n\n[[ ## completed ## ]]" + content = put_in(@content, ["choices", Access.at(0), "delta", "content"], answer) + program = Imp.predict("question -> answer", lm: lm(sse([content, @finish, @usage]))) + + {:ok, run} = Imp.Run.start(%Collecting{program: program}, %{question: "Capital of France?"}) + assert {:ok, prediction} = Task.await(run.task) + events = Imp.Run.events(run) + Imp.Run.stop(run) + + assert Imp.get(prediction, :answer) == "Paris" + assert [response] = Enum.filter(events, &(&1.kind == :model_response)) + assert %{input_tokens: 7, output_tokens: 1, total_tokens: 8} = response.metadata.usage + assert response.metadata.usage["cost"] == 0.00042 + assert response.metadata.cost == 0.00042 + end + + defp run(program) do + {:ok, run} = Imp.Run.start(%Collecting{program: program}, %{question: "Capital of France?"}) + + {result, usage} = Imp.Usage.track(fn -> Task.await(run.task) end) + events = Imp.Run.events(run) + Imp.Run.stop(run) + {result, usage, events} + end + + test "a stream that fails after the provider reported usage records that spend" do + error = %{"error" => %{"message" => "upstream died"}} + program = Imp.predict("question -> answer", lm: lm(sse([@content, @usage, error], false))) + + assert {{:error, %Imp.LMError{message: message}}, _usage, events} = run(program) + assert message =~ "upstream died" + + assert [response] = Enum.filter(events, &(&1.kind == :model_response)) + assert %Imp.LMError{} = response.error + assert %{input_tokens: 7, output_tokens: 1, total_tokens: 8} = response.metadata.usage + assert response.metadata.cost == 0.00042 + end + + test "a failure that carries a partial response counts its usage and returns the reason" do + lm = Imp.req_llm(%{provider: :openrouter, id: "m", model: "m"}, cache: false) + request = Imp.LM.new_request(lm, [%{role: :user, content: "ping"}], [], "test") + + {:ok, partial} = + Imp.Core.response(%{ + __imp_lm_output__: "po", + __imp_lm_metadata__: %{ + req_llm: %{ + provider: "openrouter", + model: "m", + usage: %{input_tokens: 7, output_tokens: 1} + } + } + }) + + assert {{:error, :broken}, usage} = + Imp.Usage.track(fn -> + Imp.LM.record(lm, request, fn _request -> {:error, :broken, partial} end) + end) + + assert usage == %{"openrouter/m" => %{input_tokens: 7, output_tokens: 1}} + end + + # A `{provider, opts}` tuple and a string-keyed map are model shapes + # ReqLLM accepts; a streamed call through either records its provider and model. + test "a streamed call records provider and model for tuple and string-keyed specs" do + for model <- [{:openai, id: "local-model"}, %{"provider" => "openai", "id" => "local-model"}] do + lm = + Imp.req_llm(model, + req_module: FinishStub, + chunks: [ + ReqLLM.StreamChunk.text("[[ ## answer ## ]]\nParis\n\n[[ ## completed ## ]]"), + ReqLLM.StreamChunk.meta(%{usage: %{input_tokens: 1}, finish_reason: :stop}) + ], + cache: false + ) + + program = Imp.predict("question -> answer", lm: lm) + assert {{:ok, _prediction}, _usage, events} = run(program) + assert [response] = Enum.filter(events, &(&1.kind == :model_response)) + + assert %{provider: "openai", model: "local-model"} = + response.metadata.response.req_llm, + inspect(model) + end + end +end diff --git a/test/streamed_run_record_test.exs b/test/streamed_run_record_test.exs index 48e2b0e1..3f31ac98 100644 --- a/test/streamed_run_record_test.exs +++ b/test/streamed_run_record_test.exs @@ -33,7 +33,7 @@ defmodule Imp.StreamedRunRecordTest do ReqLLM.StreamChunk.text(String.slice(@answer, 20..-1//1)), ReqLLM.StreamChunk.meta(%{usage: @usage, finish_reason: "stop"}) ], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages) @@ -204,7 +204,7 @@ defmodule Imp.StreamedRunRecordTest do [ReqLLM.StreamChunk.text(text)] ++ Enum.map(calls, &ReqLLM.StreamChunk.tool_call(&1.name, &1.arguments, %{id: &1.id})) ++ [ReqLLM.StreamChunk.meta(%{finish_reason: "stop"})], - metadata_handle: self(), + metadata_handle: elem(ReqLLM.StreamResponse.MetadataHandle.start_link(fn -> %{} end), 1), cancel: fn -> :ok end, model: model, context: ReqLLM.Context.new(messages)