Coverage for src / lexigram / contracts / ai / relay / dto / gemini.py: 41%

229 statements  

« prev     ^ index     » next       coverage.py v7.13.5, created at 2026-08-19 05:41 +0800

1"""Gemini ``generateContent`` wire DTO family. 

2 

3Field names follow the Gemini API camelCase wire format; the DTO layer 

4accepts documented snake_case aliases only where noted. 

5""" 

6 

7from __future__ import annotations 

8 

9from dataclasses import dataclass, field 

10from typing import Any 

11 

12from lexigram.contracts.ai.relay.dto.common import require_field 

13 

14__all__ = [ 

15 "GeminiCandidate", 

16 "GeminiContent", 

17 "GeminiGroundingMetadata", 

18 "GeminiPart", 

19 "GeminiPromptFeedback", 

20 "GeminiRequest", 

21 "GeminiResponse", 

22 "GeminiSafetyRating", 

23 "GeminiUsageMetadata", 

24] 

25 

26 

27@dataclass(frozen=True) 

28class GeminiPart: 

29 """A part inside a Gemini content. 

30 

31 Attributes: 

32 text: Text payload, or ``None``. 

33 inline_data: ``{"mime_type": ..., "data": base64}`` or ``None``. 

34 file_data: ``{"mime_type": ..., "file_uri": ...}`` or ``None`` 

35 (already-resolved file reference). 

36 function_call: ``{"name": ..., "args": {...}}`` or ``None``. 

37 function_response: ``{"name": ..., "response": {...}}`` or ``None``. 

38 thought: Whether this is a thinking part. 

39 thought_signature: Thought signature for thinking parts, or ``None``. 

40 passthrough: Unknown fields preserved verbatim. 

41 """ 

42 

43 text: str | None = None 

44 inline_data: dict[str, Any] | None = None 

45 file_data: dict[str, Any] | None = None 

46 function_call: dict[str, Any] | None = None 

47 function_response: dict[str, Any] | None = None 

48 thought: bool = False 

49 thought_signature: str | None = None 

50 passthrough: dict[str, Any] = field(default_factory=dict) 

51 

52 def to_dict(self) -> dict[str, Any]: 

53 """Serialize to wire dict (camelCase field names).""" 

54 data: dict[str, Any] = {**self.passthrough} 

55 if self.text is not None: 

56 data["text"] = self.text 

57 if self.inline_data is not None: 

58 data["inlineData"] = self.inline_data 

59 if self.file_data is not None: 

60 data["fileData"] = self.file_data 

61 if self.function_call is not None: 

62 data["functionCall"] = self.function_call 

63 if self.function_response is not None: 

64 data["functionResponse"] = self.function_response 

65 if self.thought: 

66 data["thought"] = True 

67 if self.thought_signature is not None: 

68 data["thoughtSignature"] = self.thought_signature 

69 return data 

70 

71 @classmethod 

72 def from_dict(cls, data: dict[str, Any]) -> GeminiPart: 

73 """Build a part from a wire dict, capturing unknown keys.""" 

74 known = { 

75 "text", 

76 "inlineData", 

77 "fileData", 

78 "functionCall", 

79 "functionResponse", 

80 "thought", 

81 "thoughtSignature", 

82 } 

83 return cls( 

84 text=data.get("text"), 

85 inline_data=data.get("inlineData"), 

86 file_data=data.get("fileData"), 

87 function_call=data.get("functionCall"), 

88 function_response=data.get("functionResponse"), 

89 thought=bool(data.get("thought", False)), 

90 thought_signature=data.get("thoughtSignature"), 

91 passthrough={k: v for k, v in data.items() if k not in known}, 

92 ) 

93 

94 

95@dataclass(frozen=True) 

96class GeminiContent: 

97 """One content turn in Gemini format. 

98 

99 Attributes: 

100 role: ``user``, ``model``, ``function``. 

101 parts: Content parts. 

102 """ 

103 

