AiProvider.kt 8.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205
  1. package ai
  2. import java.net.URI
  3. import java.net.http.HttpClient
  4. import java.net.http.HttpRequest
  5. import java.net.http.HttpResponse
  6. import java.time.Duration
  7. /**
  8. * A pluggable AI backend used by [RunAiAgentTask].
  9. *
  10. * Implementations: [OllamaProvider] (local), [MistralProvider] (cloud),
  11. * [ClaudeCodeProvider] (claude-code CLI). Selection is handled by [AiProviders.select].
  12. */
  13. interface AiProvider {
  14. /** Human readable provider id, e.g. "claude-code". */
  15. val name: String
  16. /** True when the backend can actually be reached / invoked right now. */
  17. fun isAvailable(): Boolean
  18. /** Whether this provider talks to a remote cloud service (vs a local process). */
  19. val isCloud: Boolean
  20. /** Send [prompt] and return the model's text answer. Throws on transport errors. */
  21. fun complete(prompt: String): String
  22. }
  23. /** Tiny stdlib-only JSON helpers — buildSrc has no JSON dependency on the classpath. */
  24. object Json {
  25. fun escape(s: String): String = buildString {
  26. for (c in s) when (c) {
  27. '\\' -> append("\\\\")
  28. '"' -> append("\\\"")
  29. '\n' -> append("\\n")
  30. '\r' -> append("\\r")
  31. '\t' -> append("\\t")
  32. else -> if (c < ' ') append("\\u%04x".format(c.code)) else append(c)
  33. }
  34. }
  35. /** Best-effort: first string value for [key] anywhere in [json] (handles \" and \\ escapes). */
  36. fun firstString(json: String, key: String): String? {
  37. val marker = "\"$key\""
  38. var at = json.indexOf(marker)
  39. while (at >= 0) {
  40. var j = at + marker.length
  41. while (j < json.length && json[j] != ':') j++
  42. j++
  43. while (j < json.length && json[j].isWhitespace()) j++
  44. if (j < json.length && json[j] == '"') {
  45. val sb = StringBuilder()
  46. j++
  47. while (j < json.length) {
  48. val c = json[j]
  49. when {
  50. c == '\\' && j + 1 < json.length -> {
  51. when (val n = json[j + 1]) {
  52. 'n' -> sb.append('\n'); 'r' -> sb.append('\r'); 't' -> sb.append('\t')
  53. '"' -> sb.append('"'); '\\' -> sb.append('\\'); '/' -> sb.append('/')
  54. 'u' -> {
  55. if (j + 5 < json.length) {
  56. sb.append(json.substring(j + 2, j + 6).toInt(16).toChar()); j += 4
  57. }
  58. }
  59. else -> sb.append(n)
  60. }
  61. j += 2
  62. }
  63. c == '"' -> return sb.toString()
  64. else -> { sb.append(c); j++ }
  65. }
  66. }
  67. }
  68. at = json.indexOf(marker, at + marker.length)
  69. }
  70. return null
  71. }
  72. }
  73. private val http: HttpClient = HttpClient.newBuilder()
  74. .connectTimeout(Duration.ofSeconds(10))
  75. .build()
  76. private fun post(url: String, body: String, headers: Map<String, String>): HttpResponse<String> {
  77. val builder = HttpRequest.newBuilder(URI.create(url))
  78. .timeout(Duration.ofMinutes(5))
  79. .header("Content-Type", "application/json")
  80. .POST(HttpRequest.BodyPublishers.ofString(body))
  81. headers.forEach { (k, v) -> builder.header(k, v) }
  82. return http.send(builder.build(), HttpResponse.BodyHandlers.ofString())
  83. }
  84. /** Local Ollama server (http://localhost:11434 by default). */
  85. class OllamaProvider(
  86. private val host: String = System.getenv("OLLAMA_HOST") ?: "http://localhost:11434",
  87. private val model: String = System.getenv("OLLAMA_MODEL") ?: "llama3.1",
  88. ) : AiProvider {
  89. override val name = "ollama"
  90. override val isCloud = false
  91. override fun isAvailable(): Boolean = runCatching {
  92. val resp = http.send(
  93. HttpRequest.newBuilder(URI.create("$host/api/tags")).timeout(Duration.ofSeconds(3)).GET().build(),
  94. HttpResponse.BodyHandlers.discarding(),
  95. )
  96. resp.statusCode() in 200..299
  97. }.getOrDefault(false)
  98. override fun complete(prompt: String): String {
  99. val body = """{"model":"${Json.escape(model)}","prompt":"${Json.escape(prompt)}","stream":false}"""
  100. val resp = post("$host/api/generate", body, emptyMap())
  101. check(resp.statusCode() in 200..299) { "ollama HTTP ${resp.statusCode()}: ${resp.body().take(300)}" }
  102. return Json.firstString(resp.body(), "response") ?: resp.body()
  103. }
  104. /** Best-effort `ollama serve` launch when the daemon is down. Returns true if it became reachable. */
  105. fun ensureServing(log: (String) -> Unit): Boolean {
  106. if (isAvailable()) return true
  107. if (which("ollama") == null) { log("ollama binary not found on PATH"); return false }
  108. log("starting `ollama serve` ...")
  109. runCatching {
  110. ProcessBuilder("ollama", "serve").redirectOutput(ProcessBuilder.Redirect.DISCARD)
  111. .redirectError(ProcessBuilder.Redirect.DISCARD).start()
  112. }.onFailure { log("failed to start ollama: ${it.message}"); return false }
  113. repeat(20) { if (isAvailable()) return true; Thread.sleep(500) }
  114. return isAvailable()
  115. }
  116. }
  117. /** Mistral cloud chat completions. */
  118. class MistralProvider(
  119. private val apiKey: String,
  120. private val model: String = System.getenv("MISTRAL_MODEL") ?: "mistral-large-latest",
  121. ) : AiProvider {
  122. override val name = "mistral"
  123. override val isCloud = true
  124. override fun isAvailable() = apiKey.isNotBlank()
  125. override fun complete(prompt: String): String {
  126. val body = """{"model":"${Json.escape(model)}","messages":[{"role":"user","content":"${Json.escape(prompt)}"}]}"""
  127. val resp = post(
  128. "https://api.mistral.ai/v1/chat/completions", body,
  129. mapOf("Authorization" to "Bearer $apiKey"),
  130. )
  131. check(resp.statusCode() in 200..299) { "mistral HTTP ${resp.statusCode()}: ${resp.body().take(300)}" }
  132. return Json.firstString(resp.body(), "content") ?: resp.body()
  133. }
  134. }
  135. /**
  136. * The claude-code CLI in headless print mode: `claude -p "<prompt>"`.
  137. * Authenticates via the user's existing claude-code login — no API key in this repo.
  138. */
  139. class ClaudeCodeProvider(
  140. private val command: List<String> = (System.getenv("CLAUDE_CODE_CMD")?.split(" ") ?: listOf("claude", "-p")),
  141. ) : AiProvider {
  142. override val name = "claude-code"
  143. override val isCloud = true
  144. override fun isAvailable() = which(command.first()) != null
  145. override fun complete(prompt: String): String {
  146. val proc = ProcessBuilder(command + prompt)
  147. .redirectErrorStream(true)
  148. .start()
  149. val out = proc.inputStream.bufferedReader().readText()
  150. proc.waitFor()
  151. return out.trim()
  152. }
  153. }
  154. /** Locate an executable on PATH (returns null when absent). */
  155. fun which(bin: String): String? {
  156. val path = System.getenv("PATH") ?: return null
  157. return path.split(java.io.File.pathSeparator).map { java.io.File(it, bin) }
  158. .firstOrNull { it.canExecute() }?.absolutePath
  159. }
  160. object AiProviders {
  161. /**
  162. * Pick a provider. Honors an explicit `provider=` credential, otherwise
  163. * prefers cloud (claude-code, then mistral) and falls back to local ollama
  164. * (starting `ollama serve` on demand), per ai-todo.txt.
  165. */
  166. fun select(creds: Map<String, String>, log: (String) -> Unit): AiProvider? {
  167. val forced = creds["provider"]?.lowercase()
  168. val claude = ClaudeCodeProvider()
  169. val mistral = creds["MISTRAL_API_KEY"]?.let { MistralProvider(it) }
  170. val ollama = OllamaProvider(
  171. host = creds["OLLAMA_HOST"] ?: System.getenv("OLLAMA_HOST") ?: "http://localhost:11434",
  172. model = creds["OLLAMA_MODEL"] ?: System.getenv("OLLAMA_MODEL") ?: "llama3.1",
  173. )
  174. when (forced) {
  175. "claude-code", "claude" -> if (claude.isAvailable()) return claude
  176. "mistral" -> if (mistral?.isAvailable() == true) return mistral
  177. "ollama" -> { if (ollama.ensureServing(log)) return ollama }
  178. }
  179. // prefer cloud
  180. if (claude.isAvailable()) return claude
  181. if (mistral?.isAvailable() == true) return mistral
  182. // local fallback — start the daemon if needed
  183. if (ollama.ensureServing(log)) return ollama
  184. return null
  185. }
  186. }