ai-lite 0.6.1 → 1.0.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/CHANGELOG.md +18 -0
- data/README.md +368 -8
- data/lib/ai_lite/gemini/chat.rb +78 -0
- data/lib/ai_lite/gemini/configuration.rb +19 -0
- data/lib/ai_lite/gemini/models.rb +102 -0
- data/lib/ai_lite/gemini/streaming.rb +148 -0
- data/lib/ai_lite/gemini/youtube.rb +36 -0
- data/lib/ai_lite/gemini.rb +92 -0
- data/lib/ai_lite/openai/configuration.rb +51 -0
- data/lib/ai_lite/openai/models.rb +96 -0
- data/lib/ai_lite/openai/streaming.rb +130 -0
- data/lib/ai_lite/openai.rb +607 -0
- data/lib/ai_lite/version.rb +1 -1
- data/lib/ai_lite.rb +32 -624
- data/test/ai_lite_test.rb +196 -72
- data/test/gemini_models_test.rb +112 -0
- data/test/gemini_streaming_test.rb +166 -0
- data/test/gemini_test.rb +200 -0
- data/test/gemini_youtube_test.rb +57 -0
- data/test/openai_models_test.rb +138 -0
- data/test/openai_streaming_test.rb +189 -0
- data/test/run.rb +3 -0
- metadata +27 -6
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
class AiLite
|
|
2
|
+
class Gemini
|
|
3
|
+
module Streaming
|
|
4
|
+
def chat_stream(message, model: nil, instructions: nil, previous_response_id: nil, max_output_tokens: nil, debug: false, options: {}, &block)
|
|
5
|
+
raise ArgumentError, "chat_stream requires a block" unless block
|
|
6
|
+
|
|
7
|
+
status = "unknown"
|
|
8
|
+
text = +""
|
|
9
|
+
response_id = nil
|
|
10
|
+
events = debug ? [] : nil
|
|
11
|
+
interaction = nil
|
|
12
|
+
raw = nil
|
|
13
|
+
callback_error = nil
|
|
14
|
+
begin
|
|
15
|
+
payload = chat_payload(message, model: model, instructions: instructions,
|
|
16
|
+
previous_response_id: previous_response_id,
|
|
17
|
+
max_output_tokens: max_output_tokens, options: options, stream: true)
|
|
18
|
+
uri = URI.parse("#{API_BASE_URL}/interactions")
|
|
19
|
+
request = Net::HTTP::Post.new(uri)
|
|
20
|
+
headers.each { |key, value| request[key] = value }
|
|
21
|
+
request["Accept"] = "text/event-stream"
|
|
22
|
+
request.body = JSON.generate(payload)
|
|
23
|
+
result = nil
|
|
24
|
+
step_types = {}
|
|
25
|
+
Net::HTTP.start(uri.host, uri.port, use_ssl: uri.scheme == "https") do |http|
|
|
26
|
+
http.open_timeout = timeout
|
|
27
|
+
http.read_timeout = timeout
|
|
28
|
+
catch(:ai_lite_gemini_finished) do
|
|
29
|
+
http.request(request) do |response|
|
|
30
|
+
status = response.code.to_i
|
|
31
|
+
unless status.between?(200, 299)
|
|
32
|
+
raw = +""
|
|
33
|
+
response.read_body { |chunk| raw << chunk }
|
|
34
|
+
raw = JSON.parse(raw)
|
|
35
|
+
result = result_envelope(status: status, error: error_message(raw), raw: raw, debug: debug)
|
|
36
|
+
throw :ai_lite_gemini_finished
|
|
37
|
+
end
|
|
38
|
+
each_gemini_stream_event(response) do |event|
|
|
39
|
+
events << event if events
|
|
40
|
+
raw = { "interaction" => interaction, "events" => events } if debug
|
|
41
|
+
case event["event_type"]
|
|
42
|
+
when "interaction.created"
|
|
43
|
+
interaction = event.fetch("interaction")
|
|
44
|
+
response_id = interaction["id"] || response_id
|
|
45
|
+
when "step.start"
|
|
46
|
+
step_types[event.fetch("index")] = event.fetch("step").fetch("type")
|
|
47
|
+
when "step.delta"
|
|
48
|
+
delta = event.fetch("delta")
|
|
49
|
+
if step_types[event["index"]] == "model_output" && delta["type"] == "text"
|
|
50
|
+
chunk = delta.fetch("text")
|
|
51
|
+
text << chunk
|
|
52
|
+
begin
|
|
53
|
+
block.call(chunk)
|
|
54
|
+
rescue StandardError => e
|
|
55
|
+
callback_error = e
|
|
56
|
+
raise
|
|
57
|
+
end
|
|
58
|
+
end
|
|
59
|
+
when "interaction.status_update"
|
|
60
|
+
response_id = event["interaction_id"] || response_id
|
|
61
|
+
if %w[failed cancelled incomplete requires_action].include?(event["status"])
|
|
62
|
+
result = result_envelope(status: status, content: text.empty? ? nil : text,
|
|
63
|
+
response_id: response_id, error: "Gemini interaction #{event['status']}", raw: raw, debug: debug)
|
|
64
|
+
throw :ai_lite_gemini_finished
|
|
65
|
+
end
|
|
66
|
+
when "interaction.completed"
|
|
67
|
+
interaction = event.fetch("interaction")
|
|
68
|
+
response_id = interaction["id"] || response_id
|
|
69
|
+
raw = { "interaction" => interaction, "events" => events } if debug
|
|
70
|
+
error = if interaction["error"]
|
|
71
|
+
error_message(interaction)
|
|
72
|
+
elsif interaction["status"] != "completed"
|
|
73
|
+
"Gemini interaction #{interaction['status'] || 'has no status'}"
|
|
74
|
+
end
|
|
75
|
+
content = text.empty? ? nil : text
|
|
76
|
+
unless error || content.nil?
|
|
77
|
+
begin
|
|
78
|
+
content = JSON.parse(text.strip)
|
|
79
|
+
rescue JSON::ParserError
|
|
80
|
+
content = text.strip
|
|
81
|
+
end
|
|
82
|
+
end
|
|
83
|
+
result = result_envelope(status: status, content: content, response_id: response_id,
|
|
84
|
+
error: error, raw: raw, debug: debug)
|
|
85
|
+
throw :ai_lite_gemini_finished
|
|
86
|
+
when "error"
|
|
87
|
+
result = result_envelope(status: status, content: text.empty? ? nil : text,
|
|
88
|
+
response_id: response_id, error: error_message(event), raw: raw, debug: debug)
|
|
89
|
+
throw :ai_lite_gemini_finished
|
|
90
|
+
end
|
|
91
|
+
end
|
|
92
|
+
end
|
|
93
|
+
end
|
|
94
|
+
end
|
|
95
|
+
result || result_envelope(status: status, content: text.empty? ? nil : text, response_id: response_id,
|
|
96
|
+
error: "Stream ended before a terminal interaction event", raw: raw, debug: debug)
|
|
97
|
+
rescue StandardError => e
|
|
98
|
+
raise if callback_error.equal?(e)
|
|
99
|
+
|
|
100
|
+
result_envelope(status: status, content: text.empty? ? nil : text, response_id: response_id,
|
|
101
|
+
error: e.message, raw: raw, debug: debug)
|
|
102
|
+
end
|
|
103
|
+
end
|
|
104
|
+
|
|
105
|
+
private
|
|
106
|
+
|
|
107
|
+
def each_gemini_stream_event(response)
|
|
108
|
+
buffer = +"".b
|
|
109
|
+
data = []
|
|
110
|
+
consume = lambda do
|
|
111
|
+
while (boundary = buffer.index(/[\r\n]/))
|
|
112
|
+
# A trailing CR might be the first byte of a CRLF split across reads.
|
|
113
|
+
break if buffer.getbyte(boundary) == 13 && boundary == buffer.bytesize - 1
|
|
114
|
+
|
|
115
|
+
width = buffer.byteslice(boundary, 2) == "\r\n" ? 2 : 1
|
|
116
|
+
line = buffer.slice!(0, boundary + width).byteslice(0, boundary)
|
|
117
|
+
if line.empty?
|
|
118
|
+
unless data.empty?
|
|
119
|
+
payload = data.join("\n").force_encoding(Encoding::UTF_8)
|
|
120
|
+
data.clear
|
|
121
|
+
unless payload == "[DONE]"
|
|
122
|
+
event = JSON.parse(payload)
|
|
123
|
+
raise "Invalid streaming event: expected an object" unless event.is_a?(Hash)
|
|
124
|
+
yield event
|
|
125
|
+
end
|
|
126
|
+
end
|
|
127
|
+
elsif line.start_with?("data:")
|
|
128
|
+
data << line.sub(/\Adata: ?/, "")
|
|
129
|
+
end
|
|
130
|
+
end
|
|
131
|
+
end
|
|
132
|
+
response.read_body do |chunk|
|
|
133
|
+
buffer << chunk.b
|
|
134
|
+
consume.call
|
|
135
|
+
end
|
|
136
|
+
# A final lone CR is a complete line delimiter at EOF.
|
|
137
|
+
if buffer.end_with?("\r")
|
|
138
|
+
buffer << "\n"
|
|
139
|
+
consume.call
|
|
140
|
+
end
|
|
141
|
+
end
|
|
142
|
+
|
|
143
|
+
end
|
|
144
|
+
|
|
145
|
+
include Streaming
|
|
146
|
+
private_constant :Streaming
|
|
147
|
+
end
|
|
148
|
+
end
|
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
class AiLite
|
|
2
|
+
class Gemini
|
|
3
|
+
module YouTube
|
|
4
|
+
def youtube(url, prompt:, **kwargs)
|
|
5
|
+
unless prompt.is_a?(String) && !prompt.strip.empty?
|
|
6
|
+
raise ArgumentError, "prompt must be a nonempty string"
|
|
7
|
+
end
|
|
8
|
+
uri = URI.parse(url) if url.is_a?(String)
|
|
9
|
+
unless uri.is_a?(URI::HTTPS) && !uri.userinfo && uri.port == 443
|
|
10
|
+
raise ArgumentError, "Expected an HTTPS YouTube video URL"
|
|
11
|
+
end
|
|
12
|
+
video_id = case uri.host&.downcase
|
|
13
|
+
when "youtu.be"
|
|
14
|
+
uri.path.delete_prefix("/")
|
|
15
|
+
when "youtube.com", "www.youtube.com", "m.youtube.com"
|
|
16
|
+
if uri.path == "/watch"
|
|
17
|
+
ids = URI.decode_www_form(uri.query.to_s).select { |key, _| key == "v" }
|
|
18
|
+
ids.first[1] if ids.length == 1
|
|
19
|
+
else
|
|
20
|
+
uri.path.match(%r{\A/(?:shorts|embed|live)/([^/]+)\z})&.captures&.first
|
|
21
|
+
end
|
|
22
|
+
end
|
|
23
|
+
unless video_id && video_id.match?(/\A[A-Za-z0-9_-]{11}\z/)
|
|
24
|
+
raise ArgumentError, "Expected a YouTube video URL with a valid video ID"
|
|
25
|
+
end
|
|
26
|
+
|
|
27
|
+
chat([{ type: "video", uri: url }, { type: "text", text: prompt }], **kwargs)
|
|
28
|
+
rescue ArgumentError, URI::InvalidURIError => e
|
|
29
|
+
result_envelope(status: "unknown", error: e.message, debug: kwargs[:debug])
|
|
30
|
+
end
|
|
31
|
+
end
|
|
32
|
+
|
|
33
|
+
include YouTube
|
|
34
|
+
private_constant :YouTube
|
|
35
|
+
end
|
|
36
|
+
end
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
require "json"
|
|
2
|
+
require "net/http"
|
|
3
|
+
require "uri"
|
|
4
|
+
require_relative "version"
|
|
5
|
+
require_relative "gemini/configuration"
|
|
6
|
+
require_relative "gemini/chat"
|
|
7
|
+
require_relative "gemini/streaming"
|
|
8
|
+
require_relative "gemini/youtube"
|
|
9
|
+
require_relative "gemini/models"
|
|
10
|
+
|
|
11
|
+
class AiLite
|
|
12
|
+
class Gemini
|
|
13
|
+
class << self
|
|
14
|
+
def configuration
|
|
15
|
+
@configuration ||= Configuration.new
|
|
16
|
+
end
|
|
17
|
+
|
|
18
|
+
def configure
|
|
19
|
+
yield configuration
|
|
20
|
+
reset_client!
|
|
21
|
+
configuration
|
|
22
|
+
end
|
|
23
|
+
|
|
24
|
+
def reset_configuration!
|
|
25
|
+
@configuration = Configuration.new
|
|
26
|
+
reset_client!
|
|
27
|
+
configuration
|
|
28
|
+
end
|
|
29
|
+
|
|
30
|
+
def client
|
|
31
|
+
@client ||= new
|
|
32
|
+
end
|
|
33
|
+
|
|
34
|
+
def reset_client!
|
|
35
|
+
@client = nil
|
|
36
|
+
end
|
|
37
|
+
|
|
38
|
+
def youtube(url, prompt:, **kwargs)
|
|
39
|
+
client.youtube(url, prompt: prompt, **kwargs)
|
|
40
|
+
end
|
|
41
|
+
|
|
42
|
+
def chat_stream(message, **kwargs, &block)
|
|
43
|
+
client.chat_stream(message, **kwargs, &block)
|
|
44
|
+
end
|
|
45
|
+
|
|
46
|
+
def chat(message, **kwargs)
|
|
47
|
+
client.chat(message, **kwargs)
|
|
48
|
+
end
|
|
49
|
+
end
|
|
50
|
+
|
|
51
|
+
attr_reader :api_key, :model, :timeout, :max_output_tokens, :headers
|
|
52
|
+
|
|
53
|
+
def initialize(api_key: nil, model: nil, timeout: nil, max_output_tokens: nil)
|
|
54
|
+
config = self.class.configuration
|
|
55
|
+
@api_key = api_key || config.api_key || ENV["GEMINI_API_KEY"]
|
|
56
|
+
raise ArgumentError, "Missing Gemini API key" if @api_key.to_s.strip.empty?
|
|
57
|
+
|
|
58
|
+
@model = model || config.model
|
|
59
|
+
@timeout = timeout || config.timeout
|
|
60
|
+
@max_output_tokens = max_output_tokens || config.max_output_tokens
|
|
61
|
+
@headers = { "x-goog-api-key" => @api_key, "Content-Type" => "application/json" }
|
|
62
|
+
end
|
|
63
|
+
|
|
64
|
+
private
|
|
65
|
+
|
|
66
|
+
def post_interaction(payload)
|
|
67
|
+
uri = URI.parse("#{API_BASE_URL}/interactions")
|
|
68
|
+
request = Net::HTTP::Post.new(uri)
|
|
69
|
+
headers.each { |key, value| request[key] = value }
|
|
70
|
+
request.body = JSON.generate(payload)
|
|
71
|
+
Net::HTTP.start(uri.host, uri.port, use_ssl: uri.scheme == "https") do |http|
|
|
72
|
+
http.open_timeout = timeout
|
|
73
|
+
http.read_timeout = timeout
|
|
74
|
+
http.request(request)
|
|
75
|
+
end
|
|
76
|
+
end
|
|
77
|
+
|
|
78
|
+
def result_envelope(status:, content: nil, error: nil, response_id: nil, raw: nil, debug: false)
|
|
79
|
+
{
|
|
80
|
+
"content" => content, "response_id" => response_id, "status" => status,
|
|
81
|
+
"error" => error, "raw" => debug ? raw : nil
|
|
82
|
+
}
|
|
83
|
+
end
|
|
84
|
+
|
|
85
|
+
def error_message(raw)
|
|
86
|
+
return raw.to_s unless raw.is_a?(Hash)
|
|
87
|
+
|
|
88
|
+
error = raw["error"]
|
|
89
|
+
error.is_a?(Hash) ? (error["message"] || error["status"] || error.to_s) : (error || raw["message"] || raw.to_s)
|
|
90
|
+
end
|
|
91
|
+
end
|
|
92
|
+
end
|
|
@@ -0,0 +1,51 @@
|
|
|
1
|
+
class AiLite
|
|
2
|
+
class OpenAI
|
|
3
|
+
API_BASE_URL = "https://api.openai.com/v1".freeze
|
|
4
|
+
DEFAULT_MODEL = "gpt-5.5".freeze
|
|
5
|
+
DEFAULT_MODERATION_MODEL = "omni-moderation-latest".freeze
|
|
6
|
+
DEFAULT_EMBEDDING_MODEL = "text-embedding-3-small".freeze
|
|
7
|
+
DEFAULT_IMAGE_MODEL = "gpt-image-2".freeze
|
|
8
|
+
DEFAULT_SPEECH_MODEL = "gpt-4o-mini-tts".freeze
|
|
9
|
+
DEFAULT_SPEECH_VOICE = "alloy".freeze
|
|
10
|
+
DEFAULT_SPEECH_FORMAT = "mp3".freeze
|
|
11
|
+
DEFAULT_TRANSCRIPTION_MODEL = "gpt-transcribe".freeze
|
|
12
|
+
DEFAULT_TIMEOUT = 120
|
|
13
|
+
DEFAULT_MAX_OUTPUT_TOKENS = 2000
|
|
14
|
+
IMAGE_MIME_TYPES = {
|
|
15
|
+
".gif" => "image/gif",
|
|
16
|
+
".jpeg" => "image/jpeg",
|
|
17
|
+
".jpg" => "image/jpeg",
|
|
18
|
+
".png" => "image/png",
|
|
19
|
+
".webp" => "image/webp"
|
|
20
|
+
}.freeze
|
|
21
|
+
AUDIO_MIME_TYPES = {
|
|
22
|
+
".flac" => "audio/flac",
|
|
23
|
+
".m4a" => "audio/mp4",
|
|
24
|
+
".mp3" => "audio/mpeg",
|
|
25
|
+
".mp4" => "audio/mp4",
|
|
26
|
+
".mpeg" => "audio/mpeg",
|
|
27
|
+
".mpga" => "audio/mpeg",
|
|
28
|
+
".ogg" => "audio/ogg",
|
|
29
|
+
".wav" => "audio/wav",
|
|
30
|
+
".webm" => "audio/webm"
|
|
31
|
+
}.freeze
|
|
32
|
+
|
|
33
|
+
class Configuration
|
|
34
|
+
attr_accessor :api_key, :model, :moderation_model, :embedding_model, :image_model, :speech_model, :speech_voice, :transcription_model, :timeout, :max_output_tokens
|
|
35
|
+
|
|
36
|
+
def initialize
|
|
37
|
+
@api_key = nil
|
|
38
|
+
@model = DEFAULT_MODEL
|
|
39
|
+
@moderation_model = DEFAULT_MODERATION_MODEL
|
|
40
|
+
@embedding_model = DEFAULT_EMBEDDING_MODEL
|
|
41
|
+
@image_model = DEFAULT_IMAGE_MODEL
|
|
42
|
+
@speech_model = DEFAULT_SPEECH_MODEL
|
|
43
|
+
@speech_voice = DEFAULT_SPEECH_VOICE
|
|
44
|
+
@transcription_model = DEFAULT_TRANSCRIPTION_MODEL
|
|
45
|
+
@timeout = DEFAULT_TIMEOUT
|
|
46
|
+
@max_output_tokens = DEFAULT_MAX_OUTPUT_TOKENS
|
|
47
|
+
end
|
|
48
|
+
end
|
|
49
|
+
|
|
50
|
+
end
|
|
51
|
+
end
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
class AiLite
|
|
2
|
+
class OpenAI
|
|
3
|
+
module Models
|
|
4
|
+
def self.configuration_snapshot(source)
|
|
5
|
+
{
|
|
6
|
+
"chat" => source.model,
|
|
7
|
+
"moderation" => source.moderation_model,
|
|
8
|
+
"embedding" => source.embedding_model,
|
|
9
|
+
"image" => source.image_model,
|
|
10
|
+
"speech" => source.speech_model,
|
|
11
|
+
"transcription" => source.transcription_model
|
|
12
|
+
}.transform_values { |value| value.is_a?(String) ? value.dup : value }
|
|
13
|
+
end
|
|
14
|
+
|
|
15
|
+
def configured_models
|
|
16
|
+
Models.configuration_snapshot(self)
|
|
17
|
+
end
|
|
18
|
+
|
|
19
|
+
def available_models(debug: false)
|
|
20
|
+
response = nil
|
|
21
|
+
raw = nil
|
|
22
|
+
uri = URI.parse("#{API_BASE_URL}/models")
|
|
23
|
+
request = Net::HTTP::Get.new(uri)
|
|
24
|
+
headers.each { |key, value| request[key] = value }
|
|
25
|
+
response = Net::HTTP.start(uri.host, uri.port, use_ssl: uri.scheme == "https") do |http|
|
|
26
|
+
http.open_timeout = timeout
|
|
27
|
+
http.read_timeout = timeout
|
|
28
|
+
http.request(request)
|
|
29
|
+
end
|
|
30
|
+
raw = response.body
|
|
31
|
+
raw = JSON.parse(raw)
|
|
32
|
+
unless success_status?(response.code.to_i)
|
|
33
|
+
return prettify_data(status: response.code.to_i, error: error_message(raw), raw: raw, debug: debug)
|
|
34
|
+
end
|
|
35
|
+
|
|
36
|
+
unless raw.is_a?(Hash) && raw["data"].is_a?(Array) && raw["data"].all? { |item|
|
|
37
|
+
item.is_a?(Hash) && item["id"].is_a?(String) && !item["id"].empty?
|
|
38
|
+
}
|
|
39
|
+
raise "Invalid model list response"
|
|
40
|
+
end
|
|
41
|
+
|
|
42
|
+
prettify_data(status: response.code.to_i, content: raw["data"].map { |item| item["id"] }.uniq,
|
|
43
|
+
raw: raw, debug: debug)
|
|
44
|
+
rescue StandardError => e
|
|
45
|
+
prettify_data(status: response_status(response), error: e.message, raw: raw, debug: debug)
|
|
46
|
+
end
|
|
47
|
+
|
|
48
|
+
# nil means the lookup failed; false means a successful lookup did not find it.
|
|
49
|
+
def model_available?(model_id)
|
|
50
|
+
unless model_id.is_a?(String) && !model_id.strip.empty?
|
|
51
|
+
raise ArgumentError, "model_id must be a nonempty string"
|
|
52
|
+
end
|
|
53
|
+
|
|
54
|
+
result = available_models
|
|
55
|
+
return nil if result["error"]
|
|
56
|
+
|
|
57
|
+
result["content"].include?(model_id)
|
|
58
|
+
end
|
|
59
|
+
|
|
60
|
+
def check_configured_models(debug: false)
|
|
61
|
+
configured = configured_models
|
|
62
|
+
result = available_models(debug: debug)
|
|
63
|
+
report = configured.each_with_object({}) do |(usage, model), entries|
|
|
64
|
+
entries[usage] = {
|
|
65
|
+
"model" => model,
|
|
66
|
+
"usage" => usage,
|
|
67
|
+
"available" => result["error"] ? nil : result["content"].include?(model),
|
|
68
|
+
"error" => result["error"]
|
|
69
|
+
}
|
|
70
|
+
end
|
|
71
|
+
result.merge("content" => report)
|
|
72
|
+
end
|
|
73
|
+
end
|
|
74
|
+
|
|
75
|
+
include Models
|
|
76
|
+
private_constant :Models
|
|
77
|
+
|
|
78
|
+
class << self
|
|
79
|
+
def configured_models
|
|
80
|
+
Models.configuration_snapshot(configuration)
|
|
81
|
+
end
|
|
82
|
+
|
|
83
|
+
def available_models(**kwargs)
|
|
84
|
+
client.available_models(**kwargs)
|
|
85
|
+
end
|
|
86
|
+
|
|
87
|
+
def model_available?(model_id)
|
|
88
|
+
client.model_available?(model_id)
|
|
89
|
+
end
|
|
90
|
+
|
|
91
|
+
def check_configured_models(**kwargs)
|
|
92
|
+
client.check_configured_models(**kwargs)
|
|
93
|
+
end
|
|
94
|
+
end
|
|
95
|
+
end
|
|
96
|
+
end
|
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
class AiLite
|
|
2
|
+
class OpenAI
|
|
3
|
+
module Streaming
|
|
4
|
+
# Text deltas are yielded immediately; the return value is the final envelope.
|
|
5
|
+
def chat_stream(message, model: nil, instructions: nil, previous_response_id: nil, max_output_tokens: nil, debug: false, options: {}, &block)
|
|
6
|
+
raise ArgumentError, "chat_stream requires a block" unless block
|
|
7
|
+
|
|
8
|
+
status = "unknown"
|
|
9
|
+
text = +""
|
|
10
|
+
response_id = nil
|
|
11
|
+
raw = nil
|
|
12
|
+
callback_error = nil
|
|
13
|
+
begin
|
|
14
|
+
payload = options.reject { |key, _| key.to_s == "stream" }.merge(
|
|
15
|
+
model: model || self.model, input: message, stream: true,
|
|
16
|
+
max_output_tokens: max_output_tokens || self.max_output_tokens
|
|
17
|
+
)
|
|
18
|
+
payload[:instructions] = instructions if instructions
|
|
19
|
+
payload[:previous_response_id] = previous_response_id if previous_response_id
|
|
20
|
+
uri = URI.parse(response_endpoint)
|
|
21
|
+
request = Net::HTTP::Post.new(uri)
|
|
22
|
+
headers.each { |key, value| request[key] = value }
|
|
23
|
+
request["Accept"] = "text/event-stream"
|
|
24
|
+
request.body = JSON.generate(payload)
|
|
25
|
+
result = nil
|
|
26
|
+
|
|
27
|
+
Net::HTTP.start(uri.host, uri.port, use_ssl: uri.scheme == "https") do |http|
|
|
28
|
+
http.open_timeout = timeout
|
|
29
|
+
http.read_timeout = timeout
|
|
30
|
+
# Exit the request block at the terminal event, closing the connection.
|
|
31
|
+
catch(:ai_lite_stream_finished) do
|
|
32
|
+
http.request(request) do |response|
|
|
33
|
+
status = response.code.to_i
|
|
34
|
+
unless success_status?(status)
|
|
35
|
+
body = +""
|
|
36
|
+
response.read_body { |chunk| body << chunk }
|
|
37
|
+
raw = body
|
|
38
|
+
raw = JSON.parse(body)
|
|
39
|
+
result = prettify_data(status: status, error: error_message(raw), raw: raw, debug: debug)
|
|
40
|
+
throw :ai_lite_stream_finished
|
|
41
|
+
end
|
|
42
|
+
|
|
43
|
+
each_stream_event(response) do |event|
|
|
44
|
+
raw = event
|
|
45
|
+
response_id = event.dig("response", "id") || response_id
|
|
46
|
+
case event["type"]
|
|
47
|
+
when "response.output_text.delta"
|
|
48
|
+
delta = event.fetch("delta")
|
|
49
|
+
text << delta
|
|
50
|
+
begin
|
|
51
|
+
block.call(delta)
|
|
52
|
+
rescue StandardError => e
|
|
53
|
+
callback_error = e
|
|
54
|
+
raise
|
|
55
|
+
end
|
|
56
|
+
when "response.completed", "response.failed", "response.incomplete"
|
|
57
|
+
raw = event.fetch("response")
|
|
58
|
+
error = case event["type"]
|
|
59
|
+
when "response.failed"
|
|
60
|
+
raw["error"] ? error_message(raw) : "Response failed"
|
|
61
|
+
when "response.incomplete"
|
|
62
|
+
"Response incomplete: #{raw.dig('incomplete_details', 'reason') || 'unknown reason'}"
|
|
63
|
+
end
|
|
64
|
+
content = error ? text : parse_content(extract_output_text(raw))
|
|
65
|
+
result = prettify_data(status: status, content: content, error: error,
|
|
66
|
+
response_id: response_id, raw: raw, debug: debug)
|
|
67
|
+
throw :ai_lite_stream_finished
|
|
68
|
+
when "error"
|
|
69
|
+
result = prettify_data(status: status, content: text, error: error_message(event),
|
|
70
|
+
response_id: response_id, raw: raw, debug: debug)
|
|
71
|
+
throw :ai_lite_stream_finished
|
|
72
|
+
end
|
|
73
|
+
end
|
|
74
|
+
end
|
|
75
|
+
end
|
|
76
|
+
end
|
|
77
|
+
result || prettify_data(status: status, content: text, response_id: response_id,
|
|
78
|
+
error: "Stream ended before a terminal response event", raw: raw, debug: debug)
|
|
79
|
+
rescue StandardError => e
|
|
80
|
+
raise if callback_error.equal?(e)
|
|
81
|
+
|
|
82
|
+
prettify_data(status: status, content: text.empty? ? nil : text, response_id: response_id,
|
|
83
|
+
error: e.message, raw: raw, debug: debug)
|
|
84
|
+
end
|
|
85
|
+
end
|
|
86
|
+
|
|
87
|
+
private
|
|
88
|
+
|
|
89
|
+
def each_stream_event(response)
|
|
90
|
+
buffer = +"".b
|
|
91
|
+
data = []
|
|
92
|
+
consume = lambda do
|
|
93
|
+
while (boundary = buffer.index(/[\r\n]/))
|
|
94
|
+
# A trailing CR might be the first byte of a CRLF split across reads.
|
|
95
|
+
break if buffer.getbyte(boundary) == 13 && boundary == buffer.bytesize - 1
|
|
96
|
+
|
|
97
|
+
width = buffer.byteslice(boundary, 2) == "\r\n" ? 2 : 1
|
|
98
|
+
line = buffer.slice!(0, boundary + width).byteslice(0, boundary)
|
|
99
|
+
if line.empty?
|
|
100
|
+
unless data.empty?
|
|
101
|
+
payload = data.join("\n").force_encoding(Encoding::UTF_8)
|
|
102
|
+
data.clear
|
|
103
|
+
unless payload == "[DONE]"
|
|
104
|
+
event = JSON.parse(payload)
|
|
105
|
+
raise "Invalid streaming event: expected an object" unless event.is_a?(Hash)
|
|
106
|
+
yield event
|
|
107
|
+
end
|
|
108
|
+
end
|
|
109
|
+
elsif line.start_with?("data:")
|
|
110
|
+
data << line.sub(/\Adata: ?/, "")
|
|
111
|
+
end
|
|
112
|
+
end
|
|
113
|
+
end
|
|
114
|
+
response.read_body do |chunk|
|
|
115
|
+
buffer << chunk.b
|
|
116
|
+
consume.call
|
|
117
|
+
end
|
|
118
|
+
# A final lone CR is a complete line delimiter at EOF.
|
|
119
|
+
if buffer.end_with?("\r")
|
|
120
|
+
buffer << "\n"
|
|
121
|
+
consume.call
|
|
122
|
+
end
|
|
123
|
+
end
|
|
124
|
+
|
|
125
|
+
end
|
|
126
|
+
|
|
127
|
+
include Streaming
|
|
128
|
+
private_constant :Streaming
|
|
129
|
+
end
|
|
130
|
+
end
|