104 role: str 

105 parts: list[GeminiPart] 

106 

107 def to_dict(self) -> dict[str, Any]: 

108 """Serialize to wire dict.""" 

109 return {"role": self.role, "parts": [p.to_dict() for p in self.parts]} 

110 

111 @classmethod 

112 def from_dict(cls, data: dict[str, Any]) -> GeminiContent: 

113 """Build a content from a wire dict.""" 

114 return cls( 

115 role=data.get("role", "user"), 

116 parts=[GeminiPart.from_dict(p) for p in data.get("parts", [])], 

117 ) 

118 

119 

120@dataclass(frozen=True) 

121class GeminiRequest: 

122 """Gemini ``generateContent`` request body. 

123 

124 Attributes: 

125 contents: Conversation turns. 

126 system_instruction: ``{"parts": [{"text": ...}]}`` or ``None``. 

127 Serialized as ``systemInstruction``; ``from_dict`` accepts 

128 both wire casings. 

129 generation_config: Generation config dict (empty when unset). 

130 Serialized as ``generationConfig``. 

131 safety_settings: Safety threshold list, or ``None``. Serialized 

132 as ``safetySettings``. 

133 tools: Tool definitions list, or ``None``. 

134 tool_config: Tool configuration dict, or ``None``. Serialized 

135 as ``toolConfig``. 

136 passthrough: Unknown fields preserved verbatim. 

137 """ 

138 

139 contents: list[GeminiContent] 

140 system_instruction: dict[str, Any] | None = None 

141 generation_config: dict[str, Any] = field(default_factory=dict) 

142 safety_settings: list[dict[str, Any]] | None = None 

143 tools: list[dict[str, Any]] | None = None 

144 tool_config: dict[str, Any] | None = None 

145 passthrough: dict[str, Any] = field(default_factory=dict) 

146 

147 def to_dict(self) -> dict[str, Any]: 

148 """Serialize to wire dict (camelCase field names).""" 

149 data: dict[str, Any] = { 

150 **self.passthrough, 

151 "contents": [c.to_dict() for c in self.contents], 

152 } 

153 if self.system_instruction is not None: 

154 data["systemInstruction"] = self.system_instruction 

155 if self.generation_config: 

156 data["generationConfig"] = self.generation_config 

157 if self.safety_settings is not None: 

158 data["safetySettings"] = self.safety_settings 

159 if self.tools is not None: 

160 data["tools"] = self.tools 

161 if self.tool_config is not None: 

162 data["toolConfig"] = self.tool_config 

163 return data 

164 

165 @classmethod 

166 def from_dict(cls, data: dict[str, Any]) -> GeminiRequest: 

167 """Build a request from a wire dict, capturing unknown keys. 

168 

169 ``systemInstruction`` and ``generationConfig`` are the canonical 

170 wire keys; snake_case ``system_instruction`` is accepted for 

171 compatibility. 

172 

173 Raises: 

174 RelayError: With code ``malformed_payload`` when ``contents`` 

175 is absent. 

176 """ 

177 known = { 

178 "contents", 

179 "systemInstruction", 

180 "system_instruction", 

181 "generationConfig", 

182 "safetySettings", 

183 "tools", 

184 "toolConfig", 

185 } 

186 system = data.get("systemInstruction", data.get("system_instruction")) 

187 return cls( 

188 contents=[ 

189 GeminiContent.from_dict(c) for c in require_field(data, "contents") 

190 ], 

191 system_instruction=system, 

192 generation_config=data.get("generationConfig", {}), 

193 safety_settings=data.get("safetySettings"), 

194 tools=data.get("tools"), 

195 tool_config=data.get("toolConfig"), 

196 passthrough={k: v for k, v in data.items() if k not in known}, 

197 ) 

198 

199 

200@dataclass(frozen=True) 

201class GeminiSafetyRating: 

