predictionguard 2.10.0__py3-none-any.whl

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.
@@ -0,0 +1,6 @@
1
+ """Create controlled and compliant AI systems with PredictionGuard."""
2
+
3
+ from .client import PredictionGuard as PredictionGuard
4
+ from .version import __version__
5
+
6
+ __version__ = __version__
@@ -0,0 +1,151 @@
1
+ import os
2
+
3
+ import requests
4
+ from typing import Optional, Union
5
+
6
+ from .src.audio import Audio
7
+ from .src.chat import Chat
8
+ from .src.completions import Completions
9
+ from .src.detokenize import Detokenize
10
+ from .src.documents import Documents
11
+ from .src.embeddings import Embeddings
12
+ from .src.rerank import Rerank
13
+ from .src.tokenize import Tokenize
14
+ from .src.translate import Translate
15
+ from .src.factuality import Factuality
16
+ from .src.toxicity import Toxicity
17
+ from .src.pii import Pii
18
+ from .src.injection import Injection
19
+ from .src.models import Models
20
+ from .version import __version__
21
+
22
+ __all__ = [
23
+ "PredictionGuard", "Chat", "Completions", "Embeddings",
24
+ "Audio", "Documents", "Rerank", "Tokenize", "Translate",
25
+ "Detokenize", "Factuality", "Toxicity", "Pii", "Injection",
26
+ "Models"
27
+ ]
28
+
29
+ class PredictionGuard:
30
+ """PredictionGuard provides access the Prediction Guard API."""
31
+
32
+ def __init__(
33
+ self,
34
+ api_key: Optional[str] = None,
35
+ url: Optional[str] = None,
36
+ timeout: Optional[Union[int, float]] = None
37
+ ) -> None:
38
+ """
39
+ :param api_key: api_key represents PG api key.
40
+ :param url: url represents the transport and domain:port
41
+ :param timeout: request timeout in seconds.
42
+ """
43
+
44
+ # Get the access api_key.
45
+ if not api_key:
46
+ api_key = os.environ.get("PREDICTIONGUARD_API_KEY")
47
+
48
+ if not api_key:
49
+ raise ValueError(
50
+ "No api_key provided or in environment. "
51
+ "Please provide the api_key as "
52
+ "client = PredictionGuard(api_key=<your_api_key>) "
53
+ "or as PREDICTIONGUARD_API_KEY in your environment."
54
+ )
55
+ self.api_key = api_key
56
+
57
+ if not url:
58
+ url = os.environ.get("PREDICTIONGUARD_URL")
59
+ if not url:
60
+ url = "https://api.predictionguard.com"
61
+ self.url = url
62
+
63
+ if not timeout:
64
+ timeout = os.environ.get("TIMEOUT")
65
+ if not timeout:
66
+ timeout = None
67
+ if timeout:
68
+ try:
69
+ timeout = float(timeout)
70
+ except ValueError:
71
+ raise ValueError(
72
+ "Timeout must be of type integer or float, not %s." % (type(timeout).__name__,)
73
+ )
74
+ except TypeError:
75
+ raise TypeError(
76
+ "Timeout should be of type integer or float, not %s." % (type(timeout).__name__,)
77
+ )
78
+ self.timeout = timeout
79
+
80
+ # Connect to Prediction Guard and set the access api_key.
81
+ self._connect_client()
82
+
83
+ # Pass Prediction Guard class variables to inner classes
84
+ self.chat: Chat = Chat(self.api_key, self.url, self.timeout)
85
+ """Chat generates chat completions based on a conversation history"""
86
+
87
+ self.completions: Completions = Completions(self.api_key, self.url, self.timeout)
88
+ """Completions generates text completions based on the provided input"""
89
+
90
+ self.embeddings: Embeddings = Embeddings(self.api_key, self.url, self.timeout)
91
+ """Embedding generates chat completions based on a conversation history."""
92
+
93
+ self.audio: Audio = Audio(self.api_key, self.url, self.timeout)
94
+ """Audio allows for the transcription of audio files."""
95
+
96
+ self.documents: Documents = Documents(self.api_key, self.url, self.timeout)
97
+ """Documents allows you to extract text from various document file types."""
98
+
99
+ self.rerank: Rerank = Rerank(self.api_key, self.url, self.timeout)
100
+ """Rerank sorts text inputs by semantic relevance to a specified query."""
101
+
102
+ self.translate: Translate = Translate(self.api_key, self.url, self.timeout)
103
+ """Translate converts text from one language to another."""
104
+
105
+ self.factuality: Factuality = Factuality(self.api_key, self.url, self.timeout)
106
+ """Factuality checks the factuality of a given text compared to a reference."""
107
+
108
+ self.toxicity: Toxicity = Toxicity(self.api_key, self.url, self.timeout)
109
+ """Toxicity checks the toxicity of a given text."""
110
+
111
+ self.pii: Pii = Pii(self.api_key, self.url, self.timeout)
112
+ """Pii replaces personal information such as names, SSNs, and emails in a given text."""
113
+
114
+ self.injection: Injection = Injection(self.api_key, self.url, self.timeout)
115
+ """Injection detects potential prompt injection attacks in a given prompt."""
116
+
117
+ self.tokenize: Tokenize = Tokenize(self.api_key, self.url, self.timeout)
118
+ """Tokenize generates tokens for input text."""
119
+
120
+ self.detokenize: Detokenize = Detokenize(self.api_key, self.url, self.timeout)
121
+ """Detokenizes generates text for input tokens."""
122
+
123
+ self.models: Models = Models(self.api_key, self.url, self.timeout)
124
+ """Models lists all of the models available in the Prediction Guard API."""
125
+
126
+ def _connect_client(self) -> None:
127
+
128
+ # Prepare the proper headers.
129
+ headers = {
130
+ "Content-Type": "application/json",
131
+ "Authorization": "Bearer " + self.api_key,
132
+ "User-Agent": "Prediction Guard Python Client: " + __version__,
133
+ }
134
+
135
+ # Try listing models to make sure we can connect.
136
+ response = requests.request("GET", self.url + "/completions", headers=headers, timeout=self.timeout)
137
+
138
+ # If the connection was unsuccessful, raise an exception.
139
+ if response.status_code == 200:
140
+ pass
141
+ elif response.status_code == 401:
142
+ raise ValueError(
143
+ "Could not connect to Prediction Guard API with the given api_key. "
144
+ "Please check your access api_key and try again."
145
+ )
146
+ elif response.status_code == 404:
147
+ raise ValueError(
148
+ "Could not connect to Prediction Guard API with given url. "
149
+ "Please check url specified, if no url specified, "
150
+ "Please contact support."
151
+ )
@@ -0,0 +1,194 @@
1
+ from typing import Any, Dict, List, Optional
2
+
3
+ import requests
4
+
5
+ from ..version import __version__
6
+
7
+
8
+ class Audio:
9
+ """
10
+ Audio generates a response based on audio data.
11
+
12
+ Usage::
13
+
14
+ import os
15
+ import json
16
+
17
+ from predictionguard import PredictionGuard
18
+
19
+ # Set your Prediction Guard token and url as an environmental variable.
20
+ os.environ["PREDICTIONGUARD_API_KEY"] = "<api key>"
21
+ os.environ["PREDICTIONGUARD_URL"] = "<url>"
22
+
23
+ # Or set your Prediction Guard token and url when initializing the PredictionGuard class.
24
+ client = PredictionGuard(
25
+ api_key="<api_key>",
26
+ url="<url>"
27
+ )
28
+
29
+ result = client.audio.transcriptions.create(
30
+ model="base",
31
+ file="sample_audio.wav"
32
+ )
33
+
34
+ print(json.dumps(
35
+ response,
36
+ sort_keys=True,
37
+ indent=4,
38
+ separators=(",", ": ")
39
+ ))
40
+ """
41
+
42
+ def __init__(self, api_key, url, timeout):
43
+ self.api_key = api_key
44
+ self.url = url
45
+ self.timeout = timeout
46
+
47
+ self.transcriptions: AudioTranscriptions = AudioTranscriptions(self.api_key, self.url, self.timeout)
48
+
49
+ class AudioTranscriptions:
50
+ def __init__(self, api_key, url, timeout):
51
+ self.api_key = api_key
52
+ self.url = url
53
+ self.timeout = timeout
54
+
55
+ def create(
56
+ self,
57
+ model: str,
58
+ file: str,
59
+ language: Optional[str] = "auto",
60
+ temperature: Optional[float] = 0.0,
61
+ prompt: Optional[str] = "",
62
+ timestamp_granularities: Optional[List[str]] = None,
63
+ diarization: Optional[bool] = False,
64
+ response_format: Optional[str] = "json",
65
+ toxicity: Optional[bool] = False,
66
+ pii: Optional[str] = "",
67
+ replace_method: Optional[str] = "",
68
+ entity_list: Optional[List[str]] = None,
69
+ injection: Optional[bool] = False,
70
+ ) -> Dict[str, Any]:
71
+ """
72
+ Creates an audio transcription request to the Prediction Guard /audio/transcriptions API
73
+
74
+ :param model: The model to use
75
+ :param file: Audio file to be transcribed
76
+ :param language: The language of the audio file
77
+ :param temperature: The temperature parameter for model transcription
78
+ :param prompt: A prompt to assist in transcription styling
79
+ :param timestamp_granularities: The timestamp granularities to populate for this transcription
80
+ :param diarization: Whether to diarize the audio
81
+ :param response_format: The response format to use
82
+ :param toxicity: Whether to check for output toxicity
83
+ :param pii: Whether to check for or replace pii
84
+ :param replace_method: Replace method for any PII that is present.
85
+ :param entity_list: List of entities to ignore in the PII check.
86
+ :param injection: Whether to check for prompt injection
87
+ :result: A dictionary containing the transcribed text.
88
+ """
89
+
90
+ # Create a list of tuples, each containing all the parameters for
91
+ # a call to _transcribe_audio
92
+ args = (
93
+ model,
94
+ file,
95
+ language,
96
+ temperature,
97
+ prompt,
98
+ timestamp_granularities,
99
+ diarization,
100
+ response_format,
101
+ pii,
102
+ replace_method,
103
+ entity_list,
104
+ injection,
105
+ toxicity,
106
+ )
107
+
108
+ # Run _transcribe_audio
109
+ choices = self._transcribe_audio(*args)
110
+ return choices
111
+
112
+ def _transcribe_audio(
113
+ self,
114
+ model,
115
+ file,
116
+ language,
117
+ temperature,
118
+ prompt,
119
+ timestamp_granularities,
120
+ diarization,
121
+ response_format,
122
+ pii,
123
+ replace_method,
124
+ entity_list,
125
+ injection,
126
+ toxicity,
127
+ ):
128
+ """
129
+ Function to transcribe an audio file.
130
+ """
131
+
132
+ headers = {
133
+ "Authorization": "Bearer " + self.api_key,
134
+ "User-Agent": "Prediction Guard Python Client: " + __version__,
135
+ "Toxicity": str(toxicity),
136
+ "Pii": str(pii),
137
+ "Replace-Method": str(replace_method),
138
+ "Entity-List": str(entity_list),
139
+ "Injection": str(injection)
140
+ }
141
+
142
+ if timestamp_granularities:
143
+ if diarization and "segment" in timestamp_granularities:
144
+ raise ValueError(
145
+ "Timestamp granularities cannot be set to "
146
+ "`segments` when using diarization."
147
+ )
148
+
149
+ if response_format != "verbose_json":
150
+ raise ValueError(
151
+ "Response format must be set to `verbose_json` "
152
+ "when using timestamp granularities."
153
+ )
154
+
155
+ if diarization and response_format != "verbose_json":
156
+ raise ValueError(
157
+ "Response format must be set to `verbose_json` "
158
+ "when using diarization."
159
+ )
160
+
161
+ with open(file, "rb") as audio_file:
162
+ files = {"file": (file, audio_file, "audio/wav")}
163
+ data = {
164
+ "model": model,
165
+ "language": language,
166
+ "temperature": temperature,
167
+ "prompt": prompt,
168
+ "timestamp_granularities[]": timestamp_granularities,
169
+ "diarization": str(diarization).lower(),
170
+ "response_format": response_format,
171
+ }
172
+
173
+ response = requests.request(
174
+ "POST", self.url + "/audio/transcriptions", headers=headers, files=files, data=data, timeout=self.timeout
175
+ )
176
+
177
+ # If the request was successful, print the proxies.
178
+ if response.status_code == 200:
179
+ ret = response.json()
180
+ return ret
181
+ elif response.status_code == 429:
182
+ raise ValueError(
183
+ "Could not connect to Prediction Guard API. "
184
+ "Too many requests, rate limit or quota exceeded."
185
+ )
186
+ else:
187
+ # Check if there is a json body in the response. Read that in,
188
+ # print out the error field in the json body, and raise an exception.
189
+ err = ""
190
+ try:
191
+ err = response.json()["error"]
192
+ except Exception:
193
+ pass
194
+ raise ValueError("Could not transcribe the audio file. " + err)