File size: 4,691 Bytes
338036b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
from ..LLMInterface import LLMInterface
from ..LLMEnums import OpenAIEnums
from openai import OpenAI, APIConnectionError
import logging
from typing import Union, List

class OpenAIProvider(LLMInterface):

     def __init__(self, api_key: str, api_url: str=None,

                  default_input_max_characters: int=1000,

                  default_generation_max_output_tokens: int=1000,

                  default_generation_tempreature: float=0.1):
          self.api_key = api_key
          self.api_url = api_url
          self.default_input_max_characters = default_input_max_characters
          self.default_generation_max_output_tokens = default_generation_max_output_tokens
          self.default_generation_tempreature = default_generation_tempreature

          self.generation_moddel_id = None

          self.embedding_model_id = None
          self.embedding_size = None

          self.client = OpenAI(
               api_key = self.api_key if self.api_key else "placeholder",
               base_url = self.api_url if self.api_url and len(self.api_url) else None

          )

          self.enums = OpenAIEnums
          self.logger = logging.getLogger(__name__)

     def set_generation_model(self, model_id: str):
          self.generation_moddel_id = model_id

     def set_embedding_model(self, model_id: str, embedding_size: int):  
          self.embedding_model_id = model_id   
          self.embedding_size = embedding_size


     def process_text(self, text: str):
          return text[:self.default_input_max_characters].strip()
     
     def generate_text(self, prompt: str,  chat_history: list=[],

                    max_output_token: int=None,temperature: float = None):
           
          if not self.client :
               self.logger.error("client was not set!")
               return None
          
          if not self.generation_moddel_id : 
               self.logger.error("embedding model was not set")
               return None
          
          max_output_token = max_output_token if max_output_token else self.default_generation_max_output_tokens
          temperature = temperature if temperature else self.default_generation_tempreature

          chat_history.append(
          self.construct_prompt(prompt=prompt, role=OpenAIEnums.USER))

          try:
               response = self.client.chat.completions.create(
                    model = self.generation_moddel_id,
                    messages = chat_history,
                    max_tokens = max_output_token,
                    temperature = temperature
               )
          except APIConnectionError as e:
               self.logger.error(
                    f"LLM connection error: could not reach '{self.api_url}'. "
                    f"Make sure Ollama (or your LLM server) is running. Details: {e}"
               )
               return None

          if not response or not response.choices or len(response.choices) == 0 or not response.choices[0]:
               self.logger.error("Error while generating text with OpenAI")
               return None

          return response.choices[0].message.content

     def embed_text(self, text: Union[str, List[str]], document_type: str =None):
          if not self.client :
               self.logger.error("OpenAI client was not set!")
               return None

          if isinstance(text, str):
               text = [text]
               
          if not self.embedding_model_id: 
               self.logger.error("embedding model for OpenAI was not set")
               return None

          try:
               response = self.client.embeddings.create(
                    model = self.embedding_model_id,
                    input=text
               )
          except APIConnectionError as e:
               self.logger.error(
                    f"LLM connection error: could not reach '{self.api_url}'. "
                    f"Make sure Ollama (or your LLM server) is running. Details: {e}"
               )
               return None

          # response validation
          if not response or not response.data or len(response.data) == 0 or not response.data[0].embedding:
               self.logger.error("Error while embedding text with OpenAI")
               return None

          # Return all vectors so the controller can embed many chunks in one request.
          return [item.embedding for item in response.data]

     def construct_prompt(self, prompt: str, role: str):
          return {
               "role": role,
               "content": prompt
          }