202 """A Gemini safety rating entry. 

203 

204 Attributes: 

205 category: Harm category. 

206 probability: Probability level. 

207 blocked: Whether the content was blocked. 

208 severity: Severity level, or ``None``. 

209 probability_score: Raw probability score, or ``None``. 

210 severity_score: Raw severity score, or ``None``. 

211 passthrough: Unknown fields preserved verbatim. 

212 """ 

213 

214 category: str = "" 

215 probability: str = "" 

216 blocked: bool = False 

217 severity: str | None = None 

218 probability_score: float | None = None 

219 severity_score: float | None = None 

220 passthrough: dict[str, Any] = field(default_factory=dict) 

221 

222 def to_dict(self) -> dict[str, Any]: 

223 """Serialize to wire dict (camelCase field names).""" 

224 data: dict[str, Any] = {**self.passthrough} 

225 if self.category: 

226 data["category"] = self.category 

227 if self.probability: 

228 data["probability"] = self.probability 

229 if self.blocked: 

230 data["blocked"] = True 

231 if self.severity is not None: 

232 data["severity"] = self.severity 

233 if self.probability_score is not None: 

234 data["probabilityScore"] = self.probability_score 

235 if self.severity_score is not None: 

236 data["severityScore"] = self.severity_score 

237 return data 

238 

239 @classmethod 

240 def from_dict(cls, data: dict[str, Any]) -> GeminiSafetyRating: 

241 """Build a rating from a wire dict, capturing unknown keys.""" 

242 known = { 

243 "category", 

244 "probability", 

245 "blocked", 

246 "severity", 

247 "probabilityScore", 

248 "severityScore", 

249 } 

250 return cls( 

251 category=data.get("category", ""), 

252 probability=data.get("probability", ""), 

253 blocked=bool(data.get("blocked", False)), 

254 severity=data.get("severity"), 

255 probability_score=data.get("probabilityScore"), 

256 severity_score=data.get("severityScore"), 

257 passthrough={k: v for k, v in data.items() if k not in known}, 

258 ) 

259 

260 

261@dataclass(frozen=True) 

262class GeminiGroundingMetadata: 

263 """Gemini grounding metadata for a candidate. 

264 

265 Attributes: 

266 grounding_chunks: Raw grounding chunk list. 

267 grounding_supports: Raw grounding support list. 

268 web_search_queries: Web search queries, if any. 

269 retrieval_metadata: Raw retrieval metadata list. 

270 passthrough: Unknown fields preserved verbatim. 

271 """ 

272 

273 grounding_chunks: list[dict[str, Any]] = field(default_factory=list) 

274 grounding_supports: list[dict[str, Any]] = field(default_factory=list) 

275 web_search_queries: list[str] = field(default_factory=list) 

276 retrieval_metadata: list[dict[str, Any]] = field(default_factory=list) 

277 passthrough: dict[str, Any] = field(default_factory=dict) 

278 

279 def to_dict(self) -> dict[str, Any]: 

280 """Serialize to wire dict (camelCase field names).""" 

281 data: dict[str, Any] = {**self.passthrough} 

282 if self.grounding_chunks: 

283 data["groundingChunks"] = self.grounding_chunks 

284 if self.grounding_supports: 

285 data["groundingSupports"] = self.grounding_supports 

286 if self.web_search_queries: 

287 data["webSearchQueries"] = self.web_search_queries 

288 if self.retrieval_metadata: 

289 data["retrievalMetadata"] = self.retrieval_metadata 

290 return data 

291 

292 @classmethod 

293 def from_dict(cls, data: dict[str, Any]) -> GeminiGroundingMetadata: 

294 """Build metadata from a wire dict, capturing unknown keys.""" 

295 known = { 

296 "groundingChunks", 

297 "groundingSupports", 

298 "webSearchQueries", 

299 "retrievalMetadata", 

300 } 

