import os
import oci
import logging
import time
from aidputils.agents.toolkit.service_metrics_util import get_service_metrics
from aidputils.agents.auth.client.generative_ai_inference_v2_client import GenerativeAiInferenceV2Client
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
_genai_client_cache = {}
[docs]
class GenAIService:
def __init__(self, config, signer, service_endpoint):
self.generative_ai_inference_client = self.get_genai_client(
endpoint=service_endpoint,
signer=signer
)
self.genai_compartment = (
config.get("llm", {}).get("compartment_id")
or os.getenv("USER_COMPARTMENT_ID")
)
_init_genai_service_metrics(self)
[docs]
def generate_embedding_text(self, embedding_model_name, list_text):
start_time = time.time()
try:
embed_text_response = self.generative_ai_inference_client.embed_text(
embed_text_details=oci.generative_ai_inference.models.EmbedTextDetails(
inputs=list_text,
serving_mode=oci.generative_ai_inference.models.OnDemandServingMode(
serving_type="ON_DEMAND", model_id=embedding_model_name, compartment_id=self.genai_compartment
),
)
)
embeddings = [[float(x) for x in inner] for inner in embed_text_response.data.embedTextResult.embeddings]
elapsed_ms = (time.time() - start_time) * 1000.0
self._record_genai_embed_success(elapsed_ms)
return embeddings
except Exception:
elapsed_ms = (time.time() - start_time) * 1000.0
self._record_genai_embed_failure(elapsed_ms)
raise
[docs]
def generate_response(self, conf, runtime_params, result_list):
try:
llm_conf = conf.get("llm")
model_provider = llm_conf.get("model_provider")
user_prompt = llm_conf.get("prompt", GenAIService.get_llm_default_prompt())
llm_context = "\n\n".join([
f"Chunk {idx + 1}: {r.body}" for idx, r in enumerate(result_list)
])
final_prompt = (
f"{user_prompt}\n\n"
f"Context:\n{llm_context}\n"
f"Question: {runtime_params.get('query', '')}"
)
default_llm_model = GenAIService.get_llm_default_model_name()
model_id = llm_conf.get("model_id", default_llm_model)
model_args = llm_conf.get("model_args") or {}
chat_detail = oci.generative_ai_inference.models.ChatDetails()
final_message = GenAIService.build_message(model_provider, final_prompt)
if GenAIService._is_model_provider_not_cohere(model_provider):
logger.info("Generating a generic chat request")
chat_request = oci.generative_ai_inference.models.GenericChatRequest()
chat_request.messages = [final_message]
else:
logger.info("Generating a cohere chat request")
chat_request = oci.generative_ai_inference.models.CohereChatRequest()
chat_request.message = final_message
for attr in ["frequency_penalty", "presence_penalty"]:
GenAIService._set_attr_if_present(chat_request, attr, model_args)
for attr in [
"max_tokens",
"temperature",
"top_p",
"top_k",
]:
GenAIService._set_attr_if_present(chat_request, attr, model_args)
chat_detail.serving_mode = oci.generative_ai_inference.models.OnDemandServingMode(model_id=model_id)
chat_detail.chat_request = chat_request
chat_detail.compartment_id = self.genai_compartment
chat_response = self.generative_ai_inference_client.chat(chat_detail)
retrieved_chunks = [{
"document_id": f"{r.url}#{r.chunkid}",
"content": r.body,
"score": r.score
} for r in result_list]
answer = self.extract_chat_response(chat_response, model_provider)
return {"answer": answer, "retrieved_chunks": retrieved_chunks}
except Exception as e:
logger.error("Error in generate_response: %s", e, exc_info=True)
raise # Re-raise after logging
[docs]
@staticmethod
def build_message(model_provider, msg_content):
if GenAIService._is_model_provider_not_cohere(model_provider):
content = oci.generative_ai_inference.models.TextContent()
content.text = msg_content
message = oci.generative_ai_inference.models.Message()
message.role = "USER"
message.content = [content]
return message
else:
return msg_content
# For cohere-like models (object with .text)
[docs]
@staticmethod
def get_llm_default_prompt():
return (
"You are a helpful and intelligent assistant.\n\n"
"Using only the information provided in the context below, answer the question in a clear, "
"well-structured, and meaningful way.\n"
"You must:\n"
"- Use only the provided context to derive your answer.\n"
"- Do not add any information that is not present in the context.\n"
'- If the context is unclear or does not contain enough information, respond with:\n'
' \"The provided context does not have sufficient information to answer this question.\"\n'
"- Rephrase and summarize the relevant points as needed to improve clarity and flow.\n"
)
[docs]
@staticmethod
def get_llm_default_model_name():
return "cohere.command-a-03-2025"
[docs]
@classmethod
def get_genai_client(cls, endpoint, signer):
"""
Return a cached GenerativeAiInferenceClient for the given endpoint and signer, or construct one if needed.
"""
key = (endpoint, id(signer))
global _genai_client_cache
if key in _genai_client_cache:
return _genai_client_cache[key]
client = GenerativeAiInferenceV2Client(
endpoint=endpoint,
signer=signer
)
_genai_client_cache[key] = client
return client
@staticmethod
def _is_model_provider_not_cohere(model_provider):
return model_provider is None or model_provider.lower() != "cohere"
@staticmethod
def _set_attr_if_present(obj, key, dict_obj):
if key in dict_obj:
setattr(obj, key, dict_obj[key])
# --- Metrics helpers for GenAI embedding ---
def _record_genai_embed_success(self, elapsed_ms: float):
try:
if getattr(self, "_service_metrics", None) and self._genai_embed_success_counter:
self._service_metrics.increment_counter(self._genai_embed_success_counter, 1)
if getattr(self, "_service_metrics", None) and self._genai_embed_latency_hist:
self._service_metrics.record_histogram(self._genai_embed_latency_hist, elapsed_ms)
except Exception:
logger.exception("Failed to record ragtool.embedding.genai success/latency")
def _record_genai_embed_failure(self, elapsed_ms: float):
try:
if getattr(self, "_service_metrics", None) and self._genai_embed_failure_counter:
self._service_metrics.increment_counter(self._genai_embed_failure_counter, 1)
if getattr(self, "_service_metrics", None) and self._genai_embed_latency_hist:
self._service_metrics.record_histogram(self._genai_embed_latency_hist, elapsed_ms)
except Exception:
logger.exception("Failed to record ragtool.embedding.genai failure/latency")
def _init_genai_service_metrics(self):
try:
self._service_metrics = get_service_metrics()
self._genai_embed_success_counter = self._service_metrics.create_counter(
"ragtool.embedding.genai.success",
description="Count of successful embedding generations via GenAI"
)
self._genai_embed_failure_counter = self._service_metrics.create_counter(
"ragtool.embedding.genai.failure",
description="Count of failed embedding generations via GenAI"
)
self._genai_embed_latency_hist = self._service_metrics.create_histogram(
"ragtool.embedding.genai.latency_ms",
description="Latency (ms) of embedding generations via GenAI"
)
except Exception:
logger.exception("Failed to initialize GenAI embedding metrics")
self._service_metrics = None
self._genai_embed_success_counter = None
self._genai_embed_failure_counter = None
self._genai_embed_latency_hist = None