api.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190
  1. from __future__ import annotations
  2. import uuid
  3. from dataclasses import dataclass, field
  4. from pathlib import Path
  5. from typing import Callable
  6. from .client import HarnessClient, HarnessConfig
  7. from .models import JsonObject, Notification
  8. @dataclass(slots=True)
  9. class DeepSeekHarnessConfig:
  10. """Configuration for launching the local DeepSeek Harness SDK runtime.
  11. The runtime inherits the caller's environment by default, so existing
  12. DEEPSEEK_API_KEY and DEEPSEEK_BASE_URL settings keep working. Use ``env`` to
  13. intentionally override or inject variables for a subprocess.
  14. """
  15. model: str = "deepseek-v4-flash"
  16. cwd: str | None = None
  17. runtime_cwd: str | None = None
  18. session_root: str | None = None
  19. cordis: str | None = None
  20. env: dict[str, str] = field(default_factory=dict)
  21. runtime_bin: str | None = None
  22. launch_args_override: tuple[str, ...] | None = None
  23. request_timeout_seconds: float | None = None
  24. shutdown_timeout_seconds: float | None = 1.0
  25. base_url: str | None = None
  26. api_key: str | None = None
  27. @dataclass(slots=True)
  28. class TurnResult:
  29. session_id: str
  30. status: str
  31. final_response: str
  32. events: list[JsonObject]
  33. notifications: list[Notification]
  34. session_root: str | None = None
  35. class DeepSeekHarness:
  36. """Synchronous high-level SDK for running DeepSeek Harness agent turns."""
  37. def __init__(self, config: DeepSeekHarnessConfig | None = None, **kwargs: object) -> None:
  38. if config is not None and kwargs:
  39. raise TypeError("pass either DeepSeekHarnessConfig or keyword options, not both")
  40. self.config = config or DeepSeekHarnessConfig(**kwargs)
  41. cwd = str(Path(self.config.cwd or Path.cwd()).resolve())
  42. runtime_cwd = str(Path(self.config.runtime_cwd).resolve()) if self.config.runtime_cwd is not None else cwd
  43. self._cwd = cwd
  44. env = dict(self.config.env)
  45. if self.config.session_root is not None:
  46. env["DSH_SESSION_ROOT"] = self.config.session_root
  47. if self.config.cordis is not None:
  48. env["DSH_CORDIS_CONFIG"] = self.config.cordis
  49. env["DSH_CWD"] = cwd
  50. if self.config.base_url is not None:
  51. env["DEEPSEEK_BASE_URL"] = self.config.base_url
  52. if self.config.api_key is not None:
  53. env["DEEPSEEK_API_KEY"] = self.config.api_key
  54. self._client = HarnessClient(
  55. HarnessConfig(
  56. runtime_bin=self.config.runtime_bin,
  57. launch_args_override=self.config.launch_args_override,
  58. cwd=runtime_cwd,
  59. env=env,
  60. request_timeout_seconds=self.config.request_timeout_seconds,
  61. shutdown_timeout_seconds=self.config.shutdown_timeout_seconds,
  62. )
  63. )
  64. self._initialized = False
  65. def __enter__(self) -> "DeepSeekHarness":
  66. self.start()
  67. return self
  68. def __exit__(self, _exc_type, _exc, _tb) -> None:
  69. self.close()
  70. @property
  71. def client(self) -> HarnessClient:
  72. return self._client
  73. def start(self) -> None:
  74. if self._initialized:
  75. return
  76. self._client.start()
  77. self._client.initialize(
  78. cwd=self._cwd,
  79. model=self.config.model,
  80. )
  81. self._initialized = True
  82. def close(self) -> None:
  83. self._client.close()
  84. self._initialized = False
  85. def start_session(self, session_id: str | None = None) -> "Session":
  86. self.start()
  87. return Session(self, session_id or f"session-{uuid.uuid4().hex}")
  88. def run(
  89. self,
  90. input: str | list[JsonObject],
  91. *,
  92. session_id: str | None = None,
  93. on_notification: Callable[[Notification], None] | None = None,
  94. ) -> TurnResult:
  95. return self.start_session(session_id).run(input, on_notification=on_notification)
  96. class Session:
  97. def __init__(self, harness: DeepSeekHarness, session_id: str) -> None:
  98. self.harness = harness
  99. self.id = session_id
  100. def run(
  101. self,
  102. input: str | list[JsonObject],
  103. *,
  104. on_notification: Callable[[Notification], None] | None = None,
  105. ) -> TurnResult:
  106. content_blocks = normalize_input(input)
  107. notifications: list[Notification] = []
  108. events: list[JsonObject] = []
  109. status = "error"
  110. finished = False
  111. def collect(notification: Notification) -> None:
  112. nonlocal finished, status
  113. notifications.append(notification)
  114. if on_notification is not None:
  115. on_notification(notification)
  116. if notification.method == "session.event":
  117. event = notification.payload.get("event")
  118. if isinstance(event, dict):
  119. events.append(event)
  120. if notification.method == "session.finished" and notification.payload.get("sessionId") == self.id:
  121. status = str(notification.payload.get("status") or "ok")
  122. finished = True
  123. with self.harness.client.subscribe_session_notifications(self.id) as subscription:
  124. self.harness.client.session_prompt(
  125. self.id,
  126. content_blocks,
  127. on_notification=collect,
  128. notification_subscription=subscription,
  129. )
  130. while not finished:
  131. notification = subscription.next()
  132. collect(notification)
  133. return TurnResult(
  134. session_id=self.id,
  135. status=status,
  136. final_response=final_response(events),
  137. events=events,
  138. notifications=notifications,
  139. session_root=self.harness.config.session_root,
  140. )
  141. def normalize_input(input: str | list[JsonObject]) -> list[JsonObject]:
  142. if isinstance(input, str):
  143. return [{"type": "text", "text": input}]
  144. return input
  145. def final_response(events: list[JsonObject]) -> str:
  146. for event in reversed(events):
  147. if event.get("type") != "assistant/message":
  148. continue
  149. data = event.get("data")
  150. if not isinstance(data, dict):
  151. continue
  152. content = data.get("content")
  153. if not isinstance(content, list):
  154. continue
  155. parts: list[str] = []
  156. for block in content:
  157. if isinstance(block, dict) and block.get("type") == "text":
  158. parts.append(str(block.get("text") or ""))
  159. return "".join(parts)
  160. return ""