301 return cls( 

302 grounding_chunks=data.get("groundingChunks", []), 

303 grounding_supports=data.get("groundingSupports", []), 

304 web_search_queries=data.get("webSearchQueries", []), 

305 retrieval_metadata=data.get("retrievalMetadata", []), 

306 passthrough={k: v for k, v in data.items() if k not in known}, 

307 ) 

308 

309 

310@dataclass(frozen=True) 

311class GeminiCandidate: 

312 """One candidate in a Gemini response. 

313 

314 Attributes: 

315 content: Candidate content, or ``None``. 

316 finish_reason: ``STOP``, ``MAX_TOKENS``, ``SAFETY``, etc. 

317 index: Candidate index. 

318 safety_ratings: Safety ratings, or ``None``. 

319 grounding_metadata: Grounding metadata, or ``None``. 

320 citation_metadata: Raw citation metadata, or ``None``. 

321 token_count: Token count for the candidate. 

322 avg_logprobs: Average log probability. 

323 passthrough: Unknown fields preserved verbatim. 

324 """ 

325 

326 content: GeminiContent | None = None 

327 finish_reason: str | None = None 

328 index: int | None = None 

329 safety_ratings: list[GeminiSafetyRating] | None = None 

330 grounding_metadata: GeminiGroundingMetadata | None = None 

331 citation_metadata: dict[str, Any] | None = None 

332 token_count: int | None = None 

333 avg_logprobs: float | None = None 

334 passthrough: dict[str, Any] = field(default_factory=dict) 

335 

336 def to_dict(self) -> dict[str, Any]: 

337 """Serialize to wire dict (camelCase field names).""" 

338 data: dict[str, Any] = {**self.passthrough} 

339 if self.content is not None: 

340 data["content"] = self.content.to_dict() 

341 if self.finish_reason is not None: 

342 data["finishReason"] = self.finish_reason 

343 if self.index is not None: 

344 data["index"] = self.index 

345 if self.safety_ratings is not None: 

346 data["safetyRatings"] = [r.to_dict() for r in self.safety_ratings] 

347 if self.grounding_metadata is not None: 

348 data["groundingMetadata"] = self.grounding_metadata.to_dict() 

349 if self.citation_metadata is not None: 

350 data["citationMetadata"] = self.citation_metadata 

351 if self.token_count is not None: 

352 data["tokenCount"] = self.token_count 

353 if self.avg_logprobs is not None: 

354 data["avgLogprobs"] = self.avg_logprobs 

355 return data 

356 

357 @classmethod 

358 def from_dict(cls, data: dict[str, Any]) -> GeminiCandidate: 

359 """Build a candidate from a wire dict, capturing unknown keys.""" 

360 known = { 

361 "content", 

362 "finishReason", 

363 "index", 

364 "safetyRatings", 

365 "groundingMetadata", 

366 "citationMetadata", 

367 "tokenCount", 

368 "avgLogprobs", 

369 } 

370 content = data.get("content") 

371 safety = data.get("safetyRatings") 

372 grounding = data.get("groundingMetadata") 

373 return cls( 

374 content=GeminiContent.from_dict(content) 

375 if isinstance(content, dict) 

376 else None, 

377 finish_reason=data.get("finishReason"), 

378 index=data.get("index"), 

379 safety_ratings=( 

380 [GeminiSafetyRating.from_dict(r) for r in safety] 

381 if isinstance(safety, list) 

382 else None 

383 ), 

384 grounding_metadata=( 

385 GeminiGroundingMetadata.from_dict(grounding) 

386 if isinstance(grounding, dict) 

387 else None 

388 ), 

389 citation_metadata=data.get("citationMetadata"), 

390 token_count=data.get("tokenCount"), 

391 avg_logprobs=data.get("avgLogprobs"), 

392 passthrough={k: v for k, v in data.items() if k not in known}, 

393 ) 

394 

395 

396@dataclass(frozen=True) 

397class GeminiUsageMetadata: 

