0
0
Fork 0
mirror of https://github.com/discourse/discourse.git synced 2026-08-06 13:08:40 +08:00
discourse/plugins/discourse-ai/lib/ai_helper/assistant.rb
Rafael dos Santos Silva 53e25da301
PERF: Speed up AI helper diff streaming (#41787)
Previously, streamed diff-type helper results (e.g. proofread) felt much
slower than other AI streams: the server cooked the entire input through
the markdown pipeline on every streamed chunk to build a `diff` payload
no client consumed, and the client animation was hard-capped at one
character per 10ms tick while recomputing the full word diff on every
word boundary.

This change drops the per-chunk `parse_diff` call (the diff is still
computed for the final update), makes the typing animation batch
characters adaptively based on how far it lags behind the received text,
and throttles diff recomputes to every 100ms.



https://github.com/user-attachments/assets/5914cf24-1b52-488e-bca0-6992a780f777
2026-07-17 10:17:54 -03:00

507 lines
15 KiB
Ruby
Vendored

# frozen_string_literal: true
module DiscourseAi
module AiHelper
class Assistant
IMAGE_CAPTION_MAX_WORDS = 50
TRANSLATE = "translate"
GENERATE_TITLES = "generate_titles"
PROOFREAD = "proofread"
MARKDOWN_TABLE = "markdown_table"
CUSTOM_PROMPT = "custom_prompt"
EXPLAIN = "explain"
ILLUSTRATE_POST = "illustrate_post"
REPLACE_DATES = "replace_dates"
IMAGE_CAPTION = "image_caption"
def self.prompt_cache
@prompt_cache ||= DiscourseAi::MultisiteHash.new("prompt_cache")
end
def self.clear_prompt_cache!
prompt_cache.flush!
end
def self.prompt_agent_ids
agents_prompt_map.keys.compact.uniq
end
def initialize(helper_llm: nil, image_caption_llm: nil)
@helper_llm = helper_llm
@image_caption_llm = image_caption_llm
end
def available_prompts(user)
key = "prompt_cache_#{I18n.locale}"
prompts = self.class.prompt_cache.fetch(key) { all_prompts }
prompts
.map do |prompt|
next if !user.in_any_groups?(prompt[:allowed_group_ids])
if prompt[:name] == ILLUSTRATE_POST
agent = AiAgent.find_by(id: SiteSetting.ai_helper_post_illustrator_agent)
next if agent.blank? || !agent_has_image_generation_tool?(agent)
end
# We cannot cache this. It depends on the user's effective_locale.
if prompt[:name] == TRANSLATE
locale = user.effective_locale
locale_hash =
LocaleSiteSetting.language_names[locale] ||
LocaleSiteSetting.language_names[locale.split("_")[0]]
translation =
I18n.t(
"discourse_ai.ai_helper.prompts.translate",
language: locale_hash["nativeName"],
) || prompt[:name]
prompt.merge(translated_name: translation)
else
prompt
end
end
.compact
end
def custom_locale_instructions(user = nil, force_default_locale)
locale = SiteSetting.default_locale
locale = user.effective_locale if !force_default_locale && user
locale_instructions(locale)
end
def locale_instructions(locale)
locale_hash = LocaleSiteSetting.language_names[locale]
if locale != "en" && locale_hash
locale_description = "#{locale_hash["name"]} (#{locale_hash["nativeName"]})"
"It is imperative that you write your answer in #{locale_description}, you are interacting with a #{locale_description} speaking user. Leave tag names in English."
else
nil
end
end
def attach_user_context(context, user = nil, force_default_locale: false)
locale = SiteSetting.default_locale
locale = user.effective_locale if user && !force_default_locale
locale_hash = LocaleSiteSetting.language_names[locale]
context.user_language = "#{locale_hash["name"]}"
if user
timezone = user&.user_option&.timezone || "UTC"
current_time = Time.now.in_time_zone(timezone)
temporal_context = {
utc_date_time: current_time.iso8601,
local_time: current_time.strftime("%H:%M"),
user: {
timezone: timezone,
weekday: current_time.strftime("%A"),
},
}
context.temporal_context = temporal_context.to_json
end
context
end
def generate_prompt(
helper_mode,
input,
user,
force_default_locale: false,
custom_prompt: nil,
&block
)
bot = build_bot(helper_mode, user)
user_input = "<input>#{input}</input>"
if helper_mode == CUSTOM_PROMPT && custom_prompt.present?
user_input = "<input>#{custom_prompt}:\n#{input}</input>"
end
context =
DiscourseAi::Agents::BotContext.new(
user: user,
skip_show_thinking: true,
feature_name: "ai_helper",
messages: [{ type: :user, content: user_input }],
format_dates: helper_mode == REPLACE_DATES,
custom_instructions: custom_locale_instructions(user, force_default_locale),
)
context = attach_user_context(context, user, force_default_locale: force_default_locale)
bad_json = false
json_summary_schema_key = bot.agent.response_format&.first.to_h
schema_key = json_summary_schema_key["key"]&.to_sym
schema_type = json_summary_schema_key["type"]
if schema_type == "array"
helper_response = []
else
helper_response = +""
end
buffer_blk =
Proc.new do |partial, _, type|
if type == :structured_output && schema_type
helper_chunk = partial.read_buffered_property(schema_key)
next if helper_chunk.nil? || helper_chunk.empty?
if schema_type == "array"
if helper_chunk.is_a?(Array)
helper_chunk.each do |item|
helper_response << item if helper_response.exclude?(item)
end
end
elsif schema_type == "string"
helper_response << helper_chunk
else
helper_response = helper_chunk
end
block.call(helper_chunk) if block && !bad_json
elsif type.blank?
# Assume response is a regular completion.
helper_response << partial
block.call(partial) if block
end
end
bot.reply(context, &buffer_blk)
helper_response
end
def generate_and_send_prompt(
helper_mode,
input,
user,
force_default_locale: false,
custom_prompt: nil
)
helper_response =
generate_prompt(
helper_mode,
input,
user,
force_default_locale: force_default_locale,
custom_prompt: custom_prompt,
)
result = { type: prompt_type(helper_mode) }
result[:suggestions] = (
if result[:type] == :list
helper_response.flatten.map { |suggestion| sanitize_result(suggestion) }
else
sanitized = sanitize_result(helper_response)
result[:diff] = parse_diff(input, sanitized) if result[:type] == :diff
[sanitized]
end
)
result
end
def stream_prompt(
helper_mode,
input,
user,
channel,
force_default_locale: false,
client_id: nil,
custom_prompt: nil
)
streamed_result = +""
start = Time.now
type = prompt_type(helper_mode)
generate_prompt(
helper_mode,
input,
user,
force_default_locale: force_default_locale,
custom_prompt: custom_prompt,
) do |partial_response|
streamed_result << partial_response
# Throttle updates
if (streamed_result.length > 10 && (Time.now - start > 0.3)) || Rails.env.test?
sanitized = sanitize_result(streamed_result)
payload = { result: sanitized, done: false }
publish_update(channel, payload, user, client_id: client_id)
start = Time.now
end
end
final_diff = parse_diff(input, streamed_result) if type == :diff
sanitized_result = sanitize_result(streamed_result)
if sanitized_result.present?
publish_update(
channel,
{ result: sanitized_result, diff: final_diff, done: true },
user,
client_id: client_id,
)
end
end
def generate_image_caption(upload, user, locale: nil, post: nil, skip_access_check: false)
bot = build_bot(IMAGE_CAPTION, user, skip_access_check: skip_access_check)
force_default_locale = false
custom_instructions =
if locale.present?
locale_instructions(locale)
else
custom_locale_instructions(user, force_default_locale)
end
context =
DiscourseAi::Agents::BotContext.new(
post: post,
user: user,
skip_show_thinking: true,
feature_name: IMAGE_CAPTION,
messages: [
{
type: :user,
content: ["Describe this image in a single sentence.", { upload_id: upload.id }],
},
],
custom_instructions: custom_instructions,
)
structured_output = nil
buffer_blk =
Proc.new do |partial, _, type|
if type == :structured_output
structured_output = partial
bot.agent.response_format&.first.to_h
end
end
bot.reply(context, llm_args: { max_tokens: 1024 }, &buffer_blk)
raw_caption = ""
if structured_output
json_summary_schema_key = bot.agent.response_format&.first.to_h
raw_caption =
structured_output.read_buffered_property(json_summary_schema_key["key"]&.to_sym)
end
raw_caption.delete("|").squish.truncate_words(IMAGE_CAPTION_MAX_WORDS)
end
def ensure_mode_access!(helper_mode, user)
ai_agent = ai_agent_for_mode(helper_mode)
return if ai_agent.nil?
raise Discourse::InvalidAccess if !user.in_any_groups?(ai_agent.allowed_group_ids.to_a)
ai_agent
end
private
def agent_has_image_generation_tool?(agent)
agent&.has_image_generation_tool?
rescue StandardError => e
Rails.logger.warn(
"Failed to check image generation tool for agent #{agent&.id}: #{e.message}",
)
false
end
def ai_agent_for_mode(helper_mode)
agent_id = agents_prompt_map(include_image_caption: true).invert[helper_mode]
raise Discourse::InvalidParameters.new(:mode) if agent_id.blank?
AiAgent.find_by(id: agent_id)
end
def build_bot(helper_mode, user, skip_access_check: false)
ai_agent =
if skip_access_check
ai_agent_for_mode(helper_mode)
else
ensure_mode_access!(helper_mode, user)
end
return if ai_agent.nil?
agent_klass = ai_agent.class_instance
return if agent_klass.nil?
llm_model = find_ai_helper_model(helper_mode, agent_klass)
DiscourseAi::Agents::Bot.as(user, agent: agent_klass.new, model: llm_model)
end
def find_ai_helper_model(helper_mode, agent_klass)
if helper_mode == IMAGE_CAPTION && @image_caption_llm.is_a?(LlmModel)
return @image_caption_llm
end
return @helper_llm if helper_mode != IMAGE_CAPTION && @helper_llm.is_a?(LlmModel)
self.class.find_ai_helper_model(helper_mode, agent_klass)
end
# Priorities are:
# 1. Agent's default LLM
# 2. SiteSetting.ai_default_llm_model (or newest LLM if not set)
def self.find_ai_helper_model(helper_mode, agent_klass)
model_id = agent_klass.default_llm_id || SiteSetting.ai_default_llm_model
if model_id.present?
LlmModel.find_by(id: model_id)
else
LlmModel.last
end
end
def self.agents_prompt_map(include_image_caption: false)
map = {
SiteSetting.ai_helper_translator_agent.to_i => TRANSLATE,
SiteSetting.ai_helper_title_suggestions_agent.to_i => GENERATE_TITLES,
SiteSetting.ai_helper_proofreader_agent.to_i => PROOFREAD,
SiteSetting.ai_helper_markdown_tables_agent.to_i => MARKDOWN_TABLE,
SiteSetting.ai_helper_custom_prompt_agent.to_i => CUSTOM_PROMPT,
SiteSetting.ai_helper_explain_agent.to_i => EXPLAIN,
SiteSetting.ai_helper_post_illustrator_agent.to_i => ILLUSTRATE_POST,
SiteSetting.ai_helper_smart_dates_agent.to_i => REPLACE_DATES,
}
if include_image_caption
image_caption_agent = SiteSetting.ai_image_caption_agent.to_i
map[image_caption_agent] = IMAGE_CAPTION if image_caption_agent
end
map
end
def agents_prompt_map(include_image_caption: false)
self.class.agents_prompt_map(include_image_caption:)
end
def all_prompts
AiAgent
.where(id: agents_prompt_map.keys)
.map do |ai_agent|
prompt_name = agents_prompt_map[ai_agent.id]
if prompt_name
{
name: prompt_name,
translated_name:
I18n.t("discourse_ai.ai_helper.prompts.#{prompt_name}", default: nil) ||
prompt_name,
prompt_type: prompt_type(prompt_name),
icon: icon_map(prompt_name),
location: location_map(prompt_name),
allowed_group_ids: ai_agent.allowed_group_ids,
}
end
end
.compact
end
SANITIZE_REGEX_STR =
%w[term context topic replyTo input output result]
.map { |tag| "<#{tag}>\\n?|\\n?</#{tag}>" }
.join("|")
SANITIZE_REGEX = Regexp.new(SANITIZE_REGEX_STR, Regexp::IGNORECASE | Regexp::MULTILINE)
def sanitize_result(result)
result.gsub(SANITIZE_REGEX, "")
end
def publish_update(channel, payload, user, client_id: nil)
# when publishing we make sure we do not keep large backlogs on the channel
# and make sure we clear the streaming info after 60 seconds
# this ensures we do not bloat redis
if client_id
MessageBus.publish(
channel,
payload,
user_ids: [user.id],
client_ids: [client_id],
max_backlog_age: 60,
)
else
MessageBus.publish(channel, payload, user_ids: [user.id], max_backlog_age: 60)
end
end
def icon_map(name)
case name
when TRANSLATE
"language"
when GENERATE_TITLES
"heading"
when PROOFREAD
"spell-check"
when MARKDOWN_TABLE
"table"
when CUSTOM_PROMPT
"comment"
when EXPLAIN
"question"
when ILLUSTRATE_POST
"images"
when REPLACE_DATES
"calendar-days"
else
nil
end
end
def location_map(name)
case name
when TRANSLATE
%w[composer post]
when GENERATE_TITLES
%w[composer]
when PROOFREAD
%w[composer post]
when MARKDOWN_TABLE
%w[composer]
when CUSTOM_PROMPT
%w[composer post]
when EXPLAIN
%w[post]
when ILLUSTRATE_POST
%w[composer]
when REPLACE_DATES
%w[composer]
else
%w[]
end
end
def prompt_type(prompt_name)
if [PROOFREAD, MARKDOWN_TABLE, REPLACE_DATES, CUSTOM_PROMPT].include?(prompt_name)
return :diff
end
return :list if [ILLUSTRATE_POST, GENERATE_TITLES].include?(prompt_name)
:text
end
def parse_diff(text, suggestion)
cooked_text = PrettyText.cook(text)
cooked_suggestion = PrettyText.cook(suggestion)
DiscourseDiff.new(cooked_text, cooked_suggestion).inline_html
end
end
end
end