gemini_genai_rs/generate/
mod.rs1mod config;
21mod response;
22
23pub use config::GenerateContentConfig;
24pub use response::{BlockReason, Candidate, GenerateContentResponse, PromptFeedback};
25
26use crate::client::Client;
27use crate::client::http::HttpError;
28use crate::protocol::types::ModelId;
29use crate::transport::auth::ServiceEndpoint;
30
31impl Client {
32 pub async fn generate_content(
34 &self,
35 prompt: impl Into<String>,
36 ) -> Result<GenerateContentResponse, GenerateError> {
37 let config = GenerateContentConfig::from_text(prompt);
38 self.generate_content_with(config, None).await
39 }
40
41 pub async fn generate_content_with(
43 &self,
44 config: GenerateContentConfig,
45 model: Option<&ModelId>,
46 ) -> Result<GenerateContentResponse, GenerateError> {
47 let model = model.unwrap_or(self.default_model());
48 let url = self.rest_url_for(ServiceEndpoint::GenerateContent, model);
49 let headers = self.auth_headers().await?;
50
51 let body = config.to_request_body();
52 let json = self
53 .http_client()
54 .post_json(&url, headers, &body)
55 .await
56 .map_err(GenerateError::from)?;
57
58 let response: GenerateContentResponse = serde_json::from_value(json)?;
59 Ok(response)
60 }
61}
62
63impl Client {
64 pub async fn stream_generate_content_with(
86 &self,
87 config: GenerateContentConfig,
88 model: Option<&ModelId>,
89 ) -> Result<
90 futures_util::stream::BoxStream<'static, Result<GenerateContentResponse, GenerateError>>,
91 GenerateError,
92 > {
93 use futures_util::StreamExt;
94
95 let model = model.unwrap_or(self.default_model());
96 let url = self.rest_url_for(ServiceEndpoint::StreamGenerateContent, model);
97 let separator = if url.contains('?') { '&' } else { '?' };
98 let url = format!("{url}{separator}alt=sse");
99 let headers = self.auth_headers().await?;
100
101 let body = config.to_request_body();
102 let events = self
103 .http_client()
104 .post_sse(&url, headers, &body)
105 .await
106 .map_err(GenerateError::from)?;
107 Ok(events
108 .map(|event| Ok(serde_json::from_value(event?)?))
109 .boxed())
110 }
111}
112
113#[derive(Debug, thiserror::Error)]
115pub enum GenerateError {
116 #[error(transparent)]
118 Http(#[from] HttpError),
119
120 #[error("Failed to parse response: {0}")]
122 Parse(#[from] serde_json::Error),
123
124 #[error("Auth error: {0}")]
126 Auth(#[from] crate::session::AuthError),
127
128 #[error("Content blocked: {reason:?}")]
130 SafetyBlocked {
131 reason: BlockReason,
133 },
134
135 #[error("Prompt blocked: {reason:?}")]
137 PromptBlocked {
138 reason: BlockReason,
140 },
141}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146
147 #[test]
148 fn generate_error_display() {
149 let err = GenerateError::SafetyBlocked {
150 reason: BlockReason::Safety,
151 };
152 assert!(err.to_string().contains("blocked"));
153 }
154
155 #[test]
156 fn generate_content_config_from_text() {
157 let config = GenerateContentConfig::from_text("Hello");
158 let body = config.to_request_body();
159 let contents = body.get("contents").unwrap();
160 assert!(contents.is_array());
161 let parts = contents[0].get("parts").unwrap();
162 assert!(parts[0].get("text").unwrap().as_str().unwrap() == "Hello");
163 }
164
165 #[test]
166 fn generate_content_config_with_system() {
167 let config = GenerateContentConfig::from_text("Hello")
168 .system_instruction("You are a helpful assistant");
169 let body = config.to_request_body();
170 assert!(body.get("systemInstruction").is_some());
171 }
172
173 #[test]
174 fn parse_generate_response() {
175 let json = serde_json::json!({
176 "candidates": [{
177 "content": {
178 "parts": [{"text": "Hello world!"}],
179 "role": "model"
180 },
181 "finishReason": "STOP",
182 "safetyRatings": [{
183 "category": "HARM_CATEGORY_HARASSMENT",
184 "probability": "NEGLIGIBLE"
185 }]
186 }],
187 "usageMetadata": {
188 "promptTokenCount": 5,
189 "candidatesTokenCount": 10,
190 "totalTokenCount": 15
191 }
192 });
193
194 let resp: GenerateContentResponse = serde_json::from_value(json).unwrap();
195 assert_eq!(resp.candidates.len(), 1);
196 assert_eq!(resp.text().unwrap(), "Hello world!");
197 assert!(resp.usage_metadata.is_some());
198 }
199}