398 """Gemini usage metadata. 

399 

400 Attributes: 

401 prompt_token_count: Input tokens. 

402 candidates_token_count: Output tokens. 

403 total_token_count: Total tokens. 

404 cached_content_token_count: Cached input tokens, or ``None``. 

405 thoughts_token_count: Thinking tokens, or ``None``. 

406 tool_use_prompt_token_count: Tokens spent on tool-use prompt parts. 

407 prompt_tokens_details: Per-input-category details, or ``None``. 

408 tool_use_prompt_tokens_details: Per-tool-call details, or ``None``. 

409 candidates_tokens_details: Per-output-category details, or ``None``. 

410 passthrough: Unknown fields preserved verbatim. 

411 """ 

412 

413 prompt_token_count: int = 0 

414 candidates_token_count: int = 0 

415 total_token_count: int = 0 

416 cached_content_token_count: int | None = None 

417 thoughts_token_count: int | None = None 

418 tool_use_prompt_token_count: int = 0 

419 prompt_tokens_details: Any = None 

420 tool_use_prompt_tokens_details: Any = None 

421 candidates_tokens_details: Any = None 

422 passthrough: dict[str, Any] = field(default_factory=dict) 

423 

424 def to_dict(self) -> dict[str, Any]: 

425 """Serialize to wire dict (camelCase field names).""" 

426 data: dict[str, Any] = { 

427 **self.passthrough, 

428 "promptTokenCount": self.prompt_token_count, 

429 "toolUsePromptTokenCount": self.tool_use_prompt_token_count, 

430 "candidatesTokenCount": self.candidates_token_count, 

431 "totalTokenCount": self.total_token_count, 

432 "promptTokensDetails": self.prompt_tokens_details, 

433 "toolUsePromptTokensDetails": self.tool_use_prompt_tokens_details, 

434 "candidatesTokensDetails": self.candidates_tokens_details, 

435 } 

436 if self.cached_content_token_count is not None: 

437 data["cachedContentTokenCount"] = self.cached_content_token_count 

438 if self.thoughts_token_count is not None: 

439 data["thoughtsTokenCount"] = self.thoughts_token_count 

440 return data 

441 

442 @classmethod 

443 def from_dict(cls, data: dict[str, Any]) -> GeminiUsageMetadata: 

444 """Build usage from a wire dict, capturing unknown keys.""" 

445 known = { 

446 "promptTokenCount", 

447 "toolUsePromptTokenCount", 

448 "candidatesTokenCount", 

449 "totalTokenCount", 

450 "cachedContentTokenCount", 

451 "thoughtsTokenCount", 

452 "promptTokensDetails", 

453 "toolUsePromptTokensDetails", 

454 "candidatesTokensDetails", 

455 } 

456 return cls( 

457 prompt_token_count=data.get("promptTokenCount", 0), 

458 candidates_token_count=data.get("candidatesTokenCount", 0), 

459 total_token_count=data.get("totalTokenCount", 0), 

460 cached_content_token_count=data.get("cachedContentTokenCount"), 

461 thoughts_token_count=data.get("thoughtsTokenCount"), 

462 tool_use_prompt_token_count=data.get("toolUsePromptTokenCount", 0), 

463 prompt_tokens_details=data.get("promptTokensDetails"), 

464 tool_use_prompt_tokens_details=data.get("toolUsePromptTokensDetails"), 

465 candidates_tokens_details=data.get("candidatesTokensDetails"), 

466 passthrough={k: v for k, v in data.items() if k not in known}, 

467 ) 

468 

469 

470@dataclass(frozen=True) 

471class GeminiPromptFeedback: 

472 """Prompt-level feedback on a Gemini response. 

473 

474 Attributes: 

475 block_reason: Block reason, or ``None``. 

476 safety_ratings: Safety ratings, or ``None``. 

477 passthrough: Unknown fields preserved verbatim. 

478 """ 

479 

480 block_reason: str | None = None 

481 safety_ratings: list[GeminiSafetyRating] | None = None 

482 passthrough: dict[str, Any] = field(default_factory=dict) 

