# frozen_string_literal: true

require "langchain/llm/google_vertex_ai"

module Langchain
  module LLM
    class SprintVertexAI < GoogleVertexAI
      DEFAULT_SCOPES = [
        "https://www.googleapis.com/auth/cloud-platform",
        "https://www.googleapis.com/auth/generative-language.retriever"
      ].freeze

      def initialize(project_id:, region:, default_options: {})
        depends_on "googleauth"

        @authorizer = ::Google::Auth.get_application_default(DEFAULT_SCOPES)
        proj_id = project_id || @authorizer.project_id || @authorizer.quota_project_id
        region_prefix = region == "global" ? "" : "#{region}-"
        @url = "https://#{region_prefix}aiplatform.googleapis.com/v1/projects/#{proj_id}/locations/#{region}/publishers/google/models/"
        @defaults = DEFAULTS.merge(default_options)

        chat_parameters.update(
          model: { default: @defaults[:chat_model] },
          temperature: { default: @defaults[:temperature] },
          safety_settings: { default: @defaults[:safety_settings] }
        )
        chat_parameters.remap(
          messages: :contents,
          system: :system_instruction,
          tool_choice: :tool_config
        )
      rescue Signet::AuthorizationError => e
        raise Langchain::LLM::ApiError, "Invalid Google Cloud credentials for Vertex AI: #{e.message}"
      end

      def default_dimensions
        @defaults[:dimensions]
      end

      def embed(text:, model: @defaults[:embedding_model], dimensions: @defaults[:dimensions])
        params = {
          instances: [{ content: text }],
          parameters: { outputDimensionality: dimensions }
        }

        parsed_response = http_post(URI("#{url}#{model}:predict"), params)
        Langchain::LLM::GoogleGeminiResponse.new(parsed_response, model: model)
      end

      def embed_texts(texts:, model: @defaults[:embedding_model], dimensions: @defaults[:dimensions])
        Array(texts).map do |text|
          embed(text:, model:, dimensions:).embedding
        end
      end

      private

      def http_post(url, params)
        http = Net::HTTP.new(url.hostname, url.port)
        http.use_ssl = url.scheme == "https"
        http.open_timeout = 15
        http.read_timeout = 600
        http.write_timeout = 600 if http.respond_to?(:write_timeout=)
        http.set_debug_output(Langchain.logger) if Langchain.logger.debug?

        request = Net::HTTP::Post.new(url)
        request.content_type = "application/json"
        request["Authorization"] = "Bearer #{@authorizer.fetch_access_token!["access_token"]}"
        request.body = params.to_json

        response = http.request(request)
        JSON.parse(response.body)
      end
    end
  end
end