483 

484 def to_dict(self) -> dict[str, Any]: 

485 """Serialize to wire dict (camelCase field names).""" 

486 data: dict[str, Any] = {**self.passthrough} 

487 if self.block_reason is not None: 

488 data["blockReason"] = self.block_reason 

489 if self.safety_ratings is not None: 

490 data["safetyRatings"] = [r.to_dict() for r in self.safety_ratings] 

491 return data 

492 

493 @classmethod 

494 def from_dict(cls, data: dict[str, Any]) -> GeminiPromptFeedback: 

495 """Build feedback from a wire dict, capturing unknown keys.""" 

496 known = {"blockReason", "safetyRatings"} 

497 safety = data.get("safetyRatings") 

498 return cls( 

499 block_reason=data.get("blockReason"), 

500 safety_ratings=( 

501 [GeminiSafetyRating.from_dict(r) for r in safety] 

502 if isinstance(safety, list) 

503 else None 

504 ), 

505 passthrough={k: v for k, v in data.items() if k not in known}, 

506 ) 

507 

508 

509@dataclass(frozen=True) 

510class GeminiResponse: 

511 """Gemini ``generateContent`` response body (also used for stream chunks). 

512 

513 Attributes: 

514 candidates: Candidate list. 

515 prompt_feedback: Prompt-level feedback, or ``None``. 

516 usage_metadata: Token usage, or ``None``. 

517 model_version: Model version, or ``None``. 

518 create_time: Response creation time, or ``None``. 

519 response_id: Response id, or ``None``. 

520 passthrough: Unknown fields preserved verbatim. 

521 """ 

522 

523 candidates: list[GeminiCandidate] = field(default_factory=list) 

524 prompt_feedback: GeminiPromptFeedback | None = None 

525 usage_metadata: GeminiUsageMetadata | None = None 

526 model_version: str | None = None 

527 create_time: str | None = None 

528 response_id: str | None = None 

529 passthrough: dict[str, Any] = field(default_factory=dict) 

530 

531 def to_dict(self) -> dict[str, Any]: 

532 """Serialize to wire dict (camelCase field names).""" 

533 data: dict[str, Any] = {**self.passthrough} 

534 if self.candidates: 

535 data["candidates"] = [c.to_dict() for c in self.candidates] 

536 if self.prompt_feedback is not None: 

537 data["promptFeedback"] = self.prompt_feedback.to_dict() 

538 if self.usage_metadata is not None: 

539 data["usageMetadata"] = self.usage_metadata.to_dict() 

540 if self.model_version is not None: 

541 data["modelVersion"] = self.model_version 

542 if self.create_time is not None: 

543 data["createTime"] = self.create_time 

544 if self.response_id is not None: 

545 data["responseId"] = self.response_id 

546 return data 

547 

548 @classmethod 

549 def from_dict(cls, data: dict[str, Any]) -> GeminiResponse: 

550 """Build a response from a wire dict, capturing unknown keys.""" 

551 known = { 

552 "candidates", 

553 "promptFeedback", 

554 "usageMetadata", 

555 "modelVersion", 

556 "createTime", 

557 "responseId", 

558 } 

559 feedback = data.get("promptFeedback") 

560 usage = data.get("usageMetadata") 

561 return cls( 

562 candidates=[ 

563 GeminiCandidate.from_dict(c) for c in data.get("candidates", []) 

564 ], 

565 prompt_feedback=( 

566 GeminiPromptFeedback.from_dict(feedback) 

567 if isinstance(feedback, dict) 

568 else None 

569 ), 

570 usage_metadata=( 

571 GeminiUsageMetadata.from_dict(usage) 

572 if isinstance(usage, dict) 

573 else None 

574 ), 

575 model_version=data.get("modelVersion"), 

576 create_time=data.get("createTime"), 

577 response_id=data.get("responseId"), 

578 passthrough={k: v for k, v in data.items() if k not in known}, 

579 )