Compare commits
645 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 11181c14d5 | |||
| 96c380caa3 | |||
| 59c77c65b1 | |||
| ae461d1c2d | |||
| 2a7cccd598 | |||
| aae8d0a379 | |||
| f3ca381ab3 | |||
| 8f6a568646 | |||
| 91e2a87be5 | |||
| 03771df508 | |||
| 90e773b30b | |||
| 40b896bcb5 | |||
| 174f03b92d | |||
| 194402fc88 | |||
| c6efdaa13f | |||
| 08aa0d0e05 | |||
| 8eff551bb6 | |||
| 862c1e9743 | |||
| d7711484ef | |||
| ba7d908276 | |||
| e57ecd4588 | |||
| 261845830d | |||
| 8efd4ad0ab | |||
| 01e674e1bc | |||
| 7f26ecc19a | |||
| 395baab7d5 | |||
| ccf4173990 | |||
| a57c52c676 | |||
| 384bc13a5b | |||
| 5622b40cf3 | |||
| 802f71bd46 | |||
| cc9dfb1cc9 | |||
| 12e4f34005 | |||
| b0f9a8d785 | |||
| 15cc389b3d | |||
| 3e67e64067 | |||
| bb9bed82e1 | |||
| 1a01fb5d74 | |||
| 087781556b | |||
| a1bfbd92ca | |||
| 421a664336 | |||
| 57176b359a | |||
| b2660fc38c | |||
| 940db565b3 | |||
| dbb6b7abde | |||
| c8c46eab16 | |||
| 0cc995eb80 | |||
| b1233d42dc | |||
| b2e28249fd | |||
| af4ae1aef5 | |||
| c46ae17084 | |||
| c21fc0a272 | |||
| 8a0ab17936 | |||
| 38668c4ce5 | |||
| 7a3fcc7c8e | |||
| 49770949da | |||
| 6f10406d5b | |||
| 9042068094 | |||
| 81ff0519dc | |||
| e0acc6155e | |||
| 690b903f85 | |||
| d2452c54d5 | |||
| a7b9e175c1 | |||
| be3dd272c3 | |||
| df54d8498c | |||
| 1d117ff277 | |||
| f086d77756 | |||
| 5dacab7e4c | |||
| bd54eaa0a4 | |||
| 2b244888ec | |||
| 7568a6bc1f | |||
| 682922f690 | |||
| 598b668131 | |||
| 6546f4022f | |||
| 7214d4099d | |||
| f2f010a350 | |||
| 7ccfe68f3f | |||
| b1dccf17ea | |||
| c063a00c8f | |||
| 9460b9ecd6 | |||
| 57461f671f | |||
| bd307f3a11 | |||
| 6eb467e70b | |||
| 6593ac9b5b | |||
| 6894d45be8 | |||
| 086425377d | |||
| de3f588fbf | |||
| 5282a5028b | |||
| 42d8b9bf10 | |||
| f356de36a6 | |||
| 8ea186441a | |||
| 477481f9d9 | |||
| 76972449c7 | |||
| 8f9dfa159e | |||
| d82f0ed2a4 | |||
| 05a5e5f8a0 | |||
| 526c571b10 | |||
| 49f23560fd | |||
| c02be519f6 | |||
| fc05b5eed2 | |||
| c438966b28 | |||
| d4f1fbd110 | |||
| cf5e0dd3bd | |||
| de3785346f | |||
| 2dc1e227eb | |||
| cd2baa9588 | |||
| 4b6a969df2 | |||
| 8bb1d6c0e3 | |||
| ac052bbb4c | |||
| b4ffb35d71 | |||
| 63969b596d | |||
| 3563c1d94f | |||
| 92d95dee68 | |||
| a27e5230c7 | |||
| 5c8975bc28 | |||
| cedb3aa744 | |||
| 044a85ccd2 | |||
| 0198e50e7f | |||
| faea53be52 | |||
| 9cffe9d457 | |||
| d348076f40 | |||
| 9285c6dad8 | |||
| b56234317a | |||
| 565d9647ac | |||
| d53bfa35c5 | |||
| fbd1d709ca | |||
| 3ce6523faf | |||
| a13904185d | |||
| 2364e6b130 | |||
| 721a03c25b | |||
| f75bfcda51 | |||
| b9ad694467 | |||
| d2283397a4 | |||
| b11932f91b | |||
| 7959495a13 | |||
| 9c7347eedb | |||
| 331056cdc8 | |||
| 385f9756c1 | |||
| 7f1aa3b0f6 | |||
| 0d51406149 | |||
| a4c9c779c9 | |||
| 4b0c91190a | |||
| 35ea2bfb52 | |||
| 8fe774b056 | |||
| ccb3083183 | |||
| f00ee1b5cc | |||
| c407d2e20f | |||
| 89b0ecdbf3 | |||
| 4e04ac5b72 | |||
| 80f1f4fa0f | |||
| 692dc491ac | |||
| 9e51ec6fdd | |||
| 22a65b640d | |||
| d7c0eec0e9 | |||
| f41584e10b | |||
| d5ffa00263 | |||
| a297445f92 | |||
| ab3787f86a | |||
| 29e7fd383d | |||
| 0d8ac4f24b | |||
| 73928a2d78 | |||
| 27a97c3097 | |||
| 5c829942d7 | |||
| 7cfec02416 | |||
| 50719ef256 | |||
| 56cc2fef85 | |||
| 52f8d3a3a5 | |||
| 26e3452ef6 | |||
| 20c06d4897 | |||
| e48bc1cb71 | |||
| da74c325d6 | |||
| 27fee3256e | |||
| 558360b558 | |||
| 49b03c36eb | |||
| 155c4eaa40 | |||
| c134a16e19 | |||
| 3831198f19 | |||
| 92df3e1844 | |||
| f704e3d761 | |||
| 65828e9666 | |||
| 28f3e81b4c | |||
| aa3dd00409 | |||
| 05f54334ba | |||
| 1e4c011b7c | |||
| 06822f236c | |||
| bd501cce34 | |||
| 58435dba52 | |||
| f4a3617646 | |||
| 7c6b6755f2 | |||
| 4b2aaea49a | |||
| 1965dd8661 | |||
| 03b4412c9a | |||
| 210e8864f6 | |||
| 7f522cb4fe | |||
| 0b7c162d1b | |||
| a2d2ddc5a2 | |||
| 2961e5ee88 | |||
| c4237fecb8 | |||
| 64721f5967 | |||
| 65db3a4fcd | |||
| ff15f515cc | |||
| d5b982c980 | |||
| 3e493233cf | |||
| 011dffdabc | |||
| 4f11de23a2 | |||
| 066f36fa13 | |||
| 632159261c | |||
| 3647ff1afa | |||
| e89b71aa60 | |||
| 729cab11be | |||
| 080a5c06f3 | |||
| 117750eda7 | |||
| 028dbe79d3 | |||
| d08ca535c6 | |||
| 5aa8353613 | |||
| 91de173c78 | |||
| 70629e254c | |||
| 68f3ab2962 | |||
| 897b444d27 | |||
| d8720a0b35 | |||
| b9e809aeb6 | |||
| 3c7e1e5b8f | |||
| 6e5844a68f | |||
| dd6cbe3b45 | |||
| 5c9cd353b9 | |||
| 4505300c8d | |||
| 29533cce7f | |||
| db72fb507f | |||
| 6e23e98cca | |||
| 4a9562229f | |||
| 10e0364163 | |||
| d3284ee881 | |||
| 6dc8a25579 | |||
| a813f8afd9 | |||
| 5bc33adeb0 | |||
| 5ae3d47945 | |||
| fe6a2c4b83 | |||
| 259842243d | |||
| fab5f85eee | |||
| cc266dd6b1 | |||
| 65cec64445 | |||
| 3503142af6 | |||
| a7d3ef8087 | |||
| 372a6272b7 | |||
| fdafebd47f | |||
| 210d71c590 | |||
| 649f2ca121 | |||
| c627aadd39 | |||
| 28f0d15cda | |||
| ac7fdcecf2 | |||
| edd3823881 | |||
| 8862531738 | |||
| 2fdf961fee | |||
| 81316ffb2d | |||
| 4a3d6c0318 | |||
| 793b3f32af | |||
| 97d00808e1 | |||
| 47e5ef6219 | |||
| a011dca693 | |||
| ac5ea37d14 | |||
| 7ebaa2cec0 | |||
| f9734c7e67 | |||
| 1c6f4a2078 | |||
| 2f1622579f | |||
| 16dc1121e4 | |||
| 9274df7405 | |||
| dfe81ad795 | |||
| 5d06d93ac5 | |||
| 7089d1c179 | |||
| 53699f9a4e | |||
| d5e50a690e | |||
| 97684eab7b | |||
| 2f855c07c7 | |||
| d3e877c845 | |||
| 973edb6aa4 | |||
| 859a0aabbb | |||
| 1dac00c262 | |||
| e515f69bdc | |||
| 4b648b3413 | |||
| 6408ccb6db | |||
| 0303c6d43d | |||
| dd81a69585 | |||
| 55fb99f0f4 | |||
| c542e6403e | |||
| 79b572cdae | |||
| 28ad1f0106 | |||
| 0fe30de92d | |||
| ccaa3113d2 | |||
| 62fe970476 | |||
| ab709d4ec0 | |||
| cc29229f35 | |||
| c5a4d559a2 | |||
| f27521d62b | |||
| 1ce2d059ca | |||
| 741eb709b9 | |||
| 174c6044a7 | |||
| b5a3f6de90 | |||
| a792615340 | |||
| 70a5f99c13 | |||
| 7be3a90c09 | |||
| d31ee86e79 | |||
| fb158f2ce7 | |||
| 21a875e026 | |||
| 46a1de9dd6 | |||
| 1bcc38ea16 | |||
| 27f8de5ad2 | |||
| bfecfe0afc | |||
| 4b25816ae0 | |||
| 438bcb82e1 | |||
| 3f23ac3657 | |||
| 4cbb19defc | |||
| 0e6fe8532d | |||
| f4bdc85dfc | |||
| a6186a67b8 | |||
| fb92972ce5 | |||
| 534a7cea06 | |||
| 6f927daa89 | |||
| 483ace6df0 | |||
| a635699539 | |||
| 01b94366e9 | |||
| 6f97b1027a | |||
| 04bd8c7841 | |||
| a93e58bccf | |||
| 17f75d4a32 | |||
| ef80273480 | |||
| 1ccdbbea90 | |||
| 017f169015 | |||
| 72b422d3a2 | |||
| f7b5cc2a1c | |||
| 8060ded803 | |||
| 6034f07ea5 | |||
| ee20adb021 | |||
| 438f9e56d5 | |||
| 9bbde69f7f | |||
| 1cbf6a8cd1 | |||
| d226aabe01 | |||
| 5ab677c40d | |||
| 1ce8837122 | |||
| 774473eb30 | |||
| cf225de644 | |||
| 20c76cd7cf | |||
| 5fb7094e6f | |||
| fb5babbcc8 | |||
| 9ad215b257 | |||
| ad2b1c9f97 | |||
| 8d7dfaafd0 | |||
| 6eceb045ab | |||
| 253abe8445 | |||
| 0d4f736eae | |||
| 1d91c028db | |||
| 83845e2655 | |||
| 1db62de1bb | |||
| 7784afa12b | |||
| 28d8b334c7 | |||
| 5ff000a9b2 | |||
| 1625caa194 | |||
| c15c0dfdf0 | |||
| bc7a0b753c | |||
| 30dde5cf83 | |||
| 9c56c73838 | |||
| ae1da96fb6 | |||
| b54acbde41 | |||
| 4a06afd0f1 | |||
| f41a9a6114 | |||
| d3e7e22bb8 | |||
| 311fb08676 | |||
| 903ef4d158 | |||
| aad431785c | |||
| 4c53bb0480 | |||
| 6817128a03 | |||
| 18b1d38aa5 | |||
| c4de31c772 | |||
| 68f432b39d | |||
| ce33de79af | |||
| 3885197559 | |||
| f3bb70c74e | |||
| 260ee95697 | |||
| a118d3ed23 | |||
| 07e57e110b | |||
| 659fbc7f84 | |||
| a91d6def40 | |||
| ed376d1b5b | |||
| 085cef4463 | |||
| 5ea9f3e938 | |||
| 833fa0abfb | |||
| bdd3e30284 | |||
| 250c41c561 | |||
| cdb5507896 | |||
| 6992e37ce5 | |||
| 0c5d0049a4 | |||
| d2555ac10f | |||
| 0a823c6e0d | |||
| e7c06f8ddb | |||
| d69ba4ec8b | |||
| 5e1235627d | |||
| 0897f3a2e1 | |||
| 8c77ece5e5 | |||
| eeb0a21275 | |||
| 13858b860a | |||
| b14dceb143 | |||
| 5f2d963a97 | |||
| c62c7d24e2 | |||
| 67a60b9ca1 | |||
| b05576cdbc | |||
| 313e48eb62 | |||
| 3be3f5adf0 | |||
| bf4e297cbb | |||
| 5fa7c684bd | |||
| c96d984869 | |||
| 32bbd1a1af | |||
| a574492123 | |||
| 7e65bfe0b7 | |||
| eaae62cf26 | |||
| 8ebc7dafc4 | |||
| f286091607 | |||
| 845fdd8710 | |||
| 540f986c3c | |||
| 78de46455f | |||
| feb39bfbd9 | |||
| 3251b8985b | |||
| f2723ad674 | |||
| 2f0108d7f2 | |||
| 409a1bb7d4 | |||
| e793bc0797 | |||
| 2f6cf824f2 | |||
| 1130ddd8e9 | |||
| 3f42259441 | |||
| d81d41f756 | |||
| c9e5d925b8 | |||
| ebaaff6083 | |||
| 2431d7e926 | |||
| 08507f5d14 | |||
| 200c5e0a12 | |||
| 29881d7392 | |||
| 931fc8a88d | |||
| 35135ee79c | |||
| abf783d2d3 | |||
| 7a6488d51d | |||
| ea00578cee | |||
| 5705720d0e | |||
| 6faad4bb98 | |||
| fad6bde5a6 | |||
| d39c82989d | |||
| 5260b17c0b | |||
| 6c07e9dda7 | |||
| f13d43a362 | |||
| ef46411796 | |||
| 317dfa1f63 | |||
| 0eb6254083 | |||
| bb7a4c12e9 | |||
| 53f6761dec | |||
| ae15c569c7 | |||
| 9c6f2838ad | |||
| 2da1d6944b | |||
| 1ccff57acd | |||
| 4b69903a94 | |||
| 709f212289 | |||
| b72db34a08 | |||
| c9b4bf0367 | |||
| 9b088aa0ed | |||
| c88e9ccc08 | |||
| 65e95ac4b3 | |||
| a1c44cf95d | |||
| 11b14d01a8 | |||
| 11587250ae | |||
| 5523a3c4a3 | |||
| f3cf1dc896 | |||
| c57e6132a5 | |||
| 7dedf2b426 | |||
| a2048e4127 | |||
| 663cc601d5 | |||
| f211ecb7d7 | |||
| 258fbc4d34 | |||
| 0e4d09f5d4 | |||
| bd167d6c85 | |||
| b5dbdc2a1e | |||
| d62c0cfb2e | |||
| 03f4134190 | |||
| 75ee7aeba9 | |||
| 6399855df7 | |||
| 4675d311fe | |||
| ce4497fab5 | |||
| abb237115f | |||
| 86a288b3a1 | |||
| 13d6852857 | |||
| e41e04a176 | |||
| e49bac3a33 | |||
| 8a00cd3508 | |||
| 2dd331be02 | |||
| 145b31dc2c | |||
| db6e309fdd | |||
| 98e133564f | |||
| d09ae76a82 | |||
| e2a5241a68 | |||
| 8873b83083 | |||
| a786e181bf | |||
| 23c0a47455 | |||
| 5fd8e1a8ad | |||
| bec3fc0b67 | |||
| b85faf7516 | |||
| 72e4353f54 | |||
| 65d53b4c24 | |||
| 1446267360 | |||
| 283ae49a91 | |||
| b62e63ac74 | |||
| 1cc7c539be | |||
| f39556afa0 | |||
| 4fc7a035a2 | |||
| 2ee50e3c4f | |||
| 1fc823b182 | |||
| 764e4a4cc8 | |||
| 8ffc294ba1 | |||
| 12bc82226f | |||
| 924ce3dd30 | |||
| 4daf290fee | |||
| 8a2b6161f1 | |||
| 616a90ef1c | |||
| 3da590acaf | |||
| eb10135314 | |||
| 84da81256a | |||
| 9b93cea12b | |||
| 91e0cfb7c1 | |||
| 358f63994b | |||
| c2c8e27b46 | |||
| 3758ae1c1d | |||
| 7b598dd6a6 | |||
| 1ccc4158d9 | |||
| 01018c0d24 | |||
| 919aa15c5e | |||
| 68181fae49 | |||
| e672f71c12 | |||
| 7a94951af9 | |||
| 2f851cdd7b | |||
| f4419dae80 | |||
| 299054edbf | |||
| 123a5de179 | |||
| 8893906e94 | |||
| 94184fe13b | |||
| f80246c48f | |||
| 61cf79e525 | |||
| d591fda338 | |||
| 8c67abf7ef | |||
| 7572fef50a | |||
| 418f2d58a7 | |||
| b50299534f | |||
| ccdf2157cb | |||
| 97ac3e0401 | |||
| 2265fb70b7 | |||
| a7f1e166d7 | |||
| 8210192b03 | |||
| 349693f9d6 | |||
| 5e83f4caba | |||
| 5271931935 | |||
| a61333d4c5 | |||
| 8118d20dfd | |||
| 29be24efe6 | |||
| 56711aa321 | |||
| 10266c660f | |||
| f29b254d9a | |||
| 63517a2a04 | |||
| 7cf55e2aa1 | |||
| b0644a2b53 | |||
| 8404a8e013 | |||
| 100b74dd63 | |||
| e987ec1e61 | |||
| 27572dbd46 | |||
| 20db6be1d3 | |||
| 20751205e4 | |||
| 7eb1d1bf4d | |||
| f1a7b957e7 | |||
| d1cc6096d2 | |||
| e3ef17507b | |||
| baf025da9a | |||
| 3840a3dfea | |||
| 250a23682c | |||
| a32fe79bdd | |||
| b923c73283 | |||
| e2fbbf1b02 | |||
| 61a3cb645f | |||
| 7b15b3b8a0 | |||
| 05b098ca63 | |||
| d23acd9f01 | |||
| 5518f4b1f9 | |||
| 64b2c56679 | |||
| f2d991b722 | |||
| e07bcf2a91 | |||
| 1898fdf33f | |||
| 53d38ea0c6 | |||
| b2f93bfda9 | |||
| f9c25de680 | |||
| 5b7120a401 | |||
| 967670bfc2 | |||
| 80baf85a2c | |||
| 9341dc474b | |||
| a507032236 | |||
| a0397e449a | |||
| 6261cbd07c | |||
| c3e473cc69 | |||
| 3d54a5e3c5 | |||
| ac485c30e8 | |||
| a267cabd6a | |||
| 12f435b82a | |||
| 5eb7fb2a58 | |||
| 10a63c785e | |||
| 905dd5568a | |||
| d3d1310986 | |||
| 6e2b37cc13 | |||
| cb7d3081a0 | |||
| fa304e5a55 | |||
| 55bdfef347 | |||
| 3d0805dd83 | |||
| 67a561a1ee | |||
| 3c29076e4c | |||
| 9f0ce14b94 | |||
| 97f104015b | |||
| f46e328545 | |||
| b677d8b08b | |||
| c4698e118d | |||
| 5810cdd7b1 | |||
| b6f778e725 | |||
| 9f12a35b7f | |||
| 0d789c33cd | |||
| 1aef7e450e | |||
| e725fb08e2 | |||
| 7d4b8e6557 | |||
| 51cf4a9acb | |||
| 57aaa6309d | |||
| e4b7a1eb62 | |||
| 3eec63a447 | |||
| f36fec0bf9 | |||
| ecefcee1d0 | |||
| 11dc675649 | |||
| 80ee8c82cd | |||
| 425e60e9f4 | |||
| 3993029466 | |||
| 068d6e135d | |||
| 9fc25349ea | |||
| 529527020b | |||
| 5b0e57eb61 | |||
| 19548016d9 | |||
| eccbee0fa1 | |||
| ad8f44af8d | |||
| edeae89b2e | |||
| f3126cba29 | |||
| b4036665c8 |
@@ -0,0 +1,65 @@
|
||||
# VCS
|
||||
.git
|
||||
.gitignore
|
||||
.gitattributes
|
||||
|
||||
# CI / project meta (image doesn't need these)
|
||||
.github/
|
||||
docs/
|
||||
CONTRIBUTING.md
|
||||
README.zh-CN.md
|
||||
LICENSE
|
||||
|
||||
# Editor / tooling state
|
||||
.vscode/
|
||||
.idea/
|
||||
.cursor/
|
||||
.codex/
|
||||
.claude/
|
||||
.agents/
|
||||
.cursorrules
|
||||
.ruff_cache/
|
||||
|
||||
# Python build artifacts and caches
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
*.egg
|
||||
build/
|
||||
dist/
|
||||
.venv/
|
||||
venv/
|
||||
.pytest_cache/
|
||||
.ruff_cache/
|
||||
.coverage
|
||||
|
||||
# Tests aren't needed at runtime
|
||||
tests/
|
||||
|
||||
# Notebooks
|
||||
*.ipynb
|
||||
.ipynb_checkpoints/
|
||||
|
||||
# Local runtime data (must never leak into the image)
|
||||
.env
|
||||
.env.*
|
||||
!.env.example
|
||||
runs/
|
||||
workspace/
|
||||
skills/
|
||||
memory/
|
||||
memories/
|
||||
media/
|
||||
conversation_history/
|
||||
.deno_cache/
|
||||
.langgraph_api/
|
||||
large_tool_results/
|
||||
*.log
|
||||
botpy.log
|
||||
|
||||
# Docker outputs themselves
|
||||
Dockerfile.*
|
||||
docker-compose*.override.yml
|
||||
|
||||
# OS
|
||||
.DS_Store
|
||||
@@ -1,27 +1,41 @@
|
||||
# EvoScientist CLI environment variables
|
||||
# The preferred configuration flow is `evosci onboard`, which writes
|
||||
# ~/.evoscientist/config/settings.yaml. Environment variables can override it.
|
||||
# EvoScientist — cp .env.example .env && fill in your keys
|
||||
#
|
||||
# LLM providers, models, and API keys are managed exclusively through the
|
||||
# Model Registry (WebUI 大模型配置 / Config API); no provider credential is
|
||||
# read from environment variables. See
|
||||
# docs/unified-model-configuration-architecture.md.
|
||||
|
||||
# Optional application directories
|
||||
# EVOSCIENTIST_HOME=~/.evoscientist
|
||||
# EVOSCIENTIST_DATA_ROOT=~/.evoscientist/data
|
||||
# Web search (optional)
|
||||
TAVILY_API_KEY= # app.tavily.com
|
||||
|
||||
# Logging
|
||||
EVOSCIENTIST_LOG_LEVEL=INFO
|
||||
# EVOSCIENTIST_LOG_DIR=~/.evoscientist/data/logs
|
||||
EVOSCIENTIST_LOG_RETENTION_DAYS=30
|
||||
# WebUI conversation workspace policy. EVOSCIENTIST_WORKSPACE_DIR is the
|
||||
# deployment root, not a per-conversation directory. In isolated modes each
|
||||
# conversation is stored under <root>/.evoscientist/conversations/<scope-id>/.
|
||||
#
|
||||
# EVOSCIENTIST_WORKSPACE_ISOLATION accepts exactly:
|
||||
# - legacy: all WebUI conversations share the deployment root. Compatibility
|
||||
# rollback only; files are visible to every conversation using this deployment.
|
||||
# - optional: default. New WebUI conversations receive isolated scope folders;
|
||||
# missing Registry/token/scope fails the request instead of silently sharing.
|
||||
# - required: isolated scopes plus strict runtime validation. It requires a
|
||||
# completed cutover and a verified OCI executor; no legacy fallback exists.
|
||||
#
|
||||
# This is a deployment-startup security setting. Change it only during a
|
||||
# maintenance window, restart backend and WebUI afterwards, and never use it to
|
||||
# convert an existing conversation between shared and isolated directories.
|
||||
EVOSCIENTIST_WORKSPACE_DIR=
|
||||
EVOSCIENTIST_WORKSPACE_ISOLATION=optional
|
||||
# Required mode supports only a single-host Registry topology in v1.
|
||||
EVOSCIENTIST_SCOPE_REGISTRY_TOPOLOGY=single-host
|
||||
# Required mode: use a pinned image digest, preserve single-host topology, and
|
||||
# keep the Code Interpreter disabled unless its scoped implementation is enabled.
|
||||
# Do not put EVOSCIENTIST_BACKEND_SERVICE_TOKEN here for a same-host `EvoSci
|
||||
# deploy`: it is generated and passed privately at startup.
|
||||
# EVOSCIENTIST_WORKSPACE_ISOLATION=required
|
||||
# EVOSCIENTIST_STRICT_EXECUTOR=oci
|
||||
# EVOSCIENTIST_STRICT_EXECUTOR_IMAGE=registry.example/evoscientist-runtime@sha256:replace-with-verified-digest
|
||||
# EVOSCIENTIST_STRICT_CODE_INTERPRETER=disabled
|
||||
|
||||
# Optional PostgreSQL checkpoint storage. Without this value the CLI uses an
|
||||
# in-memory checkpointer for the current process.
|
||||
# EVOSCIENTIST_SESSION_DB_URL=postgresql://user:password@localhost:5432/evoscientist
|
||||
|
||||
# Model provider examples. Provider-specific settings can also be configured
|
||||
# through `evosci onboard`.
|
||||
# OPENAI_API_KEY=sk-...
|
||||
# OPENAI_BASE_URL=https://api.openai.com/v1
|
||||
# ANTHROPIC_API_KEY=sk-ant-...
|
||||
# GOOGLE_API_KEY=...
|
||||
# TAVILY_API_KEY=tvly-...
|
||||
|
||||
# Optional default model
|
||||
# DEFAULT_MODEL=openai/gpt-5.4
|
||||
# Conversation workspace isolation retention defaults (used by workspace_maintenance.py).
|
||||
EVOSCIENTIST_DRAFT_WORKSPACE_TTL_HOURS=24
|
||||
EVOSCIENTIST_WORKSPACE_TRASH_RETENTION_DAYS=7
|
||||
|
||||
|
After Width: | Height: | Size: 234 KiB |
@@ -5,5 +5,5 @@
|
||||
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
|
||||
<text x="85" y="22" text-anchor="middle"
|
||||
font-family="Inter, -apple-system, system-ui, sans-serif"
|
||||
font-size="13" font-weight="700" fill="#ffffff">v0.0.7</text>
|
||||
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 555 B After Width: | Height: | Size: 555 B |
@@ -5,5 +5,5 @@
|
||||
<rect x="54" y="5" width="62" height="24" rx="6" fill="#2563eb"/>
|
||||
<text x="85" y="22" text-anchor="middle"
|
||||
font-family="Inter, -apple-system, system-ui, sans-serif"
|
||||
font-size="13" font-weight="700" fill="#ffffff">v0.0.7</text>
|
||||
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 555 B After Width: | Height: | Size: 555 B |
|
Before Width: | Height: | Size: 654 KiB |
|
After Width: | Height: | Size: 213 KiB |
|
Before Width: | Height: | Size: 428 KiB After Width: | Height: | Size: 287 KiB |
@@ -0,0 +1,20 @@
|
||||
version: 2
|
||||
updates:
|
||||
# Base images in Dockerfile (BASE_IMAGE / NODE_IMAGE ARG defaults).
|
||||
# Dependabot reads `FROM`, `COPY --from=`, and ARG-bound base refs, and
|
||||
# bumps both the @sha256 digest and the trailing # vX.Y.Z comment.
|
||||
- package-ecosystem: "docker"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
open-pull-requests-limit: 5
|
||||
commit-message:
|
||||
prefix: "chore(docker)"
|
||||
labels:
|
||||
- "dependencies"
|
||||
- "docker"
|
||||
# Single PR per cadence rather than one per image — keeps reviewer load low
|
||||
# and lets us validate trixie/uv/node bumps as a coherent set.
|
||||
groups:
|
||||
base-images:
|
||||
patterns: ["*"]
|
||||
@@ -0,0 +1,67 @@
|
||||
name: Docker
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: ["main"]
|
||||
tags: ["v*"]
|
||||
pull_request:
|
||||
paths:
|
||||
- "Dockerfile"
|
||||
- ".dockerignore"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
- "EvoScientist/**"
|
||||
- ".github/workflows/docker.yml"
|
||||
workflow_dispatch:
|
||||
|
||||
concurrency:
|
||||
group: docker-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
env:
|
||||
REGISTRY: ghcr.io
|
||||
IMAGE_NAME: ${{ github.repository }}
|
||||
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 45
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
steps:
|
||||
- uses: actions/checkout@93cb6efe18208431cddfb8368fd83d5badbf9bfd # v5.0.1
|
||||
|
||||
- uses: docker/setup-qemu-action@c7c53464625b32c7a7e944ae62b3e17d2b600130 # v3.7.0
|
||||
- uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
||||
|
||||
- name: Log in to ${{ env.REGISTRY }}
|
||||
if: github.event_name != 'pull_request'
|
||||
uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3.7.0
|
||||
with:
|
||||
registry: ${{ env.REGISTRY }}
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Extract image metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@c299e40c65443455700f0fdfc63efafe5b349051 # v5.10.0
|
||||
with:
|
||||
images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}
|
||||
tags: |
|
||||
type=ref,event=branch
|
||||
type=ref,event=pr
|
||||
type=semver,pattern={{version}}
|
||||
type=semver,pattern={{major}}.{{minor}}
|
||||
type=raw,value=latest,enable={{is_default_branch}}
|
||||
|
||||
- name: Build and push
|
||||
uses: docker/build-push-action@10e90e3645eae34f1e60eeb005ba3a3d33f178e8 # v6.19.2
|
||||
with:
|
||||
context: .
|
||||
platforms: linux/amd64,linux/arm64
|
||||
push: ${{ github.event_name != 'pull_request' }}
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
@@ -7,13 +7,25 @@ on:
|
||||
|
||||
jobs:
|
||||
pytest:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
# ``fail-fast: false`` so a single failing (os, python-version) cell
|
||||
# doesn't cancel the rest of the matrix. Useful while the Windows
|
||||
# leg is being brought up — we want to see all four cell results
|
||||
# in one CI run instead of playing whack-a-mole one failure at a
|
||||
# time. See #207.
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-latest, macos-latest, windows-latest]
|
||||
python-version: ["3.11", "3.12"]
|
||||
runs-on: ${{ matrix.os }}
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
- name: Check out shared usage fixtures
|
||||
uses: actions/checkout@v5
|
||||
with:
|
||||
repository: EvoScientist/EvoScientist-WebUI
|
||||
path: EvoScientist-WebUI
|
||||
- uses: astral-sh/setup-uv@v6
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
@@ -22,3 +34,23 @@ jobs:
|
||||
run: uv sync --dev
|
||||
- name: Run pytest
|
||||
run: uv run pytest -v --timeout=30
|
||||
env:
|
||||
EVOSCIENTIST_USAGE_FIXTURES: ${{ github.workspace }}/EvoScientist-WebUI/docs/schemas/fixtures
|
||||
|
||||
usage-spool-benchmark:
|
||||
timeout-minutes: 15
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-latest, macos-latest, windows-latest]
|
||||
runs-on: ${{ matrix.os }}
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
- uses: astral-sh/setup-uv@v6
|
||||
with:
|
||||
python-version: "3.11"
|
||||
cache-dependency-glob: "**/pyproject.toml"
|
||||
- name: Install dependencies
|
||||
run: uv sync --dev
|
||||
- name: Verify durable spool latency
|
||||
run: uv run python scripts/benchmark_usage_spool.py
|
||||
|
||||
@@ -9,13 +9,13 @@ dist/
|
||||
build/
|
||||
*.egg
|
||||
*.pytest_cache/
|
||||
.benchmarks/
|
||||
.coverage
|
||||
.ipynb_checkpoints/
|
||||
|
||||
# Environment
|
||||
.env
|
||||
.env.*
|
||||
.env_*
|
||||
!.env.example
|
||||
.venv/
|
||||
venv/
|
||||
@@ -36,7 +36,9 @@ bridge/package-lock.json
|
||||
.langgraph_api/
|
||||
workspace/
|
||||
skills/
|
||||
!EvoScientist/skills/
|
||||
memory/
|
||||
!EvoScientist/memory/
|
||||
media/
|
||||
conversation_history/
|
||||
.deno_cache/
|
||||
@@ -46,22 +48,8 @@ conversation_history/
|
||||
*meals/
|
||||
botpy.log
|
||||
large_tool_results/
|
||||
runs/
|
||||
|
||||
# Docker runtime data
|
||||
docker/data/
|
||||
|
||||
# Project-level config data
|
||||
.data/
|
||||
|
||||
# Sensitive / credentials (never commit)
|
||||
postgresql:*
|
||||
_s3_backup/
|
||||
|
||||
# Local scratch / tooling data
|
||||
.superpowers/
|
||||
.test-home/
|
||||
tmp/
|
||||
|
||||
# Root-level debug scratch scripts (proper tests live in tests/)
|
||||
/test_*.py
|
||||
/research_lookup_temp.py
|
||||
# local runtime artifacts (scope tokens, control DBs)
|
||||
.evoscientist/
|
||||
.release-state.json
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
repos:
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
# Ruff version.
|
||||
rev: v0.9.9
|
||||
rev: v0.15.17
|
||||
hooks:
|
||||
# Run the linter.
|
||||
- id: ruff
|
||||
- id: ruff-check
|
||||
args: [ --fix ]
|
||||
# Run the formatter.
|
||||
- id: ruff-format
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
# Task 1 报告 — RegistryV4 schema + SQLite 存储层 + 统一错误码
|
||||
|
||||
## 实现摘要
|
||||
|
||||
按简报与设计文档 v1.1.0(4.2、4.3、5.1、8.2、9.5 节)在仓库内新建独立子包
|
||||
`EvoScientist/model_registry/`,未改动任何现有文件(顶层 `EvoScientist/__init__.py`
|
||||
采用惰性导出,无需修改)。
|
||||
|
||||
1. **`schemas.py` — 唯一 RegistryV4 schema(Pydantic v2)**
|
||||
- `ProviderId`/`ModelKey`/`CredentialId`/`AdapterId` 共用锚定全匹配模式
|
||||
`\A[a-z0-9][a-z0-9._-]{0,63}\z`(pydantic 的 pattern 约束默认是子串搜索,必须锚定)。
|
||||
- `upstream_model_id` 仅限长 1–300、保留大小写;`ValidatedEndpoint` 限长 2048 且要求
|
||||
http(s) scheme(EndpointPolicy 属后续任务)。
|
||||
- `RegistryV4 { version: Literal[4], revision: PositiveInt, state, defaults, providers }`;
|
||||
schema 级校验:Provider ID 唯一、Provider 内模型 key 唯一、defaults 引用必须存在、
|
||||
`active` 状态下 primary 非空且引用已启用模型(auxiliary 同理)。
|
||||
- `ProviderConfig.runtime`:`timeout_seconds` [10,600](缺省 120)、`max_retries` [0,5]
|
||||
(缺省 2)、`default_temperature` [0,2]|null、`default_top_p` (0,1]|null、
|
||||
`default_reasoning_effort` 缺省 `auto`。
|
||||
- `ModelConfig.runtime`:`limit_mode` combined 要求 `context_window_tokens`、input_only
|
||||
要求 `max_input_tokens`(model_validator);`min_effective_input_tokens` 缺省 4096、
|
||||
下限 1024;三个 `fixed_*_reserve_tokens` 非负;temperature/top_p/reasoning_effort 与
|
||||
`declared_capabilities` 按 4.3 定义。
|
||||
- `AuthConfig`:`mode=none` 时 `credential_id` 必须为 null(model_validator)。
|
||||
- 6.2/6.4/9.1 结构:`AdapterParameterSpec`(含 `connection` 可选块,对应 6.2 示例
|
||||
`chat_model`/`model_field`/`base_url_field`)、`AuthSpec`、`ParameterRule`、
|
||||
`ModelAvailability`(`VerificationInfo`)、`ResolvedModelConfig`
|
||||
(`auth_ref={mode, credential_id?, credential_revision?}`,无任何 secret 字段)、
|
||||
`CredentialStatus`、`CredentialWrite`(9.2 的 `operation: replace`)。
|
||||
2. **`errors.py` — 统一错误码与载荷**
|
||||
- 23 个稳定错误码常量 + `ERROR_HTTP_STATUS` 映射,覆盖 9.5 总表全部 17 组
|
||||
(409×6、401×1、404×1、422×15)。
|
||||
- `ErrorPayload {code, message, details:[{path, code}], request_id}` 与 9.5 结构一致;
|
||||
`ModelRegistryError` 携带 code/`http_status`/`payload()`,未知 code 直接拒绝。
|
||||
3. **`store.py` — `ModelRuntimeStore`**
|
||||
- 数据库 `<config_dir>/model-runtime.sqlite3`(默认 `~/.config/evoscientist`,可注入);
|
||||
目录 0700、文件 0600、WAL、外键、`busy_timeout=30000`。
|
||||
- 手写 DDL(`CREATE TABLE IF NOT EXISTS`,无 alembic):`registry_state`(单行)、
|
||||
`credential_pointers`、`credential_versions`((credential_id, revision) 主键)、
|
||||
`model_verifications`(五元组主键,upsert 只留最近一次)、`run_runtime_snapshots`
|
||||
(含部分唯一索引 `UNIQUE(deployment_id, thread_id, run_request_id)
|
||||
WHERE status IN ('prepared','bound')`)、`delegation_jtis`。
|
||||
- `load_registry()` 无行时返回 bootstrap/revision=1 空 RegistryV4;
|
||||
`save_registry(expected_revision=, registry=, credential_writes=)` 在
|
||||
`BEGIN IMMEDIATE` 事务内校验 revision(不符抛 `REGISTRY_REVISION_CONFLICT`)、
|
||||
写入不可变凭据版本、registry revision+1;首次同时具备已启用模型+有效 primary+已配置
|
||||
凭据(或 `mode=none`)时原子转为 `active`;任一失败整体回滚。每次保存强制执行
|
||||
9.2 第 7 条(defaults 必须引用已启用模型,违反抛 `MODEL_DISABLED`)。
|
||||
- 凭据:`write_credential_version`(递增 revision;重写相同当前密钥幂等返回原
|
||||
revision)、`resolve_credential`(不存在/已销毁抛
|
||||
`RUN_CREDENTIAL_REVISION_UNAVAILABLE`)、`retire_credential_version`、
|
||||
`credential_status`(hint 末 4 位 `...abcd`,短于 4 字符的密钥 hint 为 null 绝不泄露;
|
||||
绝不返回明文)。
|
||||
- 快照:`insert_run_snapshot`/`set_run_snapshot_status`/`get_run_snapshot`
|
||||
(状态机 prepared|bound|expired|aborted,部分唯一索引行为由测试覆盖)。
|
||||
- `check_shared_storage()`:探测 `BEGIN IMMEDIATE` 写锁能力,失败抛
|
||||
`SharedStorageError`(多节点不共享持久卷时启动失败)。
|
||||
4. **`hashing.py` — `configuration_hash(provider, model)`**
|
||||
- 覆盖 adapter、base_url、upstream_model_id、Provider 与 Model 全部运行参数(含声明
|
||||
能力与限制),`json.dumps(sort_keys=True, separators=(",", ":"))` 规范化后 SHA-256。
|
||||
|
||||
## 文件清单
|
||||
|
||||
新增(无修改既有文件):
|
||||
- `EvoScientist/model_registry/__init__.py`
|
||||
- `EvoScientist/model_registry/schemas.py`
|
||||
- `EvoScientist/model_registry/errors.py`
|
||||
- `EvoScientist/model_registry/store.py`
|
||||
- `EvoScientist/model_registry/hashing.py`
|
||||
- `tests/test_model_registry_schemas.py`
|
||||
- `tests/test_model_registry_store.py`
|
||||
- `.superpowers/sdd/briefs/task-1-report.md`(本文件)
|
||||
|
||||
## 测试命令与输出
|
||||
|
||||
TDD 流程:先写两个测试文件并确认失败(`ModuleNotFoundError: No module named
|
||||
'EvoScientist.model_registry'`),再实现。
|
||||
|
||||
```
|
||||
$ .venv/bin/python -m pytest tests/test_model_registry_schemas.py tests/test_model_registry_store.py -x -q
|
||||
........................................................................ [ 80%]
|
||||
................. [100%]
|
||||
89 passed in 0.22s
|
||||
```
|
||||
|
||||
全量回归(无既有失败,无回归):
|
||||
|
||||
```
|
||||
$ .venv/bin/python -m pytest tests/ -x -q
|
||||
........sssss....................................... [100%]
|
||||
2922 passed, 10 skipped, 1 warning in 73.00s
|
||||
```
|
||||
|
||||
(warning 为 `test_langgraph_dev_http.py` 的 StarletteDeprecationWarning,既有、与本任务无关。)
|
||||
|
||||
lint 与格式:
|
||||
|
||||
```
|
||||
$ .venv/bin/ruff check EvoScientist/model_registry tests/test_model_registry_schemas.py tests/test_model_registry_store.py
|
||||
All checks passed!
|
||||
$ .venv/bin/ruff format --check ... # 已格式化
|
||||
```
|
||||
|
||||
## 自我审查发现(已处理)
|
||||
|
||||
1. **测试副作用污染真实配置目录**:初版 `test_default_config_dir` 未注入路径,运行时在
|
||||
真实 `~/.config/evoscientist/` 创建了空的 `model-runtime.sqlite3`。已确认该库所有表
|
||||
为空(确为测试副产物)后删除(含 -wal/-shm),并把测试改为 monkeypatch
|
||||
`store.DEFAULT_CONFIG_DIR` 到 tmp_path,此后测试不再触碰真实 home。
|
||||
2. **共享 `Field()` 实例**:初版 `_TEMPERATURE`/`_TOP_P` 在三个模型间复用同一 FieldInfo,
|
||||
已改为各字段独立 `Field(...)`,规避 pydantic 共享元数据的潜在风险。
|
||||
3. **docstring 混入中文**:`save_registry` 一处 docstring 误用中文,已改为英文以符合
|
||||
仓库注释惯例。
|
||||
4. **并发 CAS 断言**:`sorted()` 大小写排序导致误报,改为不区分大小写排序。
|
||||
5. ruff 修复:`datetime.UTC` 别名、导入排序、`pytest.raises` 增加 `match=`。
|
||||
|
||||
## 遗留疑虑
|
||||
|
||||
1. **`connection` 块为可选**:6.2 正文的"至少包含"清单未列 `connection`,但 glm-5.2
|
||||
示例契约包含它。schema 将其建模为可选字段,Task 2 落地五种 Adapter 契约时若确认
|
||||
每个契约都有 connection,可考虑收紧为必填。
|
||||
2. **激活就绪判定中的"满足 AuthSpec"**:4.3 要求激活时认证状态满足 AuthSpec,但
|
||||
Adapter 契约属 Task 2。当前 store 仅做存储层判定(`mode=none` 或凭据已配置);
|
||||
AuthSpec 级别校验(如 credential_kind 匹配)需在 Task 2/3 的 API 层补充。
|
||||
3. **幂等语义解释**:简报称 `write_credential_version` 为"幂等 prepare",实现为"重写
|
||||
与当前版本完全相同的密钥时返回现有 revision";不同密钥轮换仍产生新 revision。若
|
||||
后续任务对幂等键有不同约定(如客户端提供 request id),需再对齐。
|
||||
4. **`save_registry` 参数为 keyword-only**:与简报签名
|
||||
`save_registry(expected_revision, registry, credential_writes=[])` 语义一致,仅调用
|
||||
形式略异。
|
||||
5. `MODEL_DISABLED` 用于 defaults 引用未启用模型的保存错误(9.2 第 7 条属 422 校验,
|
||||
总表无更贴切码);引用不存在模型由 schema 层先行拒绝,store 内同名分支仅作防御。
|
||||
@@ -0,0 +1,89 @@
|
||||
# Task 4 报告 — ModelRegistryResolver + 运行快照服务
|
||||
|
||||
- 状态:DONE
|
||||
- 分支:`feature/unified-model-config`
|
||||
- Commit:`b1233d4` `feat(model-registry): add ModelRegistryResolver and run snapshot service`
|
||||
- 规格来源:`docs/unified-model-configuration-architecture.md` v1.1.1,章节 4.3 / 5.2 / 6.1 / 6.4 / 6.5 / 8.1 / 8.2
|
||||
|
||||
## 交付物
|
||||
|
||||
### 新建 `EvoScientist/model_registry/resolver.py`
|
||||
|
||||
`ModelRegistryResolver(store, *, specs=None)`:
|
||||
|
||||
- `resolve(model_ref, role="primary", *, registry=None) -> ResolvedModelConfig`,校验顺序:
|
||||
1. Provider 存在(否则 404 `MODEL_NOT_FOUND`)且 enabled(否则 422 `MODEL_NOT_AVAILABLE`);
|
||||
2. 模型存在(`MODEL_NOT_FOUND`)且 enabled(否则 422 `MODEL_DISABLED`);
|
||||
3. `find_adapter_spec` 匹配契约:未开放 Adapter 抛 `ADAPTER_NOT_SUPPORTED`,无匹配契约抛 `MODEL_NOT_AVAILABLE`(只可保持 configured);
|
||||
4. `resolve_parameters` 重放保存期契约校验(AuthSpec、能力、参数范围/互斥);
|
||||
5. 凭据:`credential_required` 时读 `credential_pointers.current_revision` 写入 `auth_ref`(指针缺失 → `CREDENTIAL_NOT_CONFIGURED`),**从不读 `secret_value`**;`mode=none` 冻结 `{mode: none}`;
|
||||
6. 当前验证五元组(`configuration_hash` + 当前凭据 revision + 当前 `adapter_spec_revision`)必须存在且 `passed`,否则 `MODEL_NOT_AVAILABLE`;
|
||||
7. `limits_status=confirmed`,否则 `MODEL_LIMITS_UNCONFIRMED`;
|
||||
8. 预算(6.5):`combined` → `context_window - max_output`;`input_only` → `max_input`;四种工具/附件组合扣除固定预留后均须 ≥ `min_effective_input_tokens`,否则 `CONTEXT_BUDGET_UNSATISFIABLE`。冻结 `resolved_input_limit`、三个固定预留与 base `message_budget`(仅扣系统预留;逐次调用的预算由 MessageBudgetMiddleware 按 6.5 公式重算,Task 6 接线);
|
||||
9. 能力按唯一规则 `protocol AND declared AND verified` 计算并冻结。
|
||||
- `resolve_for_test(model_ref)`:仅放宽"模型必须 enabled",其余校验(含 Provider enabled、凭据、验证记录、限制、预算)全部执行(9.4)。
|
||||
- `compute_availability(registry, verifications, *, credential_revisions=None) -> list[ModelAvailability]`:4.3 判定顺序 `unavailable → enabled → verified → verification_failed → verification_stale → configured`,首个命中生效;`verification_stale` 优先于 `configured`;`selectable` 仅 `enabled` 为 true。`verification` 报告当前记录的 passed/failed 或最新旧记录的 stale;`effective_capabilities` 仅在当前记录 passed 时非全 false。reason_code 取值:`PROVIDER_DISABLED` / `NO_ADAPTER_CONTRACT` / `MODEL_DISABLED` / 记录的错误码 / `VERIFICATION_STALE`。
|
||||
- 6.1 角色映射实现为 `snapshots.config_for_role`(映射对象是快照而非 Registry,故放在快照模块):`primary → snapshot.primary`;`auxiliary`/`summary`/`tool_selector → snapshot.auxiliary ?? snapshot.primary`;未知角色 `ValueError`。Resolver 不猜测其他映射。
|
||||
|
||||
### 新建 `EvoScientist/model_registry/snapshots.py`
|
||||
|
||||
`SnapshotService(store, resolver)`(Task 5 HTTP API 与 Task 7 本地入口共用):
|
||||
|
||||
- `create(SnapshotCreateRequest)`:bootstrap → 422 `MODEL_REGISTRY_NOT_READY`;`primary=null`(inherit)解析 Registry `defaults.primary`,`auxiliary=null` 解析 `defaults.auxiliary`;冻结两个完整 `ResolvedModelConfig`(含 `adapter_spec_revision`、固定预留、能力、`auth_ref` 凭据版本)与 `registry_revision`、`model_selection_revision` 写入 `payload_json`,绝无 secret。`selection_hash` = 解析前 `{primary, auxiliary}`(inherit 以 null 参与)固定字段序 JSON 的 SHA-256;`model_selection_revision` 仅审计、不参与哈希。
|
||||
- 幂等:同一三元组哈希相同返回原快照(`created=False`,200 语义),不同抛 409 `RUN_REQUEST_CONFLICT`;并发创建撞部分唯一索引时回退到同一幂等比较。`expired`/`aborted` 不占三元组,同三元组可重建(`created=True`,201)。
|
||||
- `bind(snapshot_id, langgraph_run_id)`:`prepared→bound` 一次(条件 UPDATE 保证原子),并把 `expires_at` 延长到 +24h;相同 run id 重复 bind 幂等成功,不同值抛 `SNAPSHOT_ALREADY_BOUND`;终态抛 `SNAPSHOT_EXPIRED`;不存在抛 `SNAPSHOT_NOT_FOUND`。
|
||||
- `abort(snapshot_id)`:仅 `prepared→aborted`;已 bound 抛 `SNAPSHOT_ALREADY_BOUND`;expired 抛 `SNAPSHOT_EXPIRED`;重复 abort 幂等成功。
|
||||
- `get(snapshot_id, *, deployment_id, thread_id)`:绑定关系不匹配按 `SNAPSHOT_NOT_FOUND` 失败(不跨线程/部署泄露存在性);终态抛 `SNAPSHOT_EXPIRED`;读取时对两个冻结配置逐一重校验 `adapter_spec_revision` 仍存在(经 Resolver 的 spec 列表),已移除抛 `ADAPTER_NOT_SUPPORTED`,绝不静默替换。
|
||||
- `cleanup_expired(now)`:到期 prepared(创建时 TTL **15 分钟**)与 bound(bind 时 **+24 小时**)置为终态 `expired`,返回迁移的 ID。
|
||||
- `resolve_snapshot_credential(snapshot, role) -> str`:经 6.1 角色映射取冻结 `auth_ref`,按冻结 `credential_revision` 每次从凭据存储解析(**无进程内密钥缓存**);版本销毁抛 `RUN_CREDENTIAL_REVISION_UNAVAILABLE`;`mode=none` 返回 `""`。
|
||||
- `public_snapshot_view(snapshot)`:仅 8.2 示例字段(`snapshot_id`、`registry_revision`、primary/auxiliary 的 `provider_id`/`model_key`/`adapter_spec_revision`/runtime 五项),无 base_url、无 secret。
|
||||
|
||||
### Task 1-3 文件的增补(均为纯新增,未改动任何既有对外行为)
|
||||
|
||||
- `errors.py`:新增 `SNAPSHOT_NOT_FOUND`(404)。9.5 承认"404 资源不存在"类别但总表只有 `MODEL_NOT_FOUND`(语义为 ModelRef);快照缺失需要独立稳定码,属对总表的增补,已同步更新 Task 1 的 `ALL_ERROR_CODES` 测试(仅加一行)。
|
||||
- `store.py`:新增 `current_credential_revision`、`list_model_verifications`(可用性判定的输入)、`find_active_run_snapshot`(三元组幂等查询)、`bind_run_snapshot`(`WHERE status='prepared'` 条件绑定 + 延长 expires_at)、`expire_due_run_snapshots`;行→dict 映射提取为共享私有helper,既有方法签名不变。
|
||||
- `__init__.py`:导出新符号。
|
||||
- **未删除** `EvoScientist/llm/runtime_snapshots.py`(Task 6 切换)。
|
||||
|
||||
### Task 3 评审接线要求的遵守
|
||||
|
||||
本任务不构造任何 HTTP client / ChatModel:`resolve`/`resolve_for_test` 只产出 `ResolvedModelConfig`,`SnapshotService` 只做冻结与读取。因此 `build_chat_model` 双传 sync+async client、ollama `max_retries` 经 client builder `retries=` 执行这两条接线要求在本任务无适用点,也未被绕过;Task 5/7 构造 client 时仍须遵守(`build_safe_http_client(policy, retries=resolved.client_options.max_retries)` 双传)。
|
||||
|
||||
## 测试(TDD)
|
||||
|
||||
先写 `tests/test_resolver.py`(37 例)与 `tests/test_snapshots.py`(34 例)并确认红灯(模块不存在),再实现至全绿。覆盖简报全部要求:
|
||||
|
||||
- resolve 正/反例:完整冻结断言(含预算数值 1048576-32768-4096)、Provider disabled→`MODEL_NOT_AVAILABLE`、模型 disabled→`MODEL_DISABLED`、无匹配契约→不可用、验证缺失/失败/五元组不一致(配置哈希、凭据轮换)→不可用、`MODEL_LIMITS_UNCONFIRMED`、`CONTEXT_BUDGET_UNSATISFIABLE`、四种角色戳记、未开放 Adapter、`AUTH_MODE_UNSUPPORTED`、mode=none 无凭据解析、resolved config 不含 secret;
|
||||
- resolve_for_test:放宽模型 enabled、其余校验不放宽;
|
||||
- ModelAvailability 六态判定顺序、stale 优先 configured、凭据轮换致 stale、selectable 仅 enabled、全 Provider 全模型覆盖;
|
||||
- 快照:inherit 解析 defaults(含 auxiliary default null 两条路径)、显式选择冻结双角色、bootstrap 422、selection_hash 解析前语义(显式等于默认仍不同哈希)、revision 不参与哈希、幂等 200/冲突 409、expired/aborted 后重建 201、payload 冻结凭据版本且无 secret、prepared TTL 15min 断言;
|
||||
- bind 一次/同 id 幂等/异 id 409/expired/aborted/未知 404;abort 规则全集;get 绑定校验(跨线程/跨部署拒绝)、expired 拒绝、spec_revision 移除后读取 `ADAPTER_NOT_SUPPORTED`;cleanup 到期迁移;bind 后 +24h 断言;
|
||||
- 凭据:冻结 revision 解析、轮换后旧版本仍可用、销毁后 `RUN_CREDENTIAL_REVISION_UNAVAILABLE`、新 store 实例(无进程缓存)仍可解析、辅助角色解析到 auxiliary 凭据、mode=none 返回空串;
|
||||
- 6.1 角色映射(auxiliary 冻结/缺省两路径、未知角色)与 `public_snapshot_view` 形状(无 secret、无 base_url)。
|
||||
|
||||
## 验证结果
|
||||
|
||||
- `.venv/bin/python -m pytest tests/test_resolver.py tests/test_snapshots.py -x -q` → **71 passed**
|
||||
- `.venv/bin/python -m pytest tests/ -x -q` → **3131 passed, 10 skipped**(无回归;期间修复一处:Task 1 错误码表测试因新增 `SNAPSHOT_NOT_FOUND` 需增补一行期望值)
|
||||
- `ruff check` 与 `ruff format --check`(model_registry 全包 + 涉及测试)→ 全净
|
||||
|
||||
## 疑虑 / 后续注意
|
||||
|
||||
1. `SNAPSHOT_NOT_FOUND` 是对设计文档 9.5 错误码总表的增补(404 类别文档已承认,但总表未列快照缺失码);Task 5 HTTP 层应直接复用。
|
||||
2. 终态(expired/aborted)快照的 bind/get 统一抛 `SNAPSHOT_EXPIRED`;文档只明文规定 expired 的情形,aborted 按同一终态语义处理。
|
||||
3. 冻结的 `budget.message_budget` 取 base(仅扣系统预留);逐次调用的 has_tools/has_attachments 重算属 Task 6 的 MessageBudgetMiddleware。
|
||||
4. `resolve` 支持 `registry=` 参数供 `create` 传入同一份 Registry,保证快照 `registry_revision` 与解析所用文档一致。
|
||||
|
||||
## 评审修复(2026-07-21,commit 见下)
|
||||
|
||||
1. **Important:`abort` read-then-write 竞态**。原实现先读后写且 `set_run_snapshot_status` 为无条件 UPDATE,并发 bind 在两次调用间提交时会把 bound 改写为 aborted 并丢失 `langgraph_run_id`。修复:store 层新增 `abort_run_snapshot`(`UPDATE ... SET status='aborted' WHERE snapshot_id=? AND status='prepared'`,按 rowcount 判定),`SnapshotService.abort` 改为与 `bind` 相同的读-条件写-失败重读循环;已 aborted 重复调用保持幂等成功。
|
||||
2. **Minor:selection_hash 期望值自证**。`test_selection_hash_uses_pre_resolution_semantics` 原先用测试内重复实现的同一序列化逻辑计算期望值(两侧同变不红)。改为钉死离线算出的 SHA-256 字面值(`{"auxiliary":null,"primary":null}` → `697c0462...55abc`,常量 `INHERIT_SELECTION_HASH`),锁住对外契约;删除测试内的重复实现。
|
||||
|
||||
新增测试(先红后绿):
|
||||
- `test_abort_run_snapshot_store_update_is_conditional`:store 层条件 UPDATE 的 rowcount 语义(prepared→True,重复/bound→False 且行不被改写)。
|
||||
- `test_abort_losing_bind_race_keeps_bound_state`:monkeypatch `get_run_snapshot` 在 abort 读与写之间插入并发 bind,断言 abort 抛 `SNAPSHOT_ALREADY_BOUND` 且行保持 bound、`langgraph_run_id` 完好(旧实现此测试必红)。
|
||||
|
||||
验证:
|
||||
- `.venv/bin/python -m pytest tests/test_snapshots.py -x -q` → **42 passed**
|
||||
- `.venv/bin/python -m pytest tests/ -x -q` → **3133 passed, 10 skipped**(无回归)
|
||||
- `ruff check` / `ruff format --check`(涉及文件)→ 全净
|
||||
@@ -69,7 +69,7 @@ EvoScientist is a multi-agent AI system for automated scientific experimentation
|
||||
| Framework | [DeepAgents](https://github.com/langchain-ai/deepagents) + [LangChain](https://python.langchain.com/) + [LangGraph](https://langchain-ai.github.io/langgraph/) |
|
||||
| Default model | `claude-sonnet-4-6` (Anthropic) |
|
||||
| Tests | ~890 across 36 files, no API keys needed |
|
||||
| Config file | `.data/.config/settings.yaml` (project root) |
|
||||
| Config file | `~/.config/evoscientist/config.yaml` |
|
||||
|
||||
### Sub-Agents (defined in `EvoScientist/subagent.yaml`)
|
||||
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
ARG BASE_IMAGE=ghcr.io/astral-sh/uv:python3.11-trixie-slim@sha256:7936cc6625ca04cafa6ecc3c2881ddfe90a747c55c74480cd4ac6ffad6a5af1e
|
||||
ARG NODE_IMAGE=node:24-trixie-slim@sha256:735dd688da64d22ebd9dd374b3e7e5a874635668fd2a6ec20ca1f99264294086
|
||||
|
||||
FROM ${NODE_IMAGE} AS nodejs
|
||||
|
||||
# ---------- Builder ----------
|
||||
FROM ${BASE_IMAGE} AS builder
|
||||
|
||||
ENV UV_COMPILE_BYTECODE=1 \
|
||||
UV_LINK_MODE=copy \
|
||||
UV_PYTHON_DOWNLOADS=never \
|
||||
UV_PROJECT_ENVIRONMENT=/opt/venv
|
||||
|
||||
WORKDIR /src
|
||||
|
||||
COPY pyproject.toml uv.lock README.md ./
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv sync --frozen --no-install-project --no-dev \
|
||||
--extra all-channels
|
||||
|
||||
COPY EvoScientist ./EvoScientist
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv sync --frozen --no-dev --no-editable \
|
||||
--extra all-channels
|
||||
|
||||
# ---------- Runtime ----------
|
||||
FROM ${BASE_IMAGE} AS runtime
|
||||
|
||||
RUN apt-get update \
|
||||
&& apt-get install -y --no-install-recommends \
|
||||
git \
|
||||
ca-certificates \
|
||||
tini \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY --from=nodejs /usr/local/bin/node /usr/local/bin/node
|
||||
COPY --from=nodejs /usr/local/lib/node_modules /usr/local/lib/node_modules
|
||||
RUN ln -sf /usr/local/lib/node_modules/npm/bin/npm-cli.js /usr/local/bin/npm \
|
||||
&& ln -sf /usr/local/lib/node_modules/npm/bin/npx-cli.js /usr/local/bin/npx
|
||||
|
||||
ARG UID=1000
|
||||
ARG GID=1000
|
||||
RUN groupadd --gid ${GID} evosci \
|
||||
&& useradd --uid ${UID} --gid ${GID} --create-home --shell /bin/bash evosci
|
||||
|
||||
COPY --from=builder /opt/venv /opt/venv
|
||||
|
||||
ENV PATH="/opt/venv/bin:/home/evosci/.evoscientist/.local/bin:${PATH}" \
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
EVOSCIENTIST_WORKSPACE_DIR=/workspace \
|
||||
EVOSCIENTIST_DATA_DIR=/home/evosci/.evoscientist \
|
||||
XDG_CONFIG_HOME=/home/evosci/.evoscientist/.config \
|
||||
UV_TOOL_DIR=/home/evosci/.evoscientist/.local/share/uv/tools \
|
||||
UV_TOOL_BIN_DIR=/home/evosci/.evoscientist/.local/bin
|
||||
|
||||
RUN mkdir -p /workspace \
|
||||
/home/evosci/.evoscientist/.config/evoscientist \
|
||||
/home/evosci/.evoscientist/.local/bin \
|
||||
/home/evosci/.evoscientist/.local/share/uv/tools \
|
||||
&& chown -R ${UID}:${GID} /workspace /home/evosci
|
||||
|
||||
USER evosci
|
||||
WORKDIR /workspace
|
||||
|
||||
LABEL org.opencontainers.image.title="EvoScientist" \
|
||||
org.opencontainers.image.description="EvoScientist agent with core + all-channels dependencies pre-installed." \
|
||||
org.opencontainers.image.source="https://github.com/EvoScientist/EvoScientist" \
|
||||
org.opencontainers.image.documentation="https://github.com/EvoScientist/EvoScientist#-docker" \
|
||||
org.opencontainers.image.licenses="Apache-2.0"
|
||||
|
||||
ENTRYPOINT ["tini", "--", "evosci"]
|
||||
@@ -9,14 +9,13 @@ from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
from ._version import __version__
|
||||
|
||||
_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
# Agent graph (lazy to avoid expensive initialization at import time)
|
||||
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
|
||||
"create_cli_agent": (".EvoScientist", "create_cli_agent"),
|
||||
# Backends
|
||||
"CustomSandboxBackend": (".backends", "CustomSandboxBackend"),
|
||||
"MemoryFilesystemBackend": (".backends", "MemoryFilesystemBackend"),
|
||||
"ReadOnlyFilesystemBackend": (".backends", "ReadOnlyFilesystemBackend"),
|
||||
# Configuration
|
||||
"EvoScientistConfig": (".config", "EvoScientistConfig"),
|
||||
@@ -24,26 +23,16 @@ _EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"save_config": (".config", "save_config"),
|
||||
"get_effective_config": (".config", "get_effective_config"),
|
||||
"get_config_path": (".config", "get_config_path"),
|
||||
# LLM
|
||||
"get_chat_model": (".llm", "get_chat_model"),
|
||||
"list_models": (".llm", "list_models"),
|
||||
# Prompts
|
||||
"get_system_prompt": (".prompts", "get_system_prompt"),
|
||||
"RESEARCHER_INSTRUCTIONS": (".prompts", "RESEARCHER_INSTRUCTIONS"),
|
||||
# Tools
|
||||
"tavily_search": (".tools", "tavily_search"),
|
||||
"web_search": (".tools", "web_search"),
|
||||
"web_extract": (".tools", "web_extract"),
|
||||
"web_crawl": (".tools", "web_crawl"),
|
||||
"think_tool": (".tools", "think_tool"),
|
||||
# Sessions
|
||||
"get_checkpointer": (".sessions", "get_checkpointer"),
|
||||
"generate_thread_id": (".sessions", "generate_thread_id"),
|
||||
"list_threads": (".sessions", "list_threads"),
|
||||
"delete_thread": (".sessions", "delete_thread"),
|
||||
"get_storage_stats": (".sessions", "get_storage_stats"),
|
||||
"get_aggregated_storage_stats": (".sessions", "get_aggregated_storage_stats"),
|
||||
"list_all_session_db_paths": (".sessions", "list_all_session_db_paths"),
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
__version__ = "0.1.19"
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Windows asyncio event-loop policy compatibility.
|
||||
|
||||
On Windows, MCP stdio servers are launched as subprocesses by the MCP SDK's
|
||||
stdio transport, which uses ``anyio.open_process`` → ``asyncio`` async
|
||||
subprocess support. ``asyncio``'s *Selector* event loop does **not** implement
|
||||
async subprocess creation, so on a Selector loop the stdio transport falls back
|
||||
to a synchronous ``subprocess.Popen`` inside an ``async`` function. Under
|
||||
``langgraph dev`` (which enables ``blockbuster`` by default to police blocking
|
||||
I/O) that synchronous call is flagged as a ``BlockingError`` — see issue #283.
|
||||
|
||||
The *Proactor* loop supports async subprocesses natively, so the fallback never
|
||||
happens and ``blockbuster`` allows the (now genuinely async) spawn.
|
||||
|
||||
Windows has defaulted to ``WindowsProactorEventLoopPolicy`` since Python 3.8, so
|
||||
this is normally a no-op. We set it explicitly anyway as a safeguard: a
|
||||
dependency, IDE, or notebook host may have installed a Selector policy earlier
|
||||
in the process, and the ``langgraph dev`` subprocess in particular runs code we
|
||||
don't fully control. Calling this at each process entrypoint — **before any
|
||||
event loop is created** — guarantees the MCP subprocess path stays async.
|
||||
|
||||
This must run at import/startup time, ahead of the first ``asyncio.run`` /
|
||||
``new_event_loop`` call; once a loop exists, swapping the policy does not change
|
||||
the already-running loop.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
|
||||
|
||||
def ensure_proactor_event_loop_policy() -> bool:
|
||||
"""Install ``WindowsProactorEventLoopPolicy`` on Windows if needed.
|
||||
|
||||
Returns ``True`` if a Proactor policy is in effect afterwards (always
|
||||
``False`` off Windows, where the concept doesn't apply). Safe and idempotent
|
||||
to call multiple times; a no-op on non-Windows platforms.
|
||||
"""
|
||||
if sys.platform != "win32":
|
||||
return False
|
||||
|
||||
import asyncio
|
||||
|
||||
proactor_policy = getattr(asyncio, "WindowsProactorEventLoopPolicy", None)
|
||||
if proactor_policy is None: # pragma: no cover - non-Windows / stripped build
|
||||
return False
|
||||
|
||||
current = asyncio.get_event_loop_policy()
|
||||
if not isinstance(current, proactor_policy):
|
||||
asyncio.set_event_loop_policy(proactor_policy())
|
||||
return True
|
||||
@@ -0,0 +1,334 @@
|
||||
"""Background OS-process execution for the sandbox.
|
||||
|
||||
A *process* here is a single detached OS process launched via ``run_in_background``
|
||||
(distinct from an async sub-agent *task* and a future cron *schedule* — the word
|
||||
"job" is intentionally never used).
|
||||
|
||||
The registry is **module-global (process-level)**: processes survive ``/new`` and
|
||||
``/resume`` within the same CLI process, but are not persisted across a CLI restart.
|
||||
The live ``Popen`` handle is held so ``poll()`` / ``returncode`` stay authoritative
|
||||
(no PID-reuse risk).
|
||||
|
||||
Command validation and cwd resolution happen at the tool layer
|
||||
(``middleware/background.py``); this module is the pure execution + tracking mechanism
|
||||
and is safe to unit-test on its own. A future scheduler (cron) would reuse ``launch``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
import psutil
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_BG_DIRNAME = ".bg_processes"
|
||||
_KILL_GRACE_SECONDS = 2.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class BgProcess:
|
||||
"""A tracked background OS process."""
|
||||
|
||||
process_id: str
|
||||
name: str
|
||||
command: str
|
||||
popen: subprocess.Popen
|
||||
pid: int
|
||||
log_path: Path
|
||||
started_at: str # ISO-8601 UTC (record/display)
|
||||
started_ts: float # epoch seconds (elapsed computation)
|
||||
origin_thread_id: str | None = None # CLI thread/session that launched it
|
||||
returncode: int | None = None
|
||||
finished_at: str | None = None
|
||||
finished_ts: float | None = None # epoch at exit; freezes elapsed once done
|
||||
stopped: bool = False # set by stop(); suppresses the completion notification
|
||||
# epoch each thread last checked this process (status/list); keyed by thread_id
|
||||
# so a check from one session can't dedup another session's completion ping.
|
||||
last_checked_by_thread: dict[str | None, float] = field(default_factory=dict)
|
||||
|
||||
|
||||
_PROCESSES: dict[str, BgProcess] = {}
|
||||
_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _now_iso() -> str:
|
||||
return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
|
||||
|
||||
def _record_exit(proc: BgProcess) -> None:
|
||||
"""Record terminal state on first observed exit. Caller MUST hold ``_LOCK``.
|
||||
|
||||
``finished_ts`` is set when the exit is first observed. The per-process daemon
|
||||
watcher (:func:`_watch`) calls this right after ``popen.wait()`` returns, so in
|
||||
practice ``finished_ts`` ≈ the real exit time. Calls from ``status`` / ``list_all`` /
|
||||
``stop`` are a fallback for the brief window before the watcher runs.
|
||||
"""
|
||||
rc = proc.popen.poll()
|
||||
if rc is not None and proc.returncode is None:
|
||||
proc.returncode = rc
|
||||
proc.finished_at = _now_iso()
|
||||
proc.finished_ts = time.time()
|
||||
|
||||
|
||||
def _elapsed(proc: BgProcess) -> int:
|
||||
"""Seconds the process has run — frozen at first-observed exit once it has exited."""
|
||||
end = proc.finished_ts if proc.finished_ts is not None else time.time()
|
||||
return int(end - proc.started_ts)
|
||||
|
||||
|
||||
def was_observed_done(process_id: str, origin_thread_id: str | None = None) -> bool:
|
||||
"""True if ``origin_thread_id`` already saw this process's completion itself.
|
||||
|
||||
i.e. the process has exited AND was checked (``status``/``list_all``) from that thread
|
||||
at or after it finished. Used to dedup the completion notification (routed to the
|
||||
launching thread), so a check from a *different* session can't suppress it.
|
||||
"""
|
||||
with _LOCK:
|
||||
proc = _PROCESSES.get(process_id)
|
||||
if proc is None or proc.finished_ts is None:
|
||||
return False
|
||||
seen_ts = proc.last_checked_by_thread.get(origin_thread_id)
|
||||
return seen_ts is not None and seen_ts >= proc.finished_ts
|
||||
|
||||
|
||||
def _read_tail(log_path: Path, tail_bytes: int) -> str:
|
||||
# Seek from the end so a huge log isn't fully read into memory on each status check.
|
||||
try:
|
||||
with log_path.open("rb") as f:
|
||||
f.seek(0, os.SEEK_END)
|
||||
size = f.tell()
|
||||
if size == 0:
|
||||
return "(no output yet)"
|
||||
if size > tail_bytes:
|
||||
f.seek(-tail_bytes, os.SEEK_END)
|
||||
return "...(truncated)...\n" + f.read().decode("utf-8", "replace")
|
||||
f.seek(0)
|
||||
data = f.read()
|
||||
except OSError:
|
||||
return "(no output captured yet)"
|
||||
return data.decode("utf-8", "replace")
|
||||
|
||||
|
||||
def _watch(proc: BgProcess, on_exit: Callable[[BgProcess], None] | None) -> None:
|
||||
"""Block until ``proc`` exits, record the exit promptly, then fire ``on_exit``.
|
||||
|
||||
Running in a daemon thread, ``popen.wait()`` lets us record ``finished_ts`` at (very
|
||||
close to) the real exit time — fixing the observation-time inflation — and gives a
|
||||
hook the CLI layer wires to a completion notification, without ``background.py``
|
||||
importing the notifier (kept decoupled via the callback).
|
||||
"""
|
||||
try:
|
||||
proc.popen.wait()
|
||||
except Exception:
|
||||
pass
|
||||
with _LOCK:
|
||||
_record_exit(proc)
|
||||
if on_exit is not None:
|
||||
try:
|
||||
on_exit(proc)
|
||||
except Exception:
|
||||
logger.warning("background on_exit callback failed", exc_info=True)
|
||||
|
||||
|
||||
def launch(
|
||||
command: str,
|
||||
cwd: str,
|
||||
name: str | None = None,
|
||||
*,
|
||||
origin_thread_id: str | None = None,
|
||||
on_exit: Callable[[BgProcess], None] | None = None,
|
||||
) -> str:
|
||||
"""Launch ``command`` detached in ``cwd``; return a short ``process_id``.
|
||||
|
||||
The command is run via ``shell=True`` with output redirected to a per-process log
|
||||
file under ``<cwd>/.bg_processes/`` and ``start_new_session=True`` so the child is a
|
||||
process-group leader (survives this call's return and can be killed as a group).
|
||||
The caller is responsible for validating ``command`` first.
|
||||
|
||||
``origin_thread_id`` records the launching CLI session so ``list_all`` can scope to it.
|
||||
``on_exit`` (optional) is called with the ``BgProcess`` from a daemon watcher thread
|
||||
once the process exits — used by the CLI layer to emit a completion notification.
|
||||
"""
|
||||
process_id = uuid.uuid4().hex[:8]
|
||||
log_dir = Path(cwd) / _BG_DIRNAME
|
||||
log_dir.mkdir(parents=True, exist_ok=True)
|
||||
log_path = log_dir / f"{process_id}.log"
|
||||
|
||||
log_file = open(log_path, "w")
|
||||
try:
|
||||
popen = subprocess.Popen(
|
||||
command,
|
||||
shell=True,
|
||||
cwd=cwd,
|
||||
stdout=log_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
stdin=subprocess.DEVNULL,
|
||||
start_new_session=True,
|
||||
)
|
||||
finally:
|
||||
# The child inherited its own dup of the fd during spawn; the parent's copy
|
||||
# is no longer needed (and must be closed so the pipe/file isn't held open).
|
||||
log_file.close()
|
||||
|
||||
proc = BgProcess(
|
||||
process_id=process_id,
|
||||
name=name or command[:40],
|
||||
command=command,
|
||||
popen=popen,
|
||||
pid=popen.pid,
|
||||
log_path=log_path,
|
||||
started_at=_now_iso(),
|
||||
started_ts=time.time(),
|
||||
origin_thread_id=origin_thread_id,
|
||||
)
|
||||
with _LOCK:
|
||||
_PROCESSES[process_id] = proc
|
||||
# Daemon watcher: records the precise exit time and fires on_exit when done.
|
||||
threading.Thread(target=_watch, args=(proc, on_exit), daemon=True).start()
|
||||
return process_id
|
||||
|
||||
|
||||
def status(
|
||||
process_id: str, *, thread_id: str | None = None, tail_bytes: int = 16_000
|
||||
) -> str:
|
||||
"""Return a human-readable status + recent output tail for ``process_id``."""
|
||||
with _LOCK:
|
||||
proc = _PROCESSES.get(process_id)
|
||||
if proc is None:
|
||||
return (
|
||||
f"No such background process: {process_id!r}. "
|
||||
"Use list_processes to see tracked processes."
|
||||
)
|
||||
_record_exit(proc)
|
||||
proc.last_checked_by_thread[thread_id] = time.time() # this thread observed it
|
||||
running = proc.returncode is None
|
||||
elapsed = _elapsed(proc)
|
||||
name, pid, command, returncode, log_path = (
|
||||
proc.name,
|
||||
proc.pid,
|
||||
proc.command,
|
||||
proc.returncode,
|
||||
proc.log_path,
|
||||
)
|
||||
if running:
|
||||
head = f"Process {process_id} (name={name!r}) RUNNING — {elapsed}s elapsed, pid {pid}."
|
||||
else:
|
||||
head = f"Process {process_id} (name={name!r}) EXITED code {returncode} after ~{elapsed}s."
|
||||
tail = _read_tail(log_path, tail_bytes) # file IO outside the lock
|
||||
return (
|
||||
f"{head}\nCommand: {command}\n--- output (last {tail_bytes} bytes) ---\n{tail}"
|
||||
)
|
||||
|
||||
|
||||
def _kill_process_tree(popen: subprocess.Popen, *, forceful: bool) -> None:
|
||||
"""Kill the process group/tree in a cross-platform way.
|
||||
|
||||
On POSIX ``start_new_session=True`` makes the child a process-group
|
||||
leader; ``os.killpg`` terminates the entire group (shell + any
|
||||
grandchildren). On Windows ``TerminateProcess`` (used by
|
||||
``Popen.terminate()`` / ``Popen.kill()``) only kills the direct
|
||||
child — it does *not* cascade to grandchildren. We use ``psutil``
|
||||
to walk the process tree and signal every descendant.
|
||||
"""
|
||||
if os.name == "nt":
|
||||
try:
|
||||
proc = psutil.Process(popen.pid)
|
||||
targets = [proc, *proc.children(recursive=True)]
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||
return
|
||||
for p in targets:
|
||||
try:
|
||||
if forceful:
|
||||
p.kill()
|
||||
else:
|
||||
p.terminate()
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied):
|
||||
pass
|
||||
else:
|
||||
sig = signal.SIGKILL if forceful else signal.SIGTERM
|
||||
try:
|
||||
os.killpg(os.getpgid(popen.pid), sig)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
|
||||
def stop(process_id: str) -> str:
|
||||
"""Terminate ``process_id`` and its process group (SIGTERM, then SIGKILL)."""
|
||||
with _LOCK:
|
||||
proc = _PROCESSES.get(process_id)
|
||||
if proc is None:
|
||||
return f"No such background process: {process_id!r}."
|
||||
if proc.popen.poll() is not None:
|
||||
_record_exit(proc)
|
||||
return f"Process {process_id} already finished (code {proc.returncode})."
|
||||
# Mark as user-stopped so the watcher's on_exit suppresses the completion
|
||||
# notification (the user already knows — no need to ping them).
|
||||
proc.stopped = True
|
||||
# The watcher's popen.wait() reaps without the lock, so a tiny PID-reuse race
|
||||
# remains (getpgid on a recycled pid). On POSIX ProcessLookupError covers the
|
||||
# common case; on Windows ``Popen.terminate()`` is a no-op on a dead handle
|
||||
# so we poll after the call instead.
|
||||
_kill_process_tree(proc.popen, forceful=False)
|
||||
if proc.popen.poll() is not None:
|
||||
_record_exit(proc)
|
||||
return f"Process {process_id} is no longer running."
|
||||
|
||||
deadline = time.time() + _KILL_GRACE_SECONDS
|
||||
while time.time() < deadline:
|
||||
with _LOCK:
|
||||
if proc.popen.poll() is not None:
|
||||
_record_exit(proc)
|
||||
break
|
||||
time.sleep(0.1)
|
||||
else:
|
||||
with _LOCK:
|
||||
if proc.popen.poll() is None:
|
||||
_kill_process_tree(proc.popen, forceful=True)
|
||||
_record_exit(proc)
|
||||
|
||||
with _LOCK:
|
||||
_record_exit(proc)
|
||||
name = proc.name
|
||||
return f"Stopped background process {process_id} (name={name!r})."
|
||||
|
||||
|
||||
def list_all(thread_id: str | None = None, *, include_all: bool = False) -> str:
|
||||
"""List tracked background processes with live statuses.
|
||||
|
||||
Scoped to the launching session (``thread_id``) unless ``include_all`` is set.
|
||||
"""
|
||||
with _LOCK:
|
||||
all_procs = list(_PROCESSES.values())
|
||||
procs = (
|
||||
all_procs
|
||||
if include_all
|
||||
else [p for p in all_procs if p.origin_thread_id == thread_id]
|
||||
)
|
||||
if not procs:
|
||||
if all_procs and not include_all:
|
||||
return (
|
||||
"No background processes in this session "
|
||||
f"({len(all_procs)} in other sessions — pass all_threads=True to see them)."
|
||||
)
|
||||
return "No background processes tracked."
|
||||
lines = []
|
||||
now = time.time()
|
||||
for p in procs:
|
||||
_record_exit(p)
|
||||
p.last_checked_by_thread[thread_id] = now # this thread observed it
|
||||
state = "RUNNING" if p.returncode is None else f"exited({p.returncode})"
|
||||
lines.append(
|
||||
f" {p.process_id} {state:12} {_elapsed(p)}s name={p.name!r}"
|
||||
)
|
||||
return f"{len(procs)} background process(es):\n" + "\n".join(lines)
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
EvoScientist provides unified integration with 10 messaging platforms. This document covers the architecture overview, message processing pipeline, capability matrix, security model, deployment guides, and troubleshooting.
|
||||
|
||||
Configuration file: `~/.config/ai4scientist/settings.yaml` (or use environment variables with the `EVOSCIENTIST_` prefix).
|
||||
Configuration file: `~/.config/evoscientist/config.yaml` (or use environment variables with the `EVOSCIENTIST_` prefix).
|
||||
|
||||
## Table of Contents
|
||||
|
||||
|
||||
@@ -2,23 +2,20 @@
|
||||
|
||||
Channels push messages to the inbound queue; the agent (or any consumer)
|
||||
reads from inbound, processes, and pushes responses to the outbound queue.
|
||||
A background dispatcher routes outbound messages to the correct channel
|
||||
via subscriber callbacks.
|
||||
``ChannelManager._dispatch_outbound`` routes outbound messages to the
|
||||
correct channel by looking up its registered :class:`Channel` instance.
|
||||
|
||||
Deduplication is handled at the Channel level (single dedup point).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from ..debug import TraceMixin, debug_trace_enabled
|
||||
from .events import InboundMessage, OutboundMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
OutboundCallback = Callable[[OutboundMessage], Awaitable[None]]
|
||||
|
||||
|
||||
class MessageBus(TraceMixin):
|
||||
"""Async message bus that decouples chat channels from the agent core."""
|
||||
@@ -28,8 +25,6 @@ class MessageBus(TraceMixin):
|
||||
def __init__(self):
|
||||
self.inbound: asyncio.Queue[InboundMessage] = asyncio.Queue(maxsize=5000)
|
||||
self.outbound: asyncio.Queue[OutboundMessage] = asyncio.Queue(maxsize=5000)
|
||||
self._outbound_subscribers: dict[str, list[OutboundCallback]] = {}
|
||||
self._running = False
|
||||
self._debug_trace = debug_trace_enabled()
|
||||
self._trace_logger = logger
|
||||
|
||||
@@ -53,58 +48,6 @@ class MessageBus(TraceMixin):
|
||||
"""Consume the next outbound message (blocks until available)."""
|
||||
return await self.outbound.get()
|
||||
|
||||
# ── subscriber routing ──
|
||||
|
||||
def subscribe_outbound(
|
||||
self,
|
||||
channel: str,
|
||||
callback: OutboundCallback,
|
||||
) -> None:
|
||||
"""Register a callback for outbound messages targeting *channel*."""
|
||||
if channel not in self._outbound_subscribers:
|
||||
self._outbound_subscribers[channel] = []
|
||||
self._outbound_subscribers[channel].append(callback)
|
||||
|
||||
async def dispatch_outbound(self) -> None:
|
||||
"""Route outbound messages to subscribed channels.
|
||||
|
||||
Run as a background task — loops until :meth:`stop` is called.
|
||||
"""
|
||||
self._running = True
|
||||
while self._running:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self.outbound.get(),
|
||||
timeout=1.0,
|
||||
)
|
||||
except TimeoutError:
|
||||
continue
|
||||
subscribers = self._outbound_subscribers.get(msg.channel, [])
|
||||
if not subscribers:
|
||||
self._trace_event(
|
||||
"bus_dispatch_drop",
|
||||
target_channel=msg.channel,
|
||||
reason="no_subscriber",
|
||||
chat_id=msg.chat_id,
|
||||
)
|
||||
logger.warning(f"No subscriber for channel: {msg.channel}")
|
||||
continue
|
||||
for callback in subscribers:
|
||||
try:
|
||||
await callback(msg)
|
||||
except Exception as e:
|
||||
self._trace_event(
|
||||
"bus_dispatch_error",
|
||||
target_channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
error_type=type(e).__name__,
|
||||
)
|
||||
logger.error(f"Error dispatching to {msg.channel}: {e}")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the dispatcher loop."""
|
||||
self._running = False
|
||||
|
||||
@property
|
||||
def inbound_size(self) -> int:
|
||||
return self.inbound.qsize()
|
||||
|
||||
@@ -160,6 +160,7 @@ QQ = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=4096,
|
||||
typing=False, # no typing API for QQ bots
|
||||
inline_buttons=True, # markdown + keyboard payload (C2C only)
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=False, # qq-botpy does not expose voice as a distinct message type
|
||||
|
||||
@@ -11,13 +11,12 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
|
||||
from .base import Channel
|
||||
from .bus import MessageBus
|
||||
from .bus.events import InboundMessage, OutboundMessage
|
||||
@@ -119,7 +118,7 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
|
||||
return True
|
||||
|
||||
try:
|
||||
from ..config.settings import load_config
|
||||
from ..config.settings import HITL_SHELL_TOOLS, load_config
|
||||
|
||||
cfg = load_config()
|
||||
except Exception:
|
||||
@@ -135,14 +134,10 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
|
||||
)
|
||||
|
||||
for req in action_requests:
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
if name != "execute":
|
||||
name = req.get("name", "")
|
||||
if name not in HITL_SHELL_TOOLS:
|
||||
continue
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
args = req.get("args", {})
|
||||
command = args.get("command", "") if isinstance(args, dict) else ""
|
||||
cmd = command.strip()
|
||||
if not any(cmd.startswith(prefix) for prefix in shell_allow_list):
|
||||
@@ -150,16 +145,18 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _format_approval_prompt(action_requests: list[dict]) -> str:
|
||||
"""Format an approval prompt as a text message for channel users."""
|
||||
def _format_approval_prompt(
|
||||
action_requests: list[dict], *, with_buttons: bool = False
|
||||
) -> str:
|
||||
"""Format an approval prompt as a text message for channel users.
|
||||
|
||||
When *with_buttons* is True, the trailing "Reply: 1=Approve..."
|
||||
instruction is dropped — the buttons replace the textual cue.
|
||||
"""
|
||||
lines = ["\u26a0\ufe0f Approval Required\n"]
|
||||
for i, req in enumerate(action_requests, 1):
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
name = req.get("name", "")
|
||||
args = req.get("args", {})
|
||||
if isinstance(args, dict):
|
||||
command = args.get("command", args.get("path", ""))
|
||||
else:
|
||||
@@ -168,6 +165,7 @@ def _format_approval_prompt(action_requests: list[dict]) -> str:
|
||||
lines.append(f" {i}. {name}: {command}")
|
||||
else:
|
||||
lines.append(f" {i}. {name}")
|
||||
if not with_buttons:
|
||||
lines.append("")
|
||||
lines.append("Reply: 1=Approve, 2=Reject, 3=Approve all")
|
||||
lines.append("(Auto-reject in 2 min if no reply)")
|
||||
@@ -189,6 +187,25 @@ def _parse_approval_reply(text: str) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _approval_prompt_metadata(
|
||||
base_metadata: dict | None, *, with_buttons: bool
|
||||
) -> dict:
|
||||
"""Outbound metadata for the HITL approval prompt.
|
||||
|
||||
When *with_buttons* is True, attaches Approve/Reject/Auto buttons whose
|
||||
values match ``_parse_approval_reply`` so a click flows through the same
|
||||
path as a typed ``"1"``/``"2"``/``"3"`` reply.
|
||||
"""
|
||||
metadata = dict(base_metadata or {})
|
||||
if with_buttons:
|
||||
metadata["buttons"] = [
|
||||
{"text": "Approve", "value": "1", "type": "primary"},
|
||||
{"text": "Reject", "value": "2", "type": "danger"},
|
||||
{"text": "Approve all", "value": "3"},
|
||||
]
|
||||
return metadata
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingInterrupt:
|
||||
"""Stored state for a pending HITL interrupt awaiting channel user reply."""
|
||||
@@ -217,9 +234,11 @@ class InboundConsumer:
|
||||
manager:
|
||||
The ChannelManager (used to look up channel instances).
|
||||
agent:
|
||||
The agent object (must support ``stream_agent_events``).
|
||||
The local agent object used by local graph gateway targets.
|
||||
thread_id:
|
||||
Default thread ID for agent conversations.
|
||||
graph_gateway:
|
||||
Gateway used for thread creation and graph streaming.
|
||||
send_thinking:
|
||||
Whether to forward thinking messages to the channel.
|
||||
on_message_received:
|
||||
@@ -250,6 +269,7 @@ class InboundConsumer:
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
*,
|
||||
graph_gateway: GraphGateway,
|
||||
send_thinking: bool = False,
|
||||
on_message_received: Callable[[InboundMessage], None] | None = None,
|
||||
on_streaming_event: Callable[[dict], None] | None = None,
|
||||
@@ -263,6 +283,7 @@ class InboundConsumer:
|
||||
self.manager = manager
|
||||
self.agent = agent
|
||||
self.thread_id = thread_id
|
||||
self.graph_gateway = graph_gateway
|
||||
self.send_thinking = send_thinking
|
||||
self._on_message_received = on_message_received
|
||||
self._on_streaming_event = on_streaming_event
|
||||
@@ -296,7 +317,7 @@ class InboundConsumer:
|
||||
# ask_user: pending reply per session_key
|
||||
self._pending_ask_user_replies: dict[str, _PendingAskUserReply] = {}
|
||||
|
||||
def _get_thread_id(self, sender_id: str) -> str:
|
||||
async def _get_thread_id(self, sender_id: str) -> str:
|
||||
"""Get or create a thread ID for the given sender.
|
||||
|
||||
Uses LRU ordering: recently accessed senders are moved to the
|
||||
@@ -312,7 +333,9 @@ class InboundConsumer:
|
||||
if self.thread_id:
|
||||
self._sessions[sender_id] = f"{self.thread_id}:{sender_id}"
|
||||
else:
|
||||
self._sessions[sender_id] = str(uuid.uuid4())
|
||||
self._sessions[sender_id] = await self.graph_gateway.create_thread(
|
||||
GraphTarget(local_graph=self.agent)
|
||||
)
|
||||
return self._sessions[sender_id]
|
||||
|
||||
def _get_channel(self, channel_name: str) -> Channel | None:
|
||||
@@ -406,7 +429,7 @@ class InboundConsumer:
|
||||
pass
|
||||
|
||||
channel = self._get_channel(msg.channel)
|
||||
thread_id = self._get_thread_id(msg.sender_id)
|
||||
thread_id = await self._get_thread_id(msg.sender_id)
|
||||
session_key = msg.session_key # "channel:chat_id"
|
||||
|
||||
# Lazily create per-chat lock; evict stale locks when too many
|
||||
@@ -449,15 +472,16 @@ class InboundConsumer:
|
||||
session_key: str,
|
||||
) -> None:
|
||||
"""Stream agent events with HITL interrupt handling."""
|
||||
from ..stream.events import stream_agent_events
|
||||
from langgraph.types import Command
|
||||
|
||||
stream_input: Any = msg.content
|
||||
_t0 = time.monotonic()
|
||||
stream_input: GraphRunInput = msg.content
|
||||
|
||||
try:
|
||||
if channel:
|
||||
await channel.start_typing(msg.chat_id)
|
||||
|
||||
_last_sent_thinking: str | None = None
|
||||
|
||||
for _hitl_round in range(_MAX_HITL_ROUNDS):
|
||||
final_content = ""
|
||||
thinking_buffer: list[str] = []
|
||||
@@ -466,14 +490,38 @@ class InboundConsumer:
|
||||
thinking_sent = False
|
||||
interrupt_data: dict | None = None
|
||||
|
||||
async def _flush_thinking_buffer(
|
||||
buffer: list[str] = thinking_buffer,
|
||||
) -> bool:
|
||||
"""Send the current thinking buffer, dedup by content."""
|
||||
nonlocal thinking_sent, _last_sent_thinking
|
||||
if not channel or thinking_sent or not buffer:
|
||||
return False
|
||||
|
||||
full_thinking = "".join(buffer).rstrip()
|
||||
buffer.clear()
|
||||
if not full_thinking or full_thinking == _last_sent_thinking:
|
||||
return False
|
||||
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id,
|
||||
full_thinking,
|
||||
msg.metadata,
|
||||
)
|
||||
thinking_sent = True
|
||||
_last_sent_thinking = full_thinking
|
||||
return True
|
||||
|
||||
async for event in _timeout_aiter(
|
||||
stream_agent_events(
|
||||
self.agent,
|
||||
stream_input,
|
||||
thread_id,
|
||||
self.graph_gateway.stream_events(
|
||||
RunRequest(
|
||||
message=stream_input,
|
||||
thread_id=thread_id,
|
||||
media=msg.media or None
|
||||
if isinstance(stream_input, str)
|
||||
else None,
|
||||
target=GraphTarget(local_graph=self.agent),
|
||||
)
|
||||
),
|
||||
self._inference_timeout,
|
||||
):
|
||||
@@ -494,16 +542,7 @@ class InboundConsumer:
|
||||
if event.get("name") == "write_todos" and not todo_sent:
|
||||
todos = event.get("args", {}).get("todos", [])
|
||||
if todos and channel:
|
||||
if thinking_buffer and not thinking_sent:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if full_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id,
|
||||
full_thinking,
|
||||
msg.metadata,
|
||||
)
|
||||
thinking_sent = True
|
||||
thinking_buffer.clear()
|
||||
await _flush_thinking_buffer()
|
||||
await channel.send_todo_message(
|
||||
msg.sender_id,
|
||||
_format_todo_list(todos),
|
||||
@@ -514,14 +553,11 @@ class InboundConsumer:
|
||||
elif event_type == "text":
|
||||
final_content += event.get("content", "")
|
||||
|
||||
elif event_type == "progress":
|
||||
# Internal agent planning — ignore for channel delivery.
|
||||
# This text should NOT appear in the user-facing response.
|
||||
pass
|
||||
|
||||
elif event_type == "subagent_text":
|
||||
sa_name = event.get("subagent", "unknown")
|
||||
instance_id = event.get("instance_id") or sa_name
|
||||
instance_id = event.get("instance_id")
|
||||
if not instance_id:
|
||||
continue
|
||||
if instance_id not in subagent_text_buffers:
|
||||
subagent_text_buffers[instance_id] = (sa_name, [])
|
||||
subagent_text_buffers[instance_id][1].append(
|
||||
@@ -540,14 +576,7 @@ class InboundConsumer:
|
||||
break # exit async for to handle ask_user
|
||||
|
||||
# Flush thinking
|
||||
if thinking_buffer and not thinking_sent and channel:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if full_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id,
|
||||
full_thinking,
|
||||
msg.metadata,
|
||||
)
|
||||
await _flush_thinking_buffer()
|
||||
|
||||
# No interrupt — normal completion
|
||||
if interrupt_data is None:
|
||||
@@ -562,13 +591,6 @@ class InboundConsumer:
|
||||
)
|
||||
await self.bus.publish_outbound(outbound)
|
||||
self._metrics.total_successes += 1
|
||||
elapsed = time.monotonic() - _t0
|
||||
logger.info(
|
||||
"stream completed: %d chars, %.2fs, session=%s",
|
||||
len(outbound.content),
|
||||
elapsed,
|
||||
session_key,
|
||||
)
|
||||
if self._on_message_sent:
|
||||
try:
|
||||
self._on_message_sent(outbound)
|
||||
@@ -583,7 +605,6 @@ class InboundConsumer:
|
||||
interrupt_data,
|
||||
session_key,
|
||||
)
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
stream_input = Command(resume=result)
|
||||
continue
|
||||
@@ -594,8 +615,6 @@ class InboundConsumer:
|
||||
|
||||
# Session auto-approve (user previously chose "Approve all")
|
||||
if session_key in self._auto_approve_sessions:
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
@@ -603,21 +622,27 @@ class InboundConsumer:
|
||||
|
||||
# Config auto-approve (auto_approve, non-execute, allow_list)
|
||||
if _should_auto_approve(action_reqs):
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
continue
|
||||
|
||||
# Needs user approval — send prompt to channel
|
||||
prompt_text = _format_approval_prompt(action_reqs)
|
||||
has_buttons = (
|
||||
channel is not None and channel.capabilities.inline_buttons
|
||||
)
|
||||
prompt_text = _format_approval_prompt(
|
||||
action_reqs, with_buttons=has_buttons
|
||||
)
|
||||
approval_metadata = _approval_prompt_metadata(
|
||||
msg.metadata, with_buttons=has_buttons
|
||||
)
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=prompt_text,
|
||||
metadata=msg.metadata,
|
||||
metadata=approval_metadata,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -629,35 +654,61 @@ class InboundConsumer:
|
||||
)
|
||||
self._pending_interrupts[session_key] = pending
|
||||
|
||||
timed_out = False
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
pending.event.wait(),
|
||||
timeout=_HITL_APPROVAL_TIMEOUT,
|
||||
)
|
||||
except TimeoutError:
|
||||
# Auto-approve on timeout
|
||||
pending.decision = "approve"
|
||||
timed_out = True
|
||||
finally:
|
||||
# Unregister BEFORE any further await so a late reply can't flip
|
||||
# the decision back to approve during the notification round-trip.
|
||||
self._pending_interrupts.pop(session_key, None)
|
||||
|
||||
decision = pending.decision or "approve"
|
||||
|
||||
if decision == "reject":
|
||||
if timed_out:
|
||||
# Reject on timeout (fail-closed; matches cli/channel.py). Decision
|
||||
# is a local constant, not pending.decision, so it can't be
|
||||
# overwritten by a late reply after we unregistered above.
|
||||
decision = "reject"
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Tool execution rejected.",
|
||||
content="⏰ Approval timed out. Action rejected.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
else:
|
||||
decision = pending.decision or "reject"
|
||||
|
||||
# Visible confirmation so the click/reply registers (QQ has no
|
||||
# message recall API for C2C). Only fires when the user
|
||||
# actually responded — silent on timeout to avoid claiming
|
||||
# the user approved when they just walked away.
|
||||
if pending.event.is_set():
|
||||
feedback_text = {
|
||||
"approve": "\u2705 已批准",
|
||||
"auto": "\u2705 已批准(后续自动通过)",
|
||||
"reject": "\u274c 已拒绝",
|
||||
}.get(decision)
|
||||
if feedback_text:
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=feedback_text,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
if decision == "reject":
|
||||
return
|
||||
|
||||
if decision == "auto":
|
||||
self._auto_approve_sessions.add(session_key)
|
||||
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
@@ -665,14 +716,9 @@ class InboundConsumer:
|
||||
|
||||
except TimeoutError:
|
||||
self._metrics.total_timeouts += 1
|
||||
elapsed = time.monotonic() - _t0
|
||||
logger.error(
|
||||
"Inference timeout (%ds idle) for %s in %s, elapsed=%.2fs, response=%d chars",
|
||||
self._inference_timeout,
|
||||
msg.sender_id,
|
||||
session_key,
|
||||
elapsed,
|
||||
len(final_content),
|
||||
f"Inference timeout ({self._inference_timeout}s idle) "
|
||||
f"for {msg.sender_id} in {session_key}"
|
||||
)
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
@@ -685,14 +731,7 @@ class InboundConsumer:
|
||||
|
||||
except Exception as e:
|
||||
self._metrics.total_failures += 1
|
||||
elapsed = time.monotonic() - _t0
|
||||
logger.error(
|
||||
"Agent error: %s | elapsed=%.2fs, response=%d chars, session=%s",
|
||||
e,
|
||||
elapsed,
|
||||
len(final_content),
|
||||
session_key,
|
||||
)
|
||||
logger.error(f"Agent error: {e}")
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import DingTalkChannel, DingTalkConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -68,7 +72,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = DingTalkConfig(
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import DiscordChannel, DiscordConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -69,7 +73,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = DiscordConfig(
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import EmailChannel, EmailConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -95,7 +99,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = EmailConfig(
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from ..channel_manager import _parse_csv, register_channel
|
||||
from .channel import FeishuChannel, FeishuConfig
|
||||
from .onboard import qr_register
|
||||
|
||||
__all__ = ["FeishuChannel", "FeishuConfig"]
|
||||
__all__ = ["FeishuChannel", "FeishuConfig", "qr_register"]
|
||||
|
||||
|
||||
def create_from_config(config) -> FeishuChannel:
|
||||
|
||||
@@ -376,6 +376,31 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
.build()
|
||||
)
|
||||
|
||||
# Silently absorb events we don't have a handler for. Feishu auto-
|
||||
# subscribes a PersonalAgent app to many event types (reactions,
|
||||
# read receipts, recalls, member changes…) that EvoScientist doesn't
|
||||
# care about. Without this wrapper, ``_do_without_validation``
|
||||
# raises ``EventException("processor not found, type: ...")``,
|
||||
# which lark-oapi's WS client (ws/client.py) catches and turns into
|
||||
# an HTTP 500 reply on the WebSocket frame — Feishu then marks the
|
||||
# event as failed and retries it. This is especially noisy because
|
||||
# our own ``_send_ack_reaction`` triggers ``im.message.reaction.
|
||||
# created_v1`` on every inbound message, causing a feedback loop.
|
||||
from lark_oapi.core.exception import EventException
|
||||
|
||||
_original_dispatch = handler._do_without_validation
|
||||
|
||||
def _silent_dispatch(payload: bytes):
|
||||
try:
|
||||
return _original_dispatch(payload)
|
||||
except EventException as exc:
|
||||
if "processor not found" in str(exc):
|
||||
logger.debug("Feishu: ignored unsubscribed event (%s)", exc)
|
||||
return None
|
||||
raise
|
||||
|
||||
handler._do_without_validation = _silent_dispatch
|
||||
|
||||
ws_client = lark.ws.Client(
|
||||
self.config.app_id,
|
||||
self.config.app_secret,
|
||||
|
||||
@@ -0,0 +1,361 @@
|
||||
"""Feishu / Lark scan-to-create (QR code onboard) flow.
|
||||
|
||||
Drives the Feishu open-platform device-code flow at
|
||||
``accounts.feishu.cn/oauth/v1/app/registration`` (and the Lark equivalent
|
||||
at ``accounts.larksuite.com``). The user scans a terminal QR code with
|
||||
Feishu / Lark mobile, the platform provisions a ``PersonalAgent``-archetype
|
||||
bot application with the required IM permissions pre-attached, and the
|
||||
poll endpoint returns ``client_id`` / ``client_secret`` — enough to fully
|
||||
configure :class:`FeishuChannel`.
|
||||
|
||||
Domain auto-switches from ``feishu`` to ``lark`` if the poll response's
|
||||
``user_info.tenant_brand`` reports a Lark tenant.
|
||||
|
||||
The HTTP shape mirrors RFC 8628 (OAuth Device Authorization Grant) with
|
||||
a vendor-specific ``action`` form field selecting init / begin / poll.
|
||||
|
||||
Style follows :mod:`EvoScientist.channels.qq.onboard` — httpx, plain
|
||||
``print`` for progress, and an optional ``qrcode`` dependency for ASCII
|
||||
rendering.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_ACCOUNTS_URLS: dict[str, str] = {
|
||||
"feishu": "https://accounts.feishu.cn",
|
||||
"lark": "https://accounts.larksuite.com",
|
||||
}
|
||||
_OPEN_URLS: dict[str, str] = {
|
||||
"feishu": "https://open.feishu.cn",
|
||||
"lark": "https://open.larksuite.com",
|
||||
}
|
||||
_REGISTRATION_PATH = "/oauth/v1/app/registration"
|
||||
|
||||
_REQUEST_TIMEOUT_S = 10.0
|
||||
_DEFAULT_POLL_INTERVAL_S = 5
|
||||
_DEFAULT_EXPIRE_S = 600
|
||||
|
||||
|
||||
def _accounts_base_url(domain: str) -> str:
|
||||
return _ACCOUNTS_URLS.get(domain, _ACCOUNTS_URLS["feishu"])
|
||||
|
||||
|
||||
def _open_base_url(domain: str) -> str:
|
||||
return _OPEN_URLS.get(domain, _OPEN_URLS["feishu"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# QR rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
try:
|
||||
import qrcode as _qrcode_mod
|
||||
except (ImportError, TypeError):
|
||||
_qrcode_mod = None # type: ignore[assignment]
|
||||
|
||||
|
||||
def _render_qr(url: str) -> bool:
|
||||
"""Render *url* as an ASCII QR in the terminal. Returns True on success."""
|
||||
if _qrcode_mod is None:
|
||||
return False
|
||||
try:
|
||||
qr = _qrcode_mod.QRCode(
|
||||
error_correction=_qrcode_mod.constants.ERROR_CORRECT_M,
|
||||
border=2,
|
||||
)
|
||||
qr.add_data(url)
|
||||
qr.make(fit=True)
|
||||
qr.print_ascii(invert=True)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registration HTTP
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _post_registration(base_url: str, body: dict[str, str]) -> dict:
|
||||
"""POST form-encoded *body* to the registration endpoint.
|
||||
|
||||
The endpoint replies with JSON even on 4xx responses (``authorization_pending``
|
||||
comes back as HTTP 400 with a parseable body), so we always read the body
|
||||
and only fall back to raising if the bytes are missing or not JSON.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
url = f"{base_url}{_REGISTRATION_PATH}"
|
||||
headers = {"Content-Type": "application/x-www-form-urlencoded"}
|
||||
with httpx.Client(timeout=_REQUEST_TIMEOUT_S, follow_redirects=True) as client:
|
||||
resp = client.post(url, data=body, headers=headers)
|
||||
# Don't raise_for_status — 4xx may still carry a usable JSON body.
|
||||
try:
|
||||
return resp.json()
|
||||
except ValueError:
|
||||
resp.raise_for_status() # re-raise underlying HTTP error
|
||||
raise # pragma: no cover — raise_for_status already raised
|
||||
|
||||
|
||||
def _init_registration(domain: str) -> None:
|
||||
"""Probe the registration environment. Raises if client_secret auth is unavailable."""
|
||||
res = _post_registration(_accounts_base_url(domain), {"action": "init"})
|
||||
methods = res.get("supported_auth_methods") or []
|
||||
if "client_secret" not in methods:
|
||||
raise RuntimeError(
|
||||
f"Feishu / Lark registration environment does not support "
|
||||
f"client_secret auth (got: {methods})"
|
||||
)
|
||||
|
||||
|
||||
def _begin_registration(domain: str) -> dict:
|
||||
"""Start the device-code flow.
|
||||
|
||||
Returns a dict with ``device_code``, ``qr_url``, ``user_code``,
|
||||
``interval``, and ``expire_in``.
|
||||
"""
|
||||
res = _post_registration(
|
||||
_accounts_base_url(domain),
|
||||
{
|
||||
"action": "begin",
|
||||
"archetype": "PersonalAgent",
|
||||
"auth_method": "client_secret",
|
||||
"request_user_info": "open_id",
|
||||
},
|
||||
)
|
||||
device_code = res.get("device_code")
|
||||
if not device_code:
|
||||
raise RuntimeError(
|
||||
f"Feishu / Lark registration did not return a device_code: {res}"
|
||||
)
|
||||
qr_url = res.get("verification_uri_complete") or ""
|
||||
sep = "&" if "?" in qr_url else "?"
|
||||
qr_url = f"{qr_url}{sep}from=evoscientist&tp=evoscientist"
|
||||
return {
|
||||
"device_code": device_code,
|
||||
"qr_url": qr_url,
|
||||
"user_code": res.get("user_code", ""),
|
||||
"interval": int(res.get("interval") or _DEFAULT_POLL_INTERVAL_S),
|
||||
"expire_in": int(res.get("expire_in") or _DEFAULT_EXPIRE_S),
|
||||
}
|
||||
|
||||
|
||||
def _poll_registration(
|
||||
*,
|
||||
device_code: str,
|
||||
interval: int,
|
||||
expire_in: int,
|
||||
domain: str,
|
||||
) -> dict | None:
|
||||
"""Poll until the user scans, or the device_code expires / is denied.
|
||||
|
||||
Auto-switches the polling domain to ``lark`` if the server reports
|
||||
``user_info.tenant_brand == "lark"`` — the credentials only resolve
|
||||
against the matching open-platform host.
|
||||
|
||||
Returns a dict with ``app_id``, ``app_secret``, ``domain``, ``open_id``
|
||||
on success, or ``None`` on timeout / explicit denial.
|
||||
"""
|
||||
deadline = time.monotonic() + expire_in
|
||||
current_domain = domain
|
||||
domain_switched = False
|
||||
poll_count = 0
|
||||
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
res = _post_registration(
|
||||
_accounts_base_url(current_domain),
|
||||
{
|
||||
"action": "poll",
|
||||
"device_code": device_code,
|
||||
"tp": "ob_app",
|
||||
},
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("[Feishu onboard] poll request error: %s", exc)
|
||||
time.sleep(interval)
|
||||
continue
|
||||
|
||||
poll_count += 1
|
||||
if poll_count == 1:
|
||||
print(" Waiting for scan…", end="", flush=True)
|
||||
elif poll_count % 6 == 0:
|
||||
print(".", end="", flush=True)
|
||||
|
||||
# Domain auto-detection — the server may still return creds in
|
||||
# this same poll, so we fall through rather than restarting.
|
||||
user_info = res.get("user_info") or {}
|
||||
if (
|
||||
user_info.get("tenant_brand") == "lark"
|
||||
and not domain_switched
|
||||
and current_domain != "lark"
|
||||
):
|
||||
current_domain = "lark"
|
||||
domain_switched = True
|
||||
|
||||
if res.get("client_id") and res.get("client_secret"):
|
||||
print() # newline after the dots
|
||||
return {
|
||||
"app_id": res["client_id"],
|
||||
"app_secret": res["client_secret"],
|
||||
"domain": current_domain,
|
||||
"open_id": user_info.get("open_id"),
|
||||
}
|
||||
|
||||
error = res.get("error", "")
|
||||
if error in {"access_denied", "expired_token"}:
|
||||
print()
|
||||
logger.warning("[Feishu onboard] Registration %s", error)
|
||||
return None
|
||||
|
||||
# authorization_pending / slow_down / unknown — keep polling
|
||||
time.sleep(interval)
|
||||
|
||||
print()
|
||||
logger.warning("[Feishu onboard] Poll timed out after %ds", expire_in)
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bot probe (best-effort, uses tenant_access_token + /bot/v3/info)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _probe_bot(app_id: str, app_secret: str, domain: str) -> dict | None:
|
||||
"""Fetch bot name / bot_open_id via the open-platform REST API.
|
||||
|
||||
Best-effort: failures return ``None`` and the caller proceeds without
|
||||
a friendly bot name. Uses raw HTTP so we don't require ``lark-oapi``
|
||||
to be installed at onboard time (it's only needed for WebSocket mode).
|
||||
"""
|
||||
import httpx
|
||||
|
||||
base = _open_base_url(domain)
|
||||
token_url = f"{base}/open-apis/auth/v3/tenant_access_token/internal"
|
||||
info_url = f"{base}/open-apis/bot/v3/info"
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=_REQUEST_TIMEOUT_S, follow_redirects=True) as client:
|
||||
tok_resp = client.post(
|
||||
token_url,
|
||||
json={"app_id": app_id, "app_secret": app_secret},
|
||||
)
|
||||
tok_data = tok_resp.json()
|
||||
if tok_data.get("code") != 0:
|
||||
logger.debug("[Feishu onboard] token fetch failed: %s", tok_data)
|
||||
return None
|
||||
token = tok_data.get("tenant_access_token")
|
||||
if not token:
|
||||
return None
|
||||
|
||||
info_resp = client.get(
|
||||
info_url,
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
info_data = info_resp.json()
|
||||
except Exception as exc:
|
||||
logger.debug("[Feishu onboard] bot probe failed: %s", exc)
|
||||
return None
|
||||
|
||||
if info_data.get("code") != 0:
|
||||
return None
|
||||
bot = info_data.get("bot") or info_data.get("data", {}).get("bot") or {}
|
||||
return {
|
||||
"bot_name": bot.get("app_name") or bot.get("bot_name"),
|
||||
"bot_open_id": bot.get("open_id"),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public entry-point
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def qr_register(
|
||||
*,
|
||||
initial_domain: str = "feishu",
|
||||
timeout_seconds: int = 600,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Run the Feishu / Lark scan-to-create QR registration flow.
|
||||
|
||||
Args:
|
||||
initial_domain: ``"feishu"`` (default, mainland) or ``"lark"`` (overseas).
|
||||
Auto-switches mid-flow if the scanning user is on the other tenant.
|
||||
timeout_seconds: Wall-clock budget for the whole flow.
|
||||
|
||||
Returns on success::
|
||||
|
||||
{
|
||||
"app_id": str,
|
||||
"app_secret": str,
|
||||
"domain": "feishu" | "lark",
|
||||
"open_id": str | None,
|
||||
"bot_name": str | None,
|
||||
"bot_open_id": str | None,
|
||||
}
|
||||
|
||||
Returns ``None`` on expected failures (network, denial, timeout).
|
||||
"""
|
||||
try:
|
||||
return _qr_register_inner(
|
||||
initial_domain=initial_domain,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("[Feishu onboard] Registration failed: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def _qr_register_inner(
|
||||
*,
|
||||
initial_domain: str,
|
||||
timeout_seconds: int,
|
||||
) -> dict[str, Any] | None:
|
||||
print(" Connecting to Feishu / Lark…", end="", flush=True)
|
||||
_init_registration(initial_domain)
|
||||
begin = _begin_registration(initial_domain)
|
||||
print(" done.")
|
||||
|
||||
print()
|
||||
qr_url = begin["qr_url"]
|
||||
if _render_qr(qr_url):
|
||||
print(
|
||||
f"\n Scan the QR code above with Feishu / Lark on your phone,\n"
|
||||
f" or open this URL directly:\n {qr_url}"
|
||||
)
|
||||
else:
|
||||
print(f" Open this URL in Feishu / Lark on your phone:\n\n {qr_url}\n")
|
||||
print(
|
||||
" Tip: pip install qrcode to display a scannable QR code here next time"
|
||||
)
|
||||
print()
|
||||
|
||||
result = _poll_registration(
|
||||
device_code=begin["device_code"],
|
||||
interval=begin["interval"],
|
||||
expire_in=min(begin["expire_in"], timeout_seconds),
|
||||
domain=initial_domain,
|
||||
)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
bot_info = _probe_bot(result["app_id"], result["app_secret"], result["domain"])
|
||||
if bot_info:
|
||||
result["bot_name"] = bot_info.get("bot_name")
|
||||
result["bot_open_id"] = bot_info.get("bot_open_id")
|
||||
else:
|
||||
result["bot_name"] = None
|
||||
result["bot_open_id"] = None
|
||||
|
||||
return result
|
||||
@@ -20,11 +20,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import FeishuChannel, FeishuConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -92,7 +96,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = FeishuConfig(
|
||||
|
||||
@@ -19,7 +19,6 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from . import IMessageChannel, IMessageConfig
|
||||
@@ -68,7 +67,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = IMessageConfig(
|
||||
|
||||
@@ -75,11 +75,13 @@ class DedupCache:
|
||||
max_size: int = _DEDUP_MAX,
|
||||
trim_to: int = _DEDUP_TRIM,
|
||||
ttl_seconds: float = _DEDUP_TTL,
|
||||
clock: Callable[[], float] | None = None,
|
||||
) -> None:
|
||||
self._seen: OrderedDict[str, float] = OrderedDict()
|
||||
self._max = max_size
|
||||
self._trim = trim_to
|
||||
self._ttl = ttl_seconds
|
||||
self._clock = clock or time.monotonic
|
||||
|
||||
# ── public API ──────────────────────────────────────────────────
|
||||
|
||||
@@ -93,15 +95,16 @@ class DedupCache:
|
||||
if not msg_id:
|
||||
return False
|
||||
|
||||
self._prune()
|
||||
now = self._clock()
|
||||
self._prune(now)
|
||||
|
||||
if msg_id in self._seen:
|
||||
# LRU: refresh position and timestamp
|
||||
self._seen.move_to_end(msg_id)
|
||||
self._seen[msg_id] = time.monotonic()
|
||||
self._seen[msg_id] = now
|
||||
return True
|
||||
|
||||
self._seen[msg_id] = time.monotonic()
|
||||
self._seen[msg_id] = now
|
||||
if len(self._seen) > self._max:
|
||||
while len(self._seen) > self._trim:
|
||||
self._seen.popitem(last=False)
|
||||
@@ -118,9 +121,9 @@ class DedupCache:
|
||||
|
||||
# ── internal ────────────────────────────────────────────────────
|
||||
|
||||
def _prune(self) -> None:
|
||||
def _prune(self, now: float | None = None) -> None:
|
||||
"""Remove entries older than *ttl_seconds*."""
|
||||
cutoff = time.monotonic() - self._ttl
|
||||
cutoff = (self._clock() if now is None else now) - self._ttl
|
||||
# OrderedDict is insertion-ordered; oldest entries are first.
|
||||
while self._seen:
|
||||
_key, ts = next(iter(self._seen.items()))
|
||||
@@ -428,11 +431,13 @@ class DedupMiddleware(InboundMiddleware):
|
||||
max_size: int = 1000,
|
||||
trim_to: int = 500,
|
||||
ttl_seconds: float = 3600.0,
|
||||
clock: Callable[[], float] | None = None,
|
||||
) -> None:
|
||||
self._cache = DedupCache(
|
||||
max_size=max_size,
|
||||
trim_to=trim_to,
|
||||
ttl_seconds=ttl_seconds,
|
||||
clock=clock,
|
||||
)
|
||||
|
||||
async def process_inbound(
|
||||
|
||||
@@ -10,8 +10,9 @@ Usage in config:
|
||||
|
||||
from ..channel_manager import _parse_csv, register_channel
|
||||
from .channel import QQChannel, QQConfig
|
||||
from .onboard import qr_register
|
||||
|
||||
__all__ = ["QQChannel", "QQConfig"]
|
||||
__all__ = ["QQChannel", "QQConfig", "qr_register"]
|
||||
|
||||
|
||||
def create_from_config(config) -> QQChannel:
|
||||
|
||||
@@ -26,6 +26,55 @@ except ImportError:
|
||||
GroupMessage = None
|
||||
|
||||
|
||||
# ── Inline keyboard (button) helpers ─────────────────────────────────
|
||||
|
||||
|
||||
def _normalize_button(btn: dict) -> tuple[str, str] | None:
|
||||
"""Return ``(label, value)`` for a button, or ``None`` if no label."""
|
||||
label = (btn.get("text") or "").strip()
|
||||
if not label:
|
||||
return None
|
||||
raw = btn.get("value")
|
||||
return label, str(raw) if raw is not None else label
|
||||
|
||||
|
||||
def _build_qq_keyboard(buttons: list[dict]) -> dict | None:
|
||||
"""Build a QQ Bot keyboard payload (one button per row).
|
||||
|
||||
Render style: 1 = primary (blue), 0 = secondary (grey) — QQ has no danger.
|
||||
``action.permission`` is required by the schema; ``type=2`` is harmless for
|
||||
C2C (the click always comes from the DM peer). Returns ``None`` if no
|
||||
button has a usable label.
|
||||
"""
|
||||
rows: list[dict] = []
|
||||
for idx, btn in enumerate(buttons):
|
||||
norm = _normalize_button(btn)
|
||||
if norm is None:
|
||||
continue
|
||||
label, value = norm
|
||||
style = 1 if btn.get("type") == "primary" else 0
|
||||
rows.append(
|
||||
{
|
||||
"buttons": [
|
||||
{
|
||||
"id": btn.get("id") or f"btn_{idx}",
|
||||
"render_data": {
|
||||
"label": label,
|
||||
"visited_label": label,
|
||||
"style": style,
|
||||
},
|
||||
"action": {
|
||||
"type": 1, # callback (server pushes interaction event)
|
||||
"permission": {"type": 2},
|
||||
"data": value,
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
return {"content": {"rows": rows}} if rows else None
|
||||
|
||||
|
||||
@dataclass
|
||||
class QQConfig(BaseChannelConfig):
|
||||
app_id: str = ""
|
||||
@@ -35,7 +84,11 @@ class QQConfig(BaseChannelConfig):
|
||||
|
||||
def _make_bot_class(channel: "QQChannel") -> "type[botpy.Client]":
|
||||
"""Create a botpy Client subclass bound to the given channel."""
|
||||
intents = botpy.Intents(public_messages=True, direct_message=True)
|
||||
intents = botpy.Intents(
|
||||
public_messages=True,
|
||||
direct_message=True,
|
||||
interaction=True, # button clicks → on_interaction_create
|
||||
)
|
||||
|
||||
class _Bot(botpy.Client):
|
||||
def __init__(self):
|
||||
@@ -50,6 +103,9 @@ def _make_bot_class(channel: "QQChannel") -> "type[botpy.Client]":
|
||||
async def on_group_at_message_create(self, message: "GroupMessage"):
|
||||
await channel._on_msg(message, "group")
|
||||
|
||||
async def on_interaction_create(self, interaction):
|
||||
await channel._on_interaction(interaction)
|
||||
|
||||
return _Bot
|
||||
|
||||
|
||||
@@ -174,6 +230,76 @@ class QQChannel(Channel):
|
||||
except Exception as e:
|
||||
logger.error(f"Error handling QQ message: {e}")
|
||||
|
||||
async def _on_interaction(self, interaction) -> None:
|
||||
"""Handle ``on_interaction_create`` (button click).
|
||||
|
||||
Surfaces the click as an :class:`InboundMessage` whose ``content`` is
|
||||
the button's ``data`` verbatim — so a "1"/"approve"/… click flows
|
||||
through ``_parse_approval_reply`` exactly like a typed reply.
|
||||
|
||||
The click runs through inbound middleware (Dedup suppresses QQ
|
||||
retries) but is published directly to the bus so the per-sender
|
||||
debounce buffer doesn't merge the click value with subsequent text.
|
||||
Group-scope clicks are ignored (DM-only by design).
|
||||
"""
|
||||
# ACK first — QQ requires a response within ~5s or the button UI
|
||||
# shows "expired". Code 0 just means "received"; downstream still
|
||||
# decides the actual approval/rejection.
|
||||
interaction_id = getattr(interaction, "id", "") or ""
|
||||
if interaction_id and self._client:
|
||||
try:
|
||||
await self._client.api.on_interaction_result(interaction_id, 0)
|
||||
except Exception as ack_exc:
|
||||
logger.debug("QQ interaction ack failed: %s", ack_exc)
|
||||
|
||||
try:
|
||||
user_openid = getattr(interaction, "user_openid", "") or ""
|
||||
if not user_openid:
|
||||
logger.debug("QQ interaction ignored (no user_openid; not C2C)")
|
||||
return
|
||||
|
||||
resolved = getattr(getattr(interaction, "data", None), "resolved", None)
|
||||
button_data = getattr(resolved, "button_data", "") or ""
|
||||
button_id = getattr(resolved, "button_id", "") or ""
|
||||
triggering_msg_id = getattr(resolved, "message_id", "") or ""
|
||||
|
||||
# QQ may serialize non-str values; coerce. Fall back to button id
|
||||
# when no data — same path as a typed reply via _parse_approval_reply.
|
||||
button_value = str(button_data) if button_data != "" else ""
|
||||
text = button_value or button_id
|
||||
|
||||
# Stable id so DedupMiddleware suppresses any QQ retry callbacks.
|
||||
message_id = (
|
||||
f"{triggering_msg_id}:action:{interaction_id}"
|
||||
if interaction_id
|
||||
else f"qq_action:{datetime.now().timestamp()}"
|
||||
)
|
||||
|
||||
raw = RawIncoming(
|
||||
sender_id=user_openid,
|
||||
chat_id=user_openid, # C2C: chat_id == user_openid
|
||||
text=text,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": user_openid,
|
||||
"msg_type": "c2c",
|
||||
"event_id": triggering_msg_id,
|
||||
"backend": "qq",
|
||||
"button_click": True,
|
||||
"button_id": button_id,
|
||||
"button_value": button_value,
|
||||
},
|
||||
is_group=False,
|
||||
was_mentioned=True,
|
||||
)
|
||||
|
||||
inbound = await self._build_inbound_async(raw)
|
||||
if inbound is not None and self._bus:
|
||||
await self._bus.publish_inbound(inbound)
|
||||
except Exception:
|
||||
logger.exception("QQ interaction handler error")
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
def _next_msg_seq(self, msg_id: str) -> int:
|
||||
@@ -195,17 +321,81 @@ class QQChannel(Channel):
|
||||
msg_type = (metadata or {}).get("msg_type", "c2c")
|
||||
msg_id = (metadata or {}).get("event_id", "")
|
||||
seq = self._next_msg_seq(msg_id)
|
||||
|
||||
# Inline keyboard is C2C-only here — group keyboards have stricter
|
||||
# permission semantics and are out of scope for now.
|
||||
buttons = (metadata or {}).get("buttons") if msg_type == "c2c" else None
|
||||
keyboard = _build_qq_keyboard(buttons) if buttons else None
|
||||
|
||||
try:
|
||||
await self._post_markdown_message(chat_id, raw_text, msg_type, msg_id, seq)
|
||||
await self._post_markdown_message(
|
||||
chat_id, raw_text, msg_type, msg_id, seq, keyboard=keyboard
|
||||
)
|
||||
return
|
||||
except Exception as exc:
|
||||
if not self._should_fallback_to_plain_text(exc):
|
||||
logger.error(
|
||||
"QQ markdown send failed with non-fallbackable error "
|
||||
"(chat_id=%s, msg_id=%s, seq=%s): %r",
|
||||
chat_id,
|
||||
msg_id,
|
||||
seq,
|
||||
exc,
|
||||
)
|
||||
raise
|
||||
self._record_markdown_fallback(chat_id, raw_text, exc)
|
||||
logger.debug("QQ markdown send failed, falling back to plain text: %s", exc)
|
||||
logger.warning(
|
||||
"QQ markdown send failed, falling back to plain text "
|
||||
"(chat_id=%s, msg_id=%s, seq=%s): %r",
|
||||
chat_id,
|
||||
msg_id,
|
||||
seq,
|
||||
exc,
|
||||
)
|
||||
|
||||
# QQ may have already consumed `seq` server-side even on failure.
|
||||
# Reusing it for the plain retry triggers "duplicate msg_seq", so
|
||||
# always advance to a fresh seq before the fallback send.
|
||||
fallback_seq = self._next_msg_seq(msg_id)
|
||||
plain_text = self._plain_formatter.format(raw_text)
|
||||
await self._post_plain_message(chat_id, plain_text, msg_type, msg_id, seq)
|
||||
# Plain-text fallback can't carry a keyboard. Append `value=label`
|
||||
# pairs so the user can still type "1"/"approve"/… instead of
|
||||
# tapping (`_parse_approval_reply` accepts the same values).
|
||||
if buttons:
|
||||
pairs = []
|
||||
for btn in buttons:
|
||||
norm = _normalize_button(btn)
|
||||
if norm is not None:
|
||||
label, value = norm
|
||||
pairs.append(f"{value}={label}")
|
||||
if pairs:
|
||||
plain_text = f"{plain_text}\n\nReply: {', '.join(pairs)}"
|
||||
try:
|
||||
await self._post_plain_message(
|
||||
chat_id, plain_text, msg_type, msg_id, fallback_seq
|
||||
)
|
||||
except Exception as plain_exc:
|
||||
logger.error(
|
||||
"QQ plain fallback also failed (chat_id=%s, msg_id=%s, seq=%s): %r",
|
||||
chat_id,
|
||||
msg_id,
|
||||
fallback_seq,
|
||||
plain_exc,
|
||||
)
|
||||
raise
|
||||
|
||||
# QQ server-side error codes / fragments that indicate the markdown
|
||||
# request itself is invalid (template not configured, format rejected,
|
||||
# content audit, etc.). Seeing any of these means we should retry with
|
||||
# plain text rather than re-raise.
|
||||
_QQ_MARKDOWN_ERROR_MARKERS: ClassVar[tuple[str, ...]] = (
|
||||
"304014", # markdown template not configured
|
||||
"304003", # invalid markdown params
|
||||
"40034059", # generic send message failed (often markdown-related)
|
||||
"模板", # CN: template (standard form)
|
||||
"模版", # CN: template (variant form)
|
||||
"审核", # CN: audit
|
||||
)
|
||||
|
||||
def _should_fallback_to_plain_text(self, exc: Exception) -> bool:
|
||||
"""Return True only for markdown compatibility/validation failures."""
|
||||
@@ -214,17 +404,20 @@ class QQChannel(Channel):
|
||||
|
||||
msg = str(exc).lower()
|
||||
compatibility_tokens = ("unsupported", "unexpected", "unknown", "invalid")
|
||||
return (
|
||||
"unexpected keyword argument" in msg
|
||||
or (
|
||||
"markdown" in msg
|
||||
and any(token in msg for token in compatibility_tokens)
|
||||
)
|
||||
or (
|
||||
"msg_type" in msg
|
||||
and any(token in msg for token in compatibility_tokens)
|
||||
)
|
||||
)
|
||||
|
||||
if "unexpected keyword argument" in msg:
|
||||
return True
|
||||
if "markdown" in msg and any(token in msg for token in compatibility_tokens):
|
||||
return True
|
||||
if "msg_type" in msg and any(token in msg for token in compatibility_tokens):
|
||||
return True
|
||||
|
||||
# QQ-specific server error codes returned by qq-botpy as strings.
|
||||
raw = str(exc)
|
||||
for marker in self._QQ_MARKDOWN_ERROR_MARKERS:
|
||||
if marker in raw or marker.lower() in msg:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _record_markdown_fallback(
|
||||
self,
|
||||
@@ -254,6 +447,7 @@ class QQChannel(Channel):
|
||||
msg_type: str,
|
||||
msg_id: str,
|
||||
seq: int,
|
||||
keyboard: dict | None = None,
|
||||
) -> None:
|
||||
payload = {
|
||||
"msg_type": 2,
|
||||
@@ -261,6 +455,8 @@ class QQChannel(Channel):
|
||||
"msg_id": msg_id,
|
||||
"msg_seq": seq,
|
||||
}
|
||||
if keyboard is not None:
|
||||
payload["keyboard"] = keyboard
|
||||
if msg_type == "group":
|
||||
await self._client.api.post_group_message(
|
||||
group_openid=chat_id,
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
"""AES-256-GCM utilities for QQ Bot scan-to-configure credential decryption.
|
||||
|
||||
Ported from hermes-agent/gateway/platforms/qqbot/crypto.py — the q.qq.com
|
||||
``create_bind_task`` / ``poll_bind_result`` flow uses AES-256-GCM to keep
|
||||
the bot's *client_secret* off the wire in plaintext.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
|
||||
|
||||
def generate_bind_key() -> str:
|
||||
"""Generate a 256-bit random AES key, base64-encoded.
|
||||
|
||||
The key is sent to ``create_bind_task`` so the server can encrypt
|
||||
the bot's *client_secret* before returning it. Only this client
|
||||
holds the key, so the secret never travels in plaintext.
|
||||
"""
|
||||
return base64.b64encode(os.urandom(32)).decode()
|
||||
|
||||
|
||||
def decrypt_secret(encrypted_base64: str, key_base64: str) -> str:
|
||||
"""Decrypt a base64-encoded AES-256-GCM ciphertext.
|
||||
|
||||
Ciphertext layout (after base64-decoding)::
|
||||
|
||||
IV (12 bytes) ‖ ciphertext (N bytes) ‖ AuthTag (16 bytes)
|
||||
|
||||
Args:
|
||||
encrypted_base64: The ``bot_encrypt_secret`` value returned by
|
||||
``poll_bind_result``.
|
||||
key_base64: The base64 AES key produced by :func:`generate_bind_key`.
|
||||
|
||||
Returns:
|
||||
The decrypted *client_secret* as a UTF-8 string.
|
||||
"""
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
key = base64.b64decode(key_base64)
|
||||
raw = base64.b64decode(encrypted_base64)
|
||||
|
||||
iv = raw[:12]
|
||||
ciphertext_with_tag = raw[12:] # AESGCM expects ciphertext + tag concatenated
|
||||
|
||||
aesgcm = AESGCM(key)
|
||||
plaintext = aesgcm.decrypt(iv, ciphertext_with_tag, None)
|
||||
return plaintext.decode("utf-8")
|
||||
@@ -0,0 +1,300 @@
|
||||
"""QQ Bot scan-to-configure (QR code onboard) flow.
|
||||
|
||||
Ported from hermes-agent/gateway/platforms/qqbot/onboard.py.
|
||||
|
||||
Calls the ``q.qq.com`` ``create_bind_task`` / ``poll_bind_result`` APIs to
|
||||
generate a QR code URL and poll for scan completion. On success the caller
|
||||
receives the bot's *app_id*, *client_secret* (decrypted locally), and the
|
||||
scanner's *user_openid* — enough to fully configure the QQ channel.
|
||||
|
||||
The bot must already be registered at https://q.qq.com — scanning binds
|
||||
the QQ user (developer / admin) to the existing application; it does not
|
||||
create a new one.
|
||||
|
||||
Reference: https://bot.q.qq.com/wiki/develop/api-v2/
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import sys
|
||||
import time
|
||||
from enum import IntEnum
|
||||
from urllib.parse import quote
|
||||
|
||||
from .crypto import decrypt_secret, generate_bind_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints / timing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# The portal domain is configurable for corporate proxies / sandbox routing.
|
||||
PORTAL_HOST = os.getenv("QQ_PORTAL_HOST", "q.qq.com")
|
||||
|
||||
ONBOARD_CREATE_PATH = "/lite/create_bind_task"
|
||||
ONBOARD_POLL_PATH = "/lite/poll_bind_result"
|
||||
QR_URL_TEMPLATE = (
|
||||
"https://q.qq.com/qqbot/openclaw/connect.html"
|
||||
"?task_id={task_id}&_wv=2&source=evoscientist"
|
||||
)
|
||||
|
||||
ONBOARD_API_TIMEOUT = 10.0
|
||||
ONBOARD_POLL_INTERVAL = 2.0
|
||||
|
||||
_MAX_REFRESHES = 3
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bind status
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BindStatus(IntEnum):
|
||||
"""Status codes returned by ``poll_bind_result``."""
|
||||
|
||||
NONE = 0
|
||||
PENDING = 1
|
||||
COMPLETED = 2
|
||||
EXPIRED = 3
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HTTP headers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_evoscientist_version() -> str:
|
||||
try:
|
||||
from importlib.metadata import version
|
||||
|
||||
return version("evoscientist")
|
||||
except Exception:
|
||||
return "dev"
|
||||
|
||||
|
||||
def _build_user_agent() -> str:
|
||||
py_version = (
|
||||
f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
|
||||
)
|
||||
os_name = platform.system().lower()
|
||||
return (
|
||||
f"EvoScientistQQ/1.0.0 (Python/{py_version}; {os_name}; "
|
||||
f"EvoScientist/{_get_evoscientist_version()})"
|
||||
)
|
||||
|
||||
|
||||
def _api_headers() -> dict[str, str]:
|
||||
"""Standard HTTP headers for q.qq.com onboard API requests.
|
||||
|
||||
``q.qq.com`` requires ``Accept: application/json`` — without it,
|
||||
the server returns a JavaScript anti-bot challenge page.
|
||||
"""
|
||||
return {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
"User-Agent": _build_user_agent(),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# QR rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
try:
|
||||
import qrcode as _qrcode_mod
|
||||
except (ImportError, TypeError):
|
||||
_qrcode_mod = None # type: ignore[assignment]
|
||||
|
||||
|
||||
def _render_qr(url: str) -> bool:
|
||||
"""Render a QR code to the terminal. Returns True on success."""
|
||||
if _qrcode_mod is None:
|
||||
return False
|
||||
try:
|
||||
qr = _qrcode_mod.QRCode(
|
||||
error_correction=_qrcode_mod.constants.ERROR_CORRECT_M,
|
||||
border=2,
|
||||
)
|
||||
qr.add_data(url)
|
||||
qr.make(fit=True)
|
||||
qr.print_ascii(invert=True)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HTTP helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _create_bind_task(timeout: float = ONBOARD_API_TIMEOUT) -> tuple[str, str]:
|
||||
"""Create a bind task and return *(task_id, aes_key_base64)*.
|
||||
|
||||
Raises:
|
||||
RuntimeError: if the API returns a non-zero ``retcode``.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
url = f"https://{PORTAL_HOST}{ONBOARD_CREATE_PATH}"
|
||||
key = generate_bind_key()
|
||||
|
||||
with httpx.Client(timeout=timeout, follow_redirects=True) as client:
|
||||
resp = client.post(url, json={"key": key}, headers=_api_headers())
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
if data.get("retcode") != 0:
|
||||
raise RuntimeError(data.get("msg", "create_bind_task failed"))
|
||||
|
||||
task_id = data.get("data", {}).get("task_id")
|
||||
if not task_id:
|
||||
raise RuntimeError("create_bind_task: missing task_id in response")
|
||||
|
||||
logger.debug("create_bind_task ok: task_id=%s", task_id)
|
||||
return task_id, key
|
||||
|
||||
|
||||
def _poll_bind_result(
|
||||
task_id: str,
|
||||
timeout: float = ONBOARD_API_TIMEOUT,
|
||||
) -> tuple[BindStatus, str, str, str]:
|
||||
"""Poll the bind result for *task_id*.
|
||||
|
||||
Returns:
|
||||
``(status, bot_appid, bot_encrypt_secret, user_openid)``.
|
||||
|
||||
Raises:
|
||||
RuntimeError: if the API returns a non-zero ``retcode``.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
url = f"https://{PORTAL_HOST}{ONBOARD_POLL_PATH}"
|
||||
|
||||
with httpx.Client(timeout=timeout, follow_redirects=True) as client:
|
||||
resp = client.post(url, json={"task_id": task_id}, headers=_api_headers())
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
if data.get("retcode") != 0:
|
||||
raise RuntimeError(data.get("msg", "poll_bind_result failed"))
|
||||
|
||||
d = data.get("data", {})
|
||||
return (
|
||||
BindStatus(d.get("status", 0)),
|
||||
str(d.get("bot_appid", "")),
|
||||
d.get("bot_encrypt_secret", ""),
|
||||
d.get("user_openid", ""),
|
||||
)
|
||||
|
||||
|
||||
def build_connect_url(task_id: str) -> str:
|
||||
"""Build the QR-code target URL for a given *task_id*."""
|
||||
return QR_URL_TEMPLATE.format(task_id=quote(task_id))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public entry-point
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def qr_register(timeout_seconds: int = 600) -> dict | None:
|
||||
"""Run the QQ Bot scan-to-configure QR registration flow.
|
||||
|
||||
Handles create → display → poll → decrypt in one call. The QR
|
||||
auto-refreshes up to ``_MAX_REFRESHES`` times if the user takes
|
||||
too long to scan.
|
||||
|
||||
Args:
|
||||
timeout_seconds: Total wall-clock budget across all refreshes.
|
||||
|
||||
Returns:
|
||||
``{"app_id": ..., "client_secret": ..., "user_openid": ...}`` on
|
||||
success, or ``None`` on failure / expiry / cancellation.
|
||||
"""
|
||||
deadline = time.monotonic() + timeout_seconds
|
||||
|
||||
for refresh_count in range(_MAX_REFRESHES + 1):
|
||||
# ── Create bind task ──
|
||||
try:
|
||||
task_id, aes_key = _create_bind_task()
|
||||
except Exception as exc:
|
||||
logger.warning("[QQ onboard] Failed to create bind task: %s", exc)
|
||||
return None
|
||||
|
||||
url = build_connect_url(task_id)
|
||||
|
||||
# ── Display QR code + URL ──
|
||||
print()
|
||||
if _render_qr(url):
|
||||
print(f" Scan the QR code above, or open this URL on your phone:\n {url}")
|
||||
else:
|
||||
print(f" Open this URL in QQ on your phone:\n {url}")
|
||||
print(" Tip: pip install qrcode to display a scannable QR code here")
|
||||
print()
|
||||
|
||||
# ── Poll loop ──
|
||||
consecutive_errors = 0
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
status, app_id, encrypted_secret, user_openid = _poll_bind_result(
|
||||
task_id
|
||||
)
|
||||
except Exception as exc:
|
||||
consecutive_errors += 1
|
||||
logger.warning(
|
||||
"[QQ onboard] poll_bind_result failed (%d consecutive): %s",
|
||||
consecutive_errors,
|
||||
exc,
|
||||
)
|
||||
if consecutive_errors >= 5:
|
||||
print(
|
||||
"\n Repeated polling failures — aborting."
|
||||
" See logs for details."
|
||||
)
|
||||
return None
|
||||
time.sleep(ONBOARD_POLL_INTERVAL)
|
||||
continue
|
||||
consecutive_errors = 0
|
||||
|
||||
if status == BindStatus.COMPLETED:
|
||||
try:
|
||||
client_secret = decrypt_secret(encrypted_secret, aes_key)
|
||||
except Exception as exc:
|
||||
logger.warning("[QQ onboard] decrypt_secret failed: %s", exc)
|
||||
return None
|
||||
print()
|
||||
print(f" QR scan complete! (App ID: {app_id})")
|
||||
if user_openid:
|
||||
print(f" Scanner's OpenID: {user_openid}")
|
||||
return {
|
||||
"app_id": app_id,
|
||||
"client_secret": client_secret,
|
||||
"user_openid": user_openid,
|
||||
}
|
||||
|
||||
if status == BindStatus.EXPIRED:
|
||||
if refresh_count >= _MAX_REFRESHES:
|
||||
logger.warning(
|
||||
"[QQ onboard] QR code expired %d times — giving up",
|
||||
_MAX_REFRESHES,
|
||||
)
|
||||
return None
|
||||
print(
|
||||
f"\n QR code expired, refreshing... "
|
||||
f"({refresh_count + 1}/{_MAX_REFRESHES})"
|
||||
)
|
||||
break # next outer iteration creates a new task
|
||||
|
||||
time.sleep(ONBOARD_POLL_INTERVAL)
|
||||
else:
|
||||
# deadline reached without completing
|
||||
logger.warning("[QQ onboard] Poll timed out after %ds", timeout_seconds)
|
||||
return None
|
||||
|
||||
return None
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import QQChannel, QQConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -64,7 +68,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = QQConfig(
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import SignalChannel, SignalConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -74,7 +78,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = SignalConfig(
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import SlackChannel, SlackConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -74,7 +78,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = SlackConfig(
|
||||
|
||||
@@ -108,8 +108,10 @@ async def _async_main(
|
||||
if use_agent:
|
||||
logger.info("Loading EvoScientist agent...")
|
||||
from ..EvoScientist import create_cli_agent
|
||||
from ..gateway import create_runtime_gateways
|
||||
|
||||
agent = create_cli_agent()
|
||||
runtime_gateways = create_runtime_gateways()
|
||||
logger.info("Agent loaded")
|
||||
|
||||
consumer = InboundConsumer(
|
||||
@@ -117,6 +119,7 @@ async def _async_main(
|
||||
manager=manager,
|
||||
agent=agent,
|
||||
thread_id="",
|
||||
graph_gateway=runtime_gateways.graph_gateway,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
manager.register_health_provider("consumer", lambda: consumer.metrics)
|
||||
|
||||
@@ -19,11 +19,15 @@ Examples:
|
||||
import argparse
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import TelegramChannel, TelegramConfig
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -59,7 +63,6 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
config = TelegramConfig(
|
||||
|
||||
@@ -5,13 +5,16 @@ Supports multiple WeChat backends:
|
||||
— Most stable, pure HTTP, no third-party dependencies
|
||||
- **wechatmp**: 微信公众号 (WeChat Official Account) via official API
|
||||
— Pure HTTP webhook, suitable for public-facing bots
|
||||
- **personal**: 个人微信 via Tencent's iLink Bot API
|
||||
— Long-poll + AES-128-ECB CDN media protocol; QR-code login required.
|
||||
Adapted from hermes-agent.
|
||||
|
||||
Both backends use httpx (already a core dependency) and receive messages
|
||||
via HTTP webhook, send replies via REST API.
|
||||
Backends 1+2 share the HTTP-webhook ``WeChatChannel``; backend 3 uses the
|
||||
long-poll ``WeixinPersonalChannel``.
|
||||
|
||||
Usage in config:
|
||||
channel_enabled = "wechat"
|
||||
wechat_backend = "wecom" # or "wechatmp"
|
||||
wechat_backend = "wecom" # or "wechatmp" or "personal"
|
||||
|
||||
# WeCom settings
|
||||
wechat_wecom_corp_id = "..."
|
||||
@@ -27,19 +30,54 @@ Usage in config:
|
||||
wechat_mp_token = "..."
|
||||
wechat_mp_encoding_aes_key = "..."
|
||||
wechat_webhook_port = 9001
|
||||
|
||||
# OR: Personal WeChat (iLink Bot)
|
||||
# First run `python -m EvoScientist.channels.wechat.serve --qr-login`
|
||||
# to obtain an account_id + token via QR-code scan.
|
||||
wechat_personal_account_id = "..."
|
||||
wechat_personal_token = "..." # optional if persisted on disk
|
||||
wechat_personal_dm_policy = "open" # open | allowlist
|
||||
wechat_personal_group_policy = "disabled"
|
||||
"""
|
||||
|
||||
from ..channel_manager import _parse_csv, register_channel
|
||||
from .channel import WeChatChannel, WeChatMPConfig, WeComConfig
|
||||
from .personal import WeixinPersonalChannel, WeixinPersonalConfig, qr_login
|
||||
|
||||
__all__ = ["WeChatChannel", "WeChatMPConfig", "WeComConfig"]
|
||||
__all__ = [
|
||||
"WeChatChannel",
|
||||
"WeChatMPConfig",
|
||||
"WeComConfig",
|
||||
"WeixinPersonalChannel",
|
||||
"WeixinPersonalConfig",
|
||||
"qr_login",
|
||||
]
|
||||
|
||||
|
||||
def create_from_config(config) -> WeChatChannel:
|
||||
backend = config.wechat_backend or "wecom"
|
||||
allowed = _parse_csv(config.wechat_allowed_senders)
|
||||
proxy = config.wechat_proxy or None
|
||||
port = int(config.wechat_webhook_port or 9001)
|
||||
def create_from_config(config):
|
||||
"""Factory dispatched on ``config.wechat_backend``."""
|
||||
backend = (getattr(config, "wechat_backend", "") or "wecom").lower()
|
||||
allowed = _parse_csv(getattr(config, "wechat_allowed_senders", ""))
|
||||
proxy = getattr(config, "wechat_proxy", "") or None
|
||||
port = int(getattr(config, "wechat_webhook_port", 9001) or 9001)
|
||||
|
||||
if backend == "personal":
|
||||
group_allowed = _parse_csv(getattr(config, "wechat_personal_group_allowed", ""))
|
||||
cfg = WeixinPersonalConfig(
|
||||
account_id=getattr(config, "wechat_personal_account_id", ""),
|
||||
token=getattr(config, "wechat_personal_token", ""),
|
||||
base_url=getattr(config, "wechat_personal_base_url", "")
|
||||
or "https://ilinkai.weixin.qq.com",
|
||||
cdn_base_url=getattr(config, "wechat_personal_cdn_base_url", "")
|
||||
or "https://novac2c.cdn.weixin.qq.com/c2c",
|
||||
dm_policy=getattr(config, "wechat_personal_dm_policy", "open") or "open",
|
||||
group_policy=getattr(config, "wechat_personal_group_policy", "disabled")
|
||||
or "disabled",
|
||||
group_allowed_senders=group_allowed,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
return WeixinPersonalChannel(cfg)
|
||||
|
||||
if backend == "wechatmp":
|
||||
mp_config = WeChatMPConfig(
|
||||
@@ -52,7 +90,7 @@ def create_from_config(config) -> WeChatChannel:
|
||||
proxy=proxy,
|
||||
)
|
||||
return WeChatChannel(mp_config, backend="wechatmp")
|
||||
else:
|
||||
|
||||
wecom_config = WeComConfig(
|
||||
corp_id=config.wechat_wecom_corp_id,
|
||||
agent_id=config.wechat_wecom_agent_id,
|
||||
|
||||
@@ -83,6 +83,53 @@ def _aes_encrypt(key: bytes, iv: bytes, plaintext: bytes) -> bytes:
|
||||
) from None
|
||||
|
||||
|
||||
def aes128_ecb_decrypt(ciphertext: bytes, key: bytes) -> bytes:
|
||||
"""AES-128-ECB decryption with PKCS#7 unpadding.
|
||||
|
||||
Used by the personal-WeChat (iLink) backend for CDN-encrypted media
|
||||
payloads. Block size is 16; the WeChat CDN protocol pads with PKCS#7.
|
||||
"""
|
||||
if _HAS_PYCRYPTO:
|
||||
cipher = AES.new(key, AES.MODE_ECB)
|
||||
padded = cipher.decrypt(ciphertext)
|
||||
else:
|
||||
try:
|
||||
import pyaes
|
||||
|
||||
decrypter = pyaes.Decrypter(pyaes.AESModeOfOperationECB(key))
|
||||
padded = decrypter.feed(ciphertext)
|
||||
padded += decrypter.feed()
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"WeChat CDN media decryption requires pycryptodome or pyaes. "
|
||||
"Install with: pip install pycryptodome"
|
||||
) from None
|
||||
|
||||
if not padded:
|
||||
return padded
|
||||
pad_len = padded[-1]
|
||||
if 1 <= pad_len <= 16 and padded.endswith(bytes([pad_len]) * pad_len):
|
||||
return padded[:-pad_len]
|
||||
return padded
|
||||
|
||||
|
||||
def parse_ilink_aes_key(aes_key_b64: str) -> bytes:
|
||||
"""Parse the iLink CDN AES key.
|
||||
|
||||
iLink encodes the 16-byte key in two formats:
|
||||
- direct base64 of 16 bytes
|
||||
- base64 of a 32-char ASCII hex string (which decodes to 16 raw bytes)
|
||||
"""
|
||||
decoded = base64.b64decode(aes_key_b64)
|
||||
if len(decoded) == 16:
|
||||
return decoded
|
||||
if len(decoded) == 32:
|
||||
text = decoded.decode("ascii", errors="ignore")
|
||||
if text and all(ch in "0123456789abcdefABCDEF" for ch in text):
|
||||
return bytes.fromhex(text)
|
||||
raise ValueError(f"unexpected aes_key format ({len(decoded)} decoded bytes)")
|
||||
|
||||
|
||||
class WeChatCrypto:
|
||||
"""Handles WeChat/WeCom message encryption and decryption.
|
||||
|
||||
|
||||
@@ -70,3 +70,40 @@ async def validate_wechat_mp(
|
||||
return False, f"Error ({data.get('errcode')}): {data.get('errmsg')}"
|
||||
except Exception as e:
|
||||
return False, f"Error: {e}"
|
||||
|
||||
|
||||
async def validate_wechat_personal(
|
||||
account_id: str,
|
||||
token: str = "",
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate that a personal WeChat (iLink) account has been logged in.
|
||||
|
||||
Personal-WeChat credentials are obtained via QR-code scan and persisted
|
||||
on disk; there is no offline credential format the user can paste in.
|
||||
This probe simply checks that:
|
||||
|
||||
- the account_id is set, and
|
||||
- either *token* is supplied inline, or a saved-account file exists for
|
||||
*account_id* under ``DATA_DIR/wechat_personal/accounts/``.
|
||||
|
||||
Online liveness is not checked because the iLink long-poll endpoint is
|
||||
not designed for cheap probes.
|
||||
"""
|
||||
if not account_id:
|
||||
return False, (
|
||||
"account_id is required. Run "
|
||||
"`python -m EvoScientist.channels.wechat.serve --qr-login` first."
|
||||
)
|
||||
|
||||
if token:
|
||||
return True, f"Personal WeChat account {account_id[:8]}… token provided"
|
||||
|
||||
from .personal import load_account
|
||||
|
||||
persisted = load_account(account_id)
|
||||
if not persisted or not persisted.get("token"):
|
||||
return False, (
|
||||
f"No saved credentials for account_id={account_id[:8]}…. "
|
||||
"Run `python -m EvoScientist.channels.wechat.serve --qr-login`."
|
||||
)
|
||||
return True, f"Personal WeChat account {account_id[:8]}… loaded from disk"
|
||||
|
||||
@@ -20,6 +20,13 @@ Usage:
|
||||
--token TOKEN \\
|
||||
--aes-key AES_KEY
|
||||
|
||||
# Personal WeChat (个人微信 via iLink Bot)
|
||||
# First, log in via QR scan to obtain credentials:
|
||||
python -m EvoScientist.channels.wechat.serve --qr-login
|
||||
# Then run with the saved account_id:
|
||||
python -m EvoScientist.channels.wechat.serve \\
|
||||
--backend personal --account-id <id>
|
||||
|
||||
Options:
|
||||
--port PORT Webhook listen port (default: 9001)
|
||||
--allow USER_ID Allowed sender (repeatable)
|
||||
@@ -28,13 +35,24 @@ Options:
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from ...logging_config import configure_logging_from_settings
|
||||
from ..bus import MessageBus
|
||||
from ..standalone import run_standalone
|
||||
from .channel import WeChatChannel, WeChatMPConfig, WeComConfig
|
||||
from .personal import (
|
||||
WeixinPersonalChannel,
|
||||
WeixinPersonalConfig,
|
||||
load_account,
|
||||
qr_login,
|
||||
)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -46,10 +64,15 @@ def parse_args():
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
choices=["wecom", "wechatmp"],
|
||||
choices=["wecom", "wechatmp", "personal"],
|
||||
default="wecom",
|
||||
help="WeChat backend type (default: wecom)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--qr-login",
|
||||
action="store_true",
|
||||
help="Run interactive QR-code login for personal WeChat and exit",
|
||||
)
|
||||
parser.add_argument("--port", type=int, default=9001, help="Webhook port")
|
||||
parser.add_argument(
|
||||
"--allow",
|
||||
@@ -85,6 +108,31 @@ def parse_args():
|
||||
mp.add_argument("--app-id", default="", help="MP App ID")
|
||||
mp.add_argument("--app-secret", default="", help="MP App Secret")
|
||||
|
||||
# Personal-WeChat settings
|
||||
personal = parser.add_argument_group("Personal WeChat (iLink Bot)")
|
||||
personal.add_argument(
|
||||
"--account-id",
|
||||
default="",
|
||||
help="iLink account_id (obtained via --qr-login)",
|
||||
)
|
||||
personal.add_argument(
|
||||
"--bot-token",
|
||||
default="",
|
||||
help="iLink bearer token; if omitted, loaded from disk via account-id",
|
||||
)
|
||||
personal.add_argument(
|
||||
"--dm-policy",
|
||||
choices=["open", "allowlist"],
|
||||
default="open",
|
||||
help="Direct-message policy (default: open)",
|
||||
)
|
||||
personal.add_argument(
|
||||
"--group-policy",
|
||||
choices=["open", "allowlist", "disabled"],
|
||||
default="disabled",
|
||||
help="Group-message policy (default: disabled — iLink rarely delivers)",
|
||||
)
|
||||
|
||||
# Shared settings
|
||||
parser.add_argument("--token", default="", help="Callback verification token")
|
||||
parser.add_argument("--aes-key", default="", help="EncodingAESKey")
|
||||
@@ -95,8 +143,14 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
"""Entry point."""
|
||||
configure_logging_from_settings(default_level=logging.INFO)
|
||||
args = parse_args()
|
||||
|
||||
if args.qr_login:
|
||||
result = asyncio.run(qr_login())
|
||||
if not result:
|
||||
raise SystemExit(1)
|
||||
return
|
||||
|
||||
allowed = set(args.allowed_senders) if args.allowed_senders else None
|
||||
allowed_channels = set(args.allowed_channels) if args.allowed_channels else None
|
||||
proxy = args.proxy or None
|
||||
@@ -113,7 +167,8 @@ def main():
|
||||
allowed_channels=allowed_channels,
|
||||
proxy=proxy,
|
||||
)
|
||||
else:
|
||||
channel = WeChatChannel(config, backend=args.backend)
|
||||
elif args.backend == "wechatmp":
|
||||
config = WeChatMPConfig(
|
||||
app_id=args.app_id,
|
||||
app_secret=args.app_secret,
|
||||
@@ -124,11 +179,31 @@ def main():
|
||||
allowed_channels=allowed_channels,
|
||||
proxy=proxy,
|
||||
)
|
||||
channel = WeChatChannel(config, backend=args.backend)
|
||||
else: # personal
|
||||
token = args.bot_token
|
||||
if not token and args.account_id:
|
||||
persisted = load_account(args.account_id)
|
||||
if persisted:
|
||||
token = persisted.get("token", "")
|
||||
if not args.account_id or not token:
|
||||
raise SystemExit(
|
||||
"Personal WeChat requires --account-id (and a saved token, "
|
||||
"obtained via --qr-login)."
|
||||
)
|
||||
personal_config = WeixinPersonalConfig(
|
||||
account_id=args.account_id,
|
||||
token=token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=allowed_channels,
|
||||
dm_policy=args.dm_policy,
|
||||
group_policy=args.group_policy,
|
||||
proxy=proxy,
|
||||
)
|
||||
channel = WeixinPersonalChannel(personal_config)
|
||||
|
||||
send_thinking = args.thinking and args.agent
|
||||
bus = MessageBus()
|
||||
channel = WeChatChannel(config, backend=args.backend)
|
||||
|
||||
run_standalone(channel, bus, use_agent=args.agent, send_thinking=send_thinking)
|
||||
|
||||
|
||||
|
||||
@@ -1,27 +1,67 @@
|
||||
"""EvoScientist CLI package."""
|
||||
"""EvoScientist CLI package.
|
||||
|
||||
# Backward-compat re-exports (tests import these from EvoScientist.cli)
|
||||
from ..stream.state import ( # noqa: F401
|
||||
StreamState,
|
||||
SubAgentState,
|
||||
_build_todo_stats,
|
||||
_parse_todo_items,
|
||||
)
|
||||
Most re-exports are served lazily through ``__getattr__`` so that a bare
|
||||
``import EvoScientist.cli`` only costs what ``main()`` actually needs. That
|
||||
keeps ``evosci --help`` fast — the heavy chat-model/TUI/langgraph imports
|
||||
only pay their cost when someone actually touches those names.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from .. import deploy as _deploy_pkg # noqa: F401 — registers `deploy` @app.command
|
||||
from . import commands # noqa: F401 — registers @app.command decorators
|
||||
from ._app import app
|
||||
from ._constants import WELCOME_SLOGANS # noqa: F401
|
||||
from .agent import _deduplicate_run_name # noqa: F401
|
||||
from .channel import _channels_is_running, _channels_stop # noqa: F401
|
||||
|
||||
# UI runtime re-exports (merged from former tui/ package)
|
||||
from .tui_runtime import ( # noqa: F401
|
||||
DEFAULT_UI_BACKEND,
|
||||
SUPPORTED_UI_BACKENDS,
|
||||
get_backend,
|
||||
normalize_ui_backend,
|
||||
resolve_ui_backend,
|
||||
run_streaming,
|
||||
)
|
||||
__all__ = [
|
||||
"DEFAULT_UI_BACKEND",
|
||||
"SUPPORTED_UI_BACKENDS",
|
||||
"WELCOME_SLOGANS",
|
||||
"StreamState",
|
||||
"SubAgentState",
|
||||
"_build_todo_stats",
|
||||
"_channels_is_running",
|
||||
"_channels_stop",
|
||||
"_deduplicate_run_name",
|
||||
"_parse_todo_items",
|
||||
"app",
|
||||
"get_backend",
|
||||
"main",
|
||||
"normalize_ui_backend",
|
||||
"resolve_ui_backend",
|
||||
"run_streaming",
|
||||
]
|
||||
|
||||
# Map attribute name -> (relative-module, attribute-in-module).
|
||||
# Paths starting with ".." reach out of this package.
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"StreamState": ("..stream.state", "StreamState"),
|
||||
"SubAgentState": ("..stream.state", "SubAgentState"),
|
||||
"_build_todo_stats": ("..stream.state", "_build_todo_stats"),
|
||||
"_parse_todo_items": ("..stream.state", "_parse_todo_items"),
|
||||
"WELCOME_SLOGANS": ("._constants", "WELCOME_SLOGANS"),
|
||||
"_deduplicate_run_name": (".agent", "_deduplicate_run_name"),
|
||||
"_channels_is_running": (".channel", "_channels_is_running"),
|
||||
"_channels_stop": (".channel", "_channels_stop"),
|
||||
"DEFAULT_UI_BACKEND": (".tui_runtime", "DEFAULT_UI_BACKEND"),
|
||||
"SUPPORTED_UI_BACKENDS": (".tui_runtime", "SUPPORTED_UI_BACKENDS"),
|
||||
"get_backend": (".tui_runtime", "get_backend"),
|
||||
"normalize_ui_backend": (".tui_runtime", "normalize_ui_backend"),
|
||||
"resolve_ui_backend": (".tui_runtime", "resolve_ui_backend"),
|
||||
"run_streaming": (".tui_runtime", "run_streaming"),
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
target = _LAZY_EXPORTS.get(name)
|
||||
if target is None:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
from importlib import import_module
|
||||
|
||||
module_path, attr = target
|
||||
module = import_module(module_path, package=__name__)
|
||||
value = getattr(module, attr)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
|
||||
def main():
|
||||
@@ -29,6 +69,12 @@ def main():
|
||||
import os
|
||||
import warnings
|
||||
|
||||
# Keep MCP stdio subprocess spawning async on Windows (see #283). Must run
|
||||
# before any event loop is created, hence at the very top of the entrypoint.
|
||||
from .._winloop import ensure_proactor_event_loop_policy
|
||||
|
||||
ensure_proactor_event_loop_policy()
|
||||
|
||||
warnings.filterwarnings("ignore", message=".*not known to support tools.*")
|
||||
warnings.filterwarnings(
|
||||
"ignore", message=".*type is unknown and inference may fail.*"
|
||||
@@ -41,9 +87,3 @@ def main():
|
||||
_log_level = os.environ.get("EVOSCIENTIST_LOG_LEVEL", "") or config.log_level
|
||||
_configure_logging()
|
||||
app()
|
||||
|
||||
|
||||
def admin_main():
|
||||
"""evo-admin CLI entry point — admin management commands."""
|
||||
from ._app import admin_app
|
||||
admin_app()
|
||||
|
||||
@@ -0,0 +1,199 @@
|
||||
"""Background MCP/agent load lifecycle shared by CLI and TUI surfaces.
|
||||
|
||||
Holds no references to Rich, prompt_toolkit, or Textual — UI-specific
|
||||
rendering and thread-hopping plug in via callbacks.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
ProgressEvent = str # "start" | "success" | "error"
|
||||
ProgressState = str # "pending" | "ok" | "error"
|
||||
|
||||
AgentT = TypeVar("AgentT")
|
||||
ProgressCallback = Callable[[ProgressEvent, str, str], None]
|
||||
SuccessCallback = Callable[[AgentT], None]
|
||||
FailureCallback = Callable[[BaseException], None]
|
||||
|
||||
|
||||
class MCPProgressTracker:
|
||||
"""Per-server MCP load progress state.
|
||||
|
||||
Reads and writes are GIL-atomic but iteration must go through
|
||||
:meth:`snapshot` — events can fire from a worker thread while the
|
||||
main thread renders.
|
||||
"""
|
||||
|
||||
__slots__ = ("progress",)
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.progress: dict[str, tuple[ProgressState, str]] = {}
|
||||
|
||||
def prime(self) -> None:
|
||||
"""Seed a ``pending`` entry for every configured server.
|
||||
|
||||
Keeps the UI's "N / M" denominator stable from the first render.
|
||||
"""
|
||||
try:
|
||||
from ..mcp import load_mcp_config
|
||||
|
||||
cfg = load_mcp_config() or {}
|
||||
self.progress = dict.fromkeys(cfg, ("pending", ""))
|
||||
except Exception:
|
||||
self.progress = {}
|
||||
|
||||
def record(
|
||||
self, event: ProgressEvent, server: str, detail: str
|
||||
) -> ProgressState | None:
|
||||
"""Apply an event and return the new state, or ``None`` if unknown."""
|
||||
if event == "start":
|
||||
self.progress.setdefault(server, ("pending", ""))
|
||||
return "pending"
|
||||
if event == "success":
|
||||
self.progress[server] = ("ok", detail)
|
||||
return "ok"
|
||||
if event == "error":
|
||||
self.progress[server] = ("error", detail)
|
||||
return "error"
|
||||
return None
|
||||
|
||||
def snapshot(self) -> list[tuple[ProgressState, str]]:
|
||||
return list(self.progress.values())
|
||||
|
||||
def totals(self) -> tuple[int, int]:
|
||||
"""``(done, total)`` — done excludes ``pending``."""
|
||||
snap = self.snapshot()
|
||||
total = len(snap)
|
||||
done = sum(1 for state, _ in snap if state != "pending")
|
||||
return done, total
|
||||
|
||||
|
||||
class BackgroundAgentLoader(Generic[AgentT]):
|
||||
"""Owns the background ``_load_agent`` task and its generation token.
|
||||
|
||||
Each :meth:`start` bumps an internal id; callbacks from a superseded
|
||||
load (the old worker thread keeps running after cancel, since
|
||||
``asyncio.to_thread`` can't preempt arbitrary Python code) compare
|
||||
against it and drop silently.
|
||||
|
||||
``on_progress`` fires on the **worker thread**; UI callers hop
|
||||
threads inside it if needed. ``on_success`` / ``on_failure`` fire
|
||||
on the event loop when the task completes.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
loader_fn: Callable[..., AgentT],
|
||||
*,
|
||||
on_progress: ProgressCallback | None = None,
|
||||
on_success: SuccessCallback | None = None,
|
||||
on_failure: FailureCallback | None = None,
|
||||
) -> None:
|
||||
self._loader_fn = loader_fn
|
||||
self._on_progress = on_progress
|
||||
self._on_success = on_success
|
||||
self._on_failure = on_failure
|
||||
self.agent: AgentT | None = None
|
||||
self._task: asyncio.Task[AgentT] | None = None
|
||||
self._load_id: int = 0
|
||||
|
||||
@property
|
||||
def task(self) -> asyncio.Task[AgentT] | None:
|
||||
return self._task
|
||||
|
||||
@property
|
||||
def is_pending(self) -> bool:
|
||||
return self.agent is None and self._task is not None and not self._task.done()
|
||||
|
||||
@property
|
||||
def needs_restart(self) -> bool:
|
||||
"""True when no load is in flight and no agent is ready.
|
||||
|
||||
Callers that want auto-retry behavior (e.g. TUI on the next
|
||||
user send after a failure) check this before :meth:`start`.
|
||||
"""
|
||||
return self.agent is None and (self._task is None or self._task.done())
|
||||
|
||||
def start(self, **loader_kwargs: Any) -> None:
|
||||
prev = self._task
|
||||
if prev is not None and not prev.done():
|
||||
prev.cancel()
|
||||
self._load_id += 1
|
||||
load_id = self._load_id
|
||||
self.agent = None
|
||||
|
||||
def _gated_progress(event: str, server: str, detail: str) -> None:
|
||||
if load_id != self._load_id:
|
||||
return
|
||||
if self._on_progress is None:
|
||||
return
|
||||
try:
|
||||
self._on_progress(event, server, detail)
|
||||
except Exception:
|
||||
_logger.debug("MCP progress callback raised", exc_info=True)
|
||||
|
||||
self._task = asyncio.create_task(
|
||||
asyncio.to_thread(
|
||||
self._loader_fn,
|
||||
on_mcp_progress=_gated_progress,
|
||||
**loader_kwargs,
|
||||
)
|
||||
)
|
||||
self._task.add_done_callback(lambda task, lid=load_id: self._on_done(task, lid))
|
||||
|
||||
def adopt(self, agent: AgentT) -> None:
|
||||
"""Install an externally-built agent and supersede any in-flight load.
|
||||
|
||||
Used by any caller that constructs a replacement agent directly:
|
||||
bumps the generation token so a late-arriving background load can't
|
||||
clobber ``self.agent`` via the done-callback, cancels the in-flight
|
||||
wrapper, and seats the new agent immediately.
|
||||
"""
|
||||
prev = self._task
|
||||
if prev is not None and not prev.done():
|
||||
prev.cancel()
|
||||
self._load_id += 1
|
||||
self._task = None
|
||||
self.agent = agent
|
||||
|
||||
async def await_ready(self) -> AgentT:
|
||||
"""Return the loaded agent; re-raises on load failure.
|
||||
|
||||
Idempotent. State transitions (setting ``self.agent``, calling
|
||||
``on_success`` / ``on_failure``) are handled exclusively by
|
||||
:meth:`_on_done`, which fires before this ``await`` resumes
|
||||
(asyncio guarantees done-callbacks run in registration order).
|
||||
"""
|
||||
if self.agent is not None:
|
||||
return self.agent
|
||||
if self._task is None:
|
||||
raise RuntimeError(
|
||||
"BackgroundAgentLoader.await_ready called before start()"
|
||||
)
|
||||
await self._task
|
||||
if self.agent is None:
|
||||
raise RuntimeError("BackgroundAgentLoader completed without an agent")
|
||||
return self.agent
|
||||
|
||||
def _on_done(self, task: asyncio.Task[AgentT], load_id: int) -> None:
|
||||
if load_id != self._load_id:
|
||||
return
|
||||
if task.cancelled():
|
||||
return
|
||||
try:
|
||||
self.agent = task.result()
|
||||
except Exception as exc:
|
||||
# Keep ``_task`` set so a later ``await_ready`` re-raises the
|
||||
# real exception instead of the "before start()" sentinel.
|
||||
self.agent = None
|
||||
if self._on_failure is not None:
|
||||
self._on_failure(exc)
|
||||
return
|
||||
if self._on_success is not None:
|
||||
self._on_success(self.agent)
|
||||
@@ -52,6 +52,18 @@ app.add_typer(mcp_app, name="mcp")
|
||||
channel_app = typer.Typer(help="Channel management commands")
|
||||
app.add_typer(channel_app, name="channel")
|
||||
|
||||
# Admin subcommand group
|
||||
admin_app = typer.Typer(help="Admin management commands")
|
||||
app.add_typer(admin_app, name="admin")
|
||||
# Sessions subcommand group — diagnostic tools for the LangGraph checkpoint DB
|
||||
sessions_app = typer.Typer(
|
||||
help="Inspect and manage the sessions DB (~/.evoscientist/sessions.db)",
|
||||
invoke_without_command=True,
|
||||
)
|
||||
app.add_typer(sessions_app, name="sessions")
|
||||
|
||||
# Configure subcommand group — re-run a single onboarding section.
|
||||
configure_app = typer.Typer(
|
||||
help=(
|
||||
"Re-run one onboarding section without going through the full wizard.\n"
|
||||
"Example: EvoSci configure provider"
|
||||
),
|
||||
)
|
||||
app.add_typer(configure_app, name="configure")
|
||||
|
||||
@@ -2,8 +2,22 @@
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
|
||||
def _agent_name() -> str:
|
||||
# Deferred import: ``sessions`` pulls in langgraph/aiosqlite (~300 ms)
|
||||
# and is only needed when ``build_metadata`` is actually called.
|
||||
from ..sessions import AGENT_NAME
|
||||
|
||||
return AGENT_NAME
|
||||
|
||||
|
||||
# Dangerous-mode warning banner — shared by Rich CLI, Textual TUI, and serve so
|
||||
# the wording never drifts. Label is rendered white-on-red, message in red.
|
||||
DANGEROUS_BANNER_LABEL = "DANGEROUS MODE"
|
||||
DANGEROUS_BANNER_MESSAGE = (
|
||||
"Real-filesystem access • the agent can read/write/delete anywhere."
|
||||
)
|
||||
|
||||
WELCOME_SLOGANS = [
|
||||
"Ready for vibe research? What do you want cooking?",
|
||||
"Science doesn't sleep. Neither do your sub-agents.",
|
||||
@@ -34,7 +48,7 @@ LOGO_GRADIENT = ["#1a237e", "#1565c0", "#1e88e5", "#42a5f5", "#64b5f6", "#90caf9
|
||||
def build_metadata(workspace_dir: str | None, model: str | None) -> dict:
|
||||
"""Build metadata dict for LangGraph checkpoint persistence."""
|
||||
return {
|
||||
"agent_name": AGENT_NAME,
|
||||
"agent_name": _agent_name(),
|
||||
"updated_at": datetime.now(UTC).isoformat(),
|
||||
"workspace_dir": workspace_dir or "",
|
||||
"model": model or "",
|
||||
|
||||
@@ -3,9 +3,13 @@
|
||||
import os
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..paths import new_run_dir
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
|
||||
def _shorten_path(path: str) -> str:
|
||||
"""Shorten absolute path to relative path from current directory."""
|
||||
@@ -58,18 +62,51 @@ def _create_session_workspace(name: str | None = None) -> str:
|
||||
return workspace_dir
|
||||
|
||||
|
||||
def _load_agent(workspace_dir: str | None = None, checkpointer=None, config=None):
|
||||
def current_model_label() -> str:
|
||||
"""Return the display label for the registry's primary default model.
|
||||
|
||||
The CLI no longer carries ``config.model``/``config.provider`` free
|
||||
strings; the status bar and run metadata label the active model as
|
||||
``provider_id/model_key`` from the Model Registry defaults. Returns
|
||||
``"unconfigured"`` when the registry is still in bootstrap.
|
||||
"""
|
||||
try:
|
||||
from ..model_registry.runtime import get_snapshot_runtime
|
||||
|
||||
primary, _ = get_snapshot_runtime().registry_default()
|
||||
except Exception:
|
||||
return "unconfigured"
|
||||
return f"{primary.provider_id}/{primary.model_key}"
|
||||
|
||||
|
||||
def _load_agent(
|
||||
workspace_dir: str | None = None,
|
||||
checkpointer=None,
|
||||
config=None,
|
||||
chat_model=None,
|
||||
*,
|
||||
on_mcp_progress=None,
|
||||
) -> "CompiledStateGraph":
|
||||
"""Load the CLI agent with optional persistent checkpointer.
|
||||
|
||||
Args:
|
||||
workspace_dir: Optional per-session workspace directory.
|
||||
checkpointer: Optional LangGraph checkpointer.
|
||||
checkpointer: Optional LangGraph checkpointer (e.g. ``AsyncSqliteSaver``).
|
||||
Falls back to ``InMemorySaver`` when ``None``.
|
||||
config: Optional pre-loaded ``EvoScientistConfig``. Forwarded to
|
||||
``create_cli_agent`` to avoid double config loading.
|
||||
chat_model: Optional pre-built chat model. Forwarded to
|
||||
``create_cli_agent``; combined with an explicit ``config`` it
|
||||
selects the pure (no module-global write) build path.
|
||||
on_mcp_progress: Optional per-server MCP progress callback.
|
||||
Signature ``(event, server_name, detail) -> None``.
|
||||
"""
|
||||
from ..EvoScientist import create_cli_agent
|
||||
|
||||
return create_cli_agent(
|
||||
workspace_dir=workspace_dir, checkpointer=checkpointer, config=config
|
||||
workspace_dir=workspace_dir,
|
||||
checkpointer=checkpointer,
|
||||
config=config,
|
||||
chat_model=chat_model,
|
||||
on_mcp_progress=on_mcp_progress,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,595 @@
|
||||
"""Async sub-agent auto-notification.
|
||||
|
||||
When a sub-agent on langgraph dev reaches a terminal state, a watcher coroutine
|
||||
pushes a lightweight notification onto a thread-safe queue. The CLI loop drains
|
||||
the queue, dedups against deepagents' async_tasks state, batches survivors,
|
||||
and injects a synthetic user message that triggers one LLM turn.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias, TypedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway, GraphTarget
|
||||
|
||||
TERMINAL_STATUSES: Final = frozenset({"success", "error", "timeout", "interrupted"})
|
||||
"""Aligned with langgraph_sdk.schema.RunStatus terminal values.
|
||||
|
||||
Cancel operations transition runs into ``interrupted`` (not ``cancelled``).
|
||||
"""
|
||||
|
||||
# How many times the watcher will re-join the SSE stream when it closes
|
||||
# cleanly but ``runs.get`` reports the run is still alive (typical cause:
|
||||
# HTTP keep-alive timeout on long static periods). Bounded to prevent an
|
||||
# unbounded loop if the server permanently misreports status.
|
||||
_MAX_RECONNECT_ATTEMPTS: Final = 10
|
||||
|
||||
|
||||
class AsyncTaskState(TypedDict, total=False):
|
||||
status: str
|
||||
last_checked_at: str
|
||||
last_updated_at: str
|
||||
|
||||
|
||||
AsyncTasksState: TypeAlias = dict[str, AsyncTaskState]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AsyncTaskNotification:
|
||||
"""A completed-async-task signal pushed by a watcher."""
|
||||
|
||||
task_id: str
|
||||
agent_name: str
|
||||
status: str # one of TERMINAL_STATUSES
|
||||
received_at: str # ISO-8601 UTC timestamp
|
||||
prompt: str = "" # original task description sent to the sub-agent
|
||||
kind: str = "agent" # "agent" (sub-agent) | "bg-process" (background shell)
|
||||
# The CLI/main-agent thread_id under which the watcher was spawned. Used
|
||||
# to route the notification back to the originating CLI session so a
|
||||
# /new between launch and completion does not inject the synthetic
|
||||
# message into an unrelated thread (where ``check_async_task`` cannot
|
||||
# find the task_id). ``None`` means "unrouted" — the notification
|
||||
# drains for any current_thread_id (back-compat for direct callers).
|
||||
origin_cli_thread_id: str | None = None
|
||||
|
||||
|
||||
# Per-thread routing: notifications with ``origin_cli_thread_id`` land in
|
||||
# the matching sub-queue. Notifications without one go to ``_unrouted_queue``
|
||||
# and drain regardless of current thread (back-compat for legacy callers
|
||||
# and direct-put test paths).
|
||||
_notifications_by_thread: dict[str, queue.Queue[AsyncTaskNotification]] = {}
|
||||
_notifications_lock = threading.Lock()
|
||||
_unrouted_queue: queue.Queue[AsyncTaskNotification] = queue.Queue()
|
||||
# Public alias for the unrouted bucket — preserved so legacy tests and any
|
||||
# external direct callers that did ``_notification_queue.put(...)`` keep
|
||||
# working unchanged. New code should call ``_enqueue`` instead.
|
||||
_notification_queue = _unrouted_queue
|
||||
|
||||
# Track active watcher tasks/futures for clean shutdown.
|
||||
# dict[handle, origin_cli_thread_id] so the consumer's batching grace loop
|
||||
# can filter for watchers tied to the current CLI thread (or unrouted)
|
||||
# without being delayed by sibling-thread watchers.
|
||||
_active_watchers: dict[object, str | None] = {}
|
||||
# Map thread_id (sub-agent thread) → current watcher handle (supports
|
||||
# replacement on update_async_task).
|
||||
_watcher_by_thread: dict[str, asyncio.Task[None]] = {}
|
||||
|
||||
|
||||
def _has_relevant_active_watchers(current_thread_id: str | None) -> bool:
|
||||
"""Are there any in-flight watchers whose notifications would drain on
|
||||
a ``consume_notifications`` call for ``current_thread_id``?
|
||||
|
||||
A watcher is relevant if its ``origin_cli_thread_id`` matches the
|
||||
current CLI thread or is ``None`` (unrouted bucket drains for any
|
||||
consumer). Sibling-thread watchers are ignored.
|
||||
"""
|
||||
if current_thread_id is None:
|
||||
return bool(_active_watchers)
|
||||
return any(
|
||||
origin == current_thread_id or origin is None
|
||||
for origin in _active_watchers.values()
|
||||
)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _enqueue(notification: AsyncTaskNotification) -> None:
|
||||
"""Route a notification to its origin-thread queue or the unrouted bucket."""
|
||||
tid = notification.origin_cli_thread_id
|
||||
if not tid:
|
||||
_unrouted_queue.put(notification)
|
||||
return
|
||||
with _notifications_lock:
|
||||
q = _notifications_by_thread.get(tid)
|
||||
if q is None:
|
||||
q = queue.Queue()
|
||||
_notifications_by_thread[tid] = q
|
||||
q.put(notification)
|
||||
|
||||
|
||||
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
|
||||
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
|
||||
|
||||
If ``current_thread_id`` is given, only the matching thread queue and
|
||||
the unrouted bucket count. With no argument, only the unrouted bucket
|
||||
counts (legacy behavior).
|
||||
"""
|
||||
if not _unrouted_queue.empty():
|
||||
return True
|
||||
if current_thread_id is None:
|
||||
return False
|
||||
with _notifications_lock:
|
||||
q = _notifications_by_thread.get(current_thread_id)
|
||||
return q is not None and not q.empty()
|
||||
|
||||
|
||||
def pending_thread_ids() -> set[str]:
|
||||
"""Return the set of thread_ids with pending routed notifications."""
|
||||
with _notifications_lock:
|
||||
return {tid for tid, q in _notifications_by_thread.items() if not q.empty()}
|
||||
|
||||
|
||||
async def read_async_tasks_from_gateway(
|
||||
gateway: GraphGateway,
|
||||
target: GraphTarget,
|
||||
thread_id: str,
|
||||
) -> AsyncTasksState:
|
||||
"""Read async_tasks state through the active graph gateway."""
|
||||
try:
|
||||
values = await gateway.get_state_values(target, thread_id)
|
||||
except Exception:
|
||||
return {}
|
||||
return values.get("async_tasks", {})
|
||||
|
||||
|
||||
async def watch_run_and_notify(
|
||||
client,
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
agent_name: str,
|
||||
prompt: str = "",
|
||||
origin_cli_thread_id: str | None = None,
|
||||
) -> None:
|
||||
"""Subscribe to a run's event stream; enqueue notification when it terminates.
|
||||
|
||||
Status detection strategy (priority order):
|
||||
|
||||
1. **In-band ``event="error"`` SSE part** — authoritative error signal
|
||||
from langgraph dev, no race against server-side state writeback.
|
||||
2. **Server-side state via ``runs.get``** — invoked after the stream
|
||||
closes (cleanly or with exception) to verify the run is actually
|
||||
done. Required because SSE long-poll can close on HTTP keep-alive
|
||||
timeout while the run is still running, which would otherwise be
|
||||
misread as ``"success"`` (observed in production with long-running
|
||||
literature search tasks under concurrency).
|
||||
3. **Re-join loop** — if ``runs.get`` reports ``pending`` / ``running``,
|
||||
the run is alive but we lost the stream; re-join up to
|
||||
``_MAX_RECONNECT_ATTEMPTS`` times before giving up.
|
||||
|
||||
The previous implementation trusted clean stream exits as success
|
||||
without any verification, which produced false-positive notifications
|
||||
when SSE keep-alive timeouts closed the stream early.
|
||||
|
||||
Race-safety note: ``runs.get`` returning ``"error"`` immediately after
|
||||
a clean stream close can be a transient state for an actually-successful
|
||||
run (server hasn't finalized the writeback). We trust the absence of
|
||||
in-band error event over a stale ``runs.get="error"`` — see the
|
||||
``status == "error" and not saw_error_event`` branch below.
|
||||
"""
|
||||
for attempt in range(_MAX_RECONNECT_ATTEMPTS + 1):
|
||||
stream_failed = False
|
||||
saw_error_event = False
|
||||
try:
|
||||
async for chunk in client.runs.join_stream(
|
||||
thread_id=thread_id, run_id=run_id, stream_mode="values"
|
||||
):
|
||||
ev = getattr(chunk, "event", None)
|
||||
data = getattr(chunk, "data", None)
|
||||
if ev == "error":
|
||||
saw_error_event = True
|
||||
logger.info(
|
||||
"Watcher saw error event for task %s: %r", thread_id, data
|
||||
)
|
||||
except Exception:
|
||||
stream_failed = True
|
||||
logger.warning(
|
||||
"Watcher stream failed for task %s", thread_id, exc_info=True
|
||||
)
|
||||
|
||||
if saw_error_event:
|
||||
status = "error"
|
||||
break
|
||||
|
||||
# Verify with server before deciding the run is done — clean stream
|
||||
# close does NOT guarantee terminal state.
|
||||
try:
|
||||
run = await client.runs.get(thread_id=thread_id, run_id=run_id)
|
||||
raw = run.get("status", "")
|
||||
except Exception:
|
||||
# Cannot verify terminal state. Defaulting to "success" here would
|
||||
# reintroduce the false-positive class this watcher exists to
|
||||
# prevent (clean stream + transient runs.get failure → unverified
|
||||
# success). Retry within the reconnect budget; on exhaustion drop
|
||||
# the notification rather than guess.
|
||||
if attempt >= _MAX_RECONNECT_ATTEMPTS:
|
||||
logger.warning(
|
||||
"Watcher runs.get failed for task %s after %d reconnects; "
|
||||
"unable to verify terminal state, skipping notification",
|
||||
thread_id,
|
||||
_MAX_RECONNECT_ATTEMPTS,
|
||||
exc_info=True,
|
||||
)
|
||||
return
|
||||
logger.warning(
|
||||
"Watcher runs.get failed for task %s; retrying after backoff "
|
||||
"(attempt %d)",
|
||||
thread_id,
|
||||
attempt + 1,
|
||||
exc_info=True,
|
||||
)
|
||||
await asyncio.sleep(min(0.25 * (attempt + 1), 2.0))
|
||||
continue
|
||||
|
||||
if raw not in TERMINAL_STATUSES:
|
||||
# Non-terminal status — includes the documented ``pending`` /
|
||||
# ``running`` values AND any future / unknown status the SDK may
|
||||
# introduce. Stream closed early but run is not done; re-join
|
||||
# unless we've exhausted attempts. Treating unknown statuses as
|
||||
# non-terminal is the safe default — better to retry once more
|
||||
# than to enqueue a false-positive on an unrecognized state.
|
||||
if attempt >= _MAX_RECONNECT_ATTEMPTS:
|
||||
logger.warning(
|
||||
"Watcher gave up on task %s after %d reconnects "
|
||||
"(server still reports %r); skipping notification",
|
||||
thread_id,
|
||||
_MAX_RECONNECT_ATTEMPTS,
|
||||
raw,
|
||||
)
|
||||
return
|
||||
logger.info(
|
||||
"Watcher SSE closed for task %s but run reports %r; "
|
||||
"re-joining (attempt %d)",
|
||||
thread_id,
|
||||
raw,
|
||||
attempt + 1,
|
||||
)
|
||||
continue
|
||||
|
||||
if raw == "error":
|
||||
# Race-safe interpretation: no in-band error event → trust the
|
||||
# absence over the server-side ``error`` (likely transient
|
||||
# writeback state for a successful run). Stream-failure path
|
||||
# is the one case where we DO trust ``error`` — the stream
|
||||
# blowing up usually means something genuinely went wrong.
|
||||
status = "error" if stream_failed else "success"
|
||||
break
|
||||
|
||||
# success / timeout / interrupted — trust authoritative terminal status.
|
||||
status = raw
|
||||
break
|
||||
else:
|
||||
# Loop exhausted without a break — should be unreachable because the
|
||||
# re-join branch returns explicitly when attempts are exhausted, but
|
||||
# guard against future refactors.
|
||||
return
|
||||
|
||||
notification = AsyncTaskNotification(
|
||||
task_id=thread_id,
|
||||
agent_name=agent_name,
|
||||
status=status,
|
||||
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
|
||||
prompt=prompt,
|
||||
origin_cli_thread_id=origin_cli_thread_id,
|
||||
)
|
||||
_enqueue(notification)
|
||||
logger.info(
|
||||
"Enqueued async notification: task=%s agent=%s status=%s origin_thread=%s",
|
||||
thread_id,
|
||||
agent_name,
|
||||
status,
|
||||
origin_cli_thread_id or "<unrouted>",
|
||||
)
|
||||
|
||||
|
||||
def spawn_watcher(
|
||||
client,
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
agent_name: str,
|
||||
prompt: str = "",
|
||||
origin_cli_thread_id: str | None = None,
|
||||
) -> asyncio.Task[None]:
|
||||
"""Spawn a watcher on the caller's asyncio loop.
|
||||
|
||||
Replacement semantics support ``update_async_task`` which creates a new
|
||||
run_id on the same thread_id — we want the new watcher to take over
|
||||
without the old (now obsolete) watcher firing a stale notification.
|
||||
Cancellation propagates ``CancelledError`` (a BaseException), which the
|
||||
watcher's ``except Exception:`` does NOT catch — so ``_enqueue(...)``
|
||||
never executes for the cancelled watcher (no stale notification).
|
||||
|
||||
``origin_cli_thread_id`` tags the resulting notification so the consumer
|
||||
only injects it back into the originating CLI session.
|
||||
|
||||
Caller must already be in a running asyncio event loop. Serve mode's
|
||||
ephemeral per-turn loop kills watchers spawned during a turn — that
|
||||
limitation is tracked separately.
|
||||
"""
|
||||
old_task = _watcher_by_thread.get(thread_id)
|
||||
if old_task is not None and not old_task.done():
|
||||
old_task.cancel()
|
||||
|
||||
task = asyncio.create_task(
|
||||
watch_run_and_notify(
|
||||
client,
|
||||
thread_id,
|
||||
run_id,
|
||||
agent_name,
|
||||
prompt,
|
||||
origin_cli_thread_id=origin_cli_thread_id,
|
||||
)
|
||||
)
|
||||
_watcher_by_thread[thread_id] = task
|
||||
_active_watchers[task] = origin_cli_thread_id
|
||||
|
||||
def _cleanup(t: asyncio.Task[None]) -> None:
|
||||
_active_watchers.pop(t, None)
|
||||
# Only remove if THIS task is still the registered one — could
|
||||
# have been replaced by a newer spawn_watcher call already.
|
||||
if _watcher_by_thread.get(thread_id) is t:
|
||||
del _watcher_by_thread[thread_id]
|
||||
|
||||
task.add_done_callback(_cleanup)
|
||||
return task
|
||||
|
||||
|
||||
def _drain_one_queue(q: queue.Queue) -> list[AsyncTaskNotification]:
|
||||
items: list[AsyncTaskNotification] = []
|
||||
while True:
|
||||
try:
|
||||
items.append(q.get_nowait())
|
||||
except queue.Empty:
|
||||
return items
|
||||
|
||||
|
||||
def drain_notifications(
|
||||
current_thread_id: str | None = None,
|
||||
) -> list[AsyncTaskNotification]:
|
||||
"""Pull pending notifications off the queue (non-blocking).
|
||||
|
||||
With ``current_thread_id``: drains the matching per-thread queue plus
|
||||
the unrouted bucket. Without it: drains EVERY queue (legacy behavior;
|
||||
used by tests and diagnostics).
|
||||
"""
|
||||
if current_thread_id is None:
|
||||
items: list[AsyncTaskNotification] = _drain_one_queue(_unrouted_queue)
|
||||
with _notifications_lock:
|
||||
queues = list(_notifications_by_thread.values())
|
||||
for q in queues:
|
||||
items.extend(_drain_one_queue(q))
|
||||
return items
|
||||
|
||||
items = _drain_one_queue(_unrouted_queue)
|
||||
with _notifications_lock:
|
||||
q = _notifications_by_thread.get(current_thread_id)
|
||||
if q is not None:
|
||||
items.extend(_drain_one_queue(q))
|
||||
return items
|
||||
|
||||
|
||||
def dedup_notifications(
|
||||
notifs: list[AsyncTaskNotification],
|
||||
async_tasks: AsyncTasksState | None,
|
||||
) -> list[AsyncTaskNotification]:
|
||||
"""Filter notifications the agent has already 'seen' via prior check.
|
||||
|
||||
Logic: skip a notification if `async_tasks[task_id]` exists with a TERMINAL
|
||||
status and `last_checked_at >= last_updated_at` (timestamps are ISO-8601
|
||||
so lexicographic comparison is correct). Also skip if `last_checked_at`
|
||||
is empty (brand-new task where agent hasn't checked yet).
|
||||
"""
|
||||
from .. import background # cli -> core import; lazy to avoid import-order issues
|
||||
|
||||
async_tasks = async_tasks or {}
|
||||
survivors: list[AsyncTaskNotification] = []
|
||||
for n in notifs:
|
||||
if n.kind == "bg-process":
|
||||
# Background process: skip if the launching session already inspected it
|
||||
# after it finished (check_process / list_processes) — mirrors the task
|
||||
# dedup below. Per-thread: another session's check doesn't suppress this.
|
||||
if background.was_observed_done(n.task_id, n.origin_cli_thread_id):
|
||||
logger.debug("Dedup: skipping shell notification for %s", n.task_id)
|
||||
continue
|
||||
survivors.append(n)
|
||||
continue
|
||||
task = async_tasks.get(n.task_id)
|
||||
if (
|
||||
task
|
||||
and task.get("status") in TERMINAL_STATUSES
|
||||
and task.get("last_checked_at", "") >= task.get("last_updated_at", "")
|
||||
and task.get("last_checked_at", "") != ""
|
||||
):
|
||||
logger.debug(
|
||||
"Dedup: skipping notification for already-checked task %s", n.task_id
|
||||
)
|
||||
continue
|
||||
survivors.append(n)
|
||||
return survivors
|
||||
|
||||
|
||||
def _render_notification_group(
|
||||
notifs: list[AsyncTaskNotification], title: str, label: str
|
||||
) -> list[tuple[str, str]]:
|
||||
"""Render one group of notifications inside a titled open-right frame.
|
||||
|
||||
Open-right compact frame; bottom matches the top's width:
|
||||
╭── ✦ Agent Teams ✦ ────
|
||||
✔ writing Task: ... success
|
||||
╰─────────────────────────
|
||||
"""
|
||||
top_divider = "╭──" + title + "────" # 4 dashes on the right (2x of left)
|
||||
bottom_divider = "╰" + "─" * (len(top_divider) - 1)
|
||||
lines: list[tuple[str, str]] = [(top_divider, "dim")]
|
||||
for n in notifs:
|
||||
# `writing-agent` → `writing`.
|
||||
name = n.agent_name.removesuffix("-agent")
|
||||
if n.status == "success":
|
||||
icon, color = "✔", "#e67e22" # carrot orange (CSS hex; Rich+Textual)
|
||||
elif n.status == "error":
|
||||
icon, color = "✗", "red"
|
||||
else: # cancelled, timeout, interrupted
|
||||
icon, color = "⚠", "yellow"
|
||||
# Collapse newlines, truncate prompt/command preview to 60 chars.
|
||||
prompt_preview = (n.prompt or "").replace("\n", " ").strip()
|
||||
if len(prompt_preview) > 60:
|
||||
prompt_preview = prompt_preview[:60] + "…"
|
||||
if prompt_preview:
|
||||
text = f" {icon} {name:18s} {label}: {prompt_preview} {n.status}"
|
||||
else:
|
||||
# Fallback: short task_id when no prompt is available
|
||||
short_tid = (
|
||||
f"{n.task_id[:8]}…{n.task_id[-4:]}"
|
||||
if len(n.task_id) > 12
|
||||
else n.task_id
|
||||
)
|
||||
text = f" {icon} {name:18s} ({short_tid}) {n.status}"
|
||||
lines.append((text, color))
|
||||
lines.append((bottom_divider, "dim"))
|
||||
return lines
|
||||
|
||||
|
||||
def format_notification_lines(
|
||||
notifs: list[AsyncTaskNotification],
|
||||
) -> list[tuple[str, str]]:
|
||||
"""Render notifications as compact tool-result-style lines for screen display.
|
||||
|
||||
Async sub-agents and background processes get SEPARATE titled frames so a shell
|
||||
background process is never mislabeled as an "Agent Team". Returns (text, rich_style)
|
||||
tuples. The LLM still receives the full ``format_batch_message`` text; this is purely
|
||||
the visual representation for the human operator.
|
||||
"""
|
||||
if not notifs:
|
||||
return []
|
||||
tasks = [n for n in notifs if n.kind == "agent"]
|
||||
shell = [n for n in notifs if n.kind == "bg-process"]
|
||||
unknown = [n for n in notifs if n.kind not in {"agent", "bg-process"}]
|
||||
lines: list[tuple[str, str]] = []
|
||||
if tasks:
|
||||
lines += _render_notification_group(tasks, " ✦ Agent Teams ✦ ", "Task")
|
||||
if shell:
|
||||
lines += _render_notification_group(shell, " ✦ Background ✦ ", "Cmd")
|
||||
if unknown:
|
||||
# Fallback so a future kind is never silently dropped from the display.
|
||||
lines += _render_notification_group(unknown, " ✦ Updates ✦ ", "Task")
|
||||
return lines
|
||||
|
||||
|
||||
def format_batch_message(notifs: list[AsyncTaskNotification]) -> str:
|
||||
"""Compose the synthetic user message that wakes the supervisor.
|
||||
|
||||
Each task is rendered as a compact JSON object (one per line) so the LLM
|
||||
can reliably parse agent name, status, and task_id without ambiguity.
|
||||
``ensure_ascii=False`` lets non-ASCII agent names pass through unchanged.
|
||||
Visual decoration lives in ``format_notification_lines``.
|
||||
"""
|
||||
if not notifs:
|
||||
return ""
|
||||
lines = ["[Async tasks update]"]
|
||||
for n in notifs:
|
||||
lines.append(
|
||||
json.dumps(
|
||||
{
|
||||
"agent": n.agent_name,
|
||||
"kind": n.kind,
|
||||
"status": n.status,
|
||||
"task_id": n.task_id,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
)
|
||||
# bg-process is inspected with check_process; sub-agents with check_async_task.
|
||||
hints: list[str] = []
|
||||
if any(n.kind == "agent" for n in notifs):
|
||||
hints.append("check_async_task (sub-agents)")
|
||||
if any(n.kind == "bg-process" for n in notifs):
|
||||
hints.append("check_process (background processes)")
|
||||
# Fallback when a batch has only unrecognized kinds (hints empty).
|
||||
hint_text = " or ".join(hints) if hints else "the appropriate status tool"
|
||||
lines.append(
|
||||
f"(Signal only — fetch full result via {hint_text} if relevant to "
|
||||
"the current step, else acknowledge & continue.)"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
# Brief grace window after the last drain: catch one final burst of arrivals
|
||||
NOTIFICATION_BATCH_GRACE_SECONDS = 0.3
|
||||
# Max time we'll wait for in-flight watchers to settle before triggering the
|
||||
# agent turn — bounds latency for long-running tasks while still batching
|
||||
# co-completing ones.
|
||||
NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS = 3.0
|
||||
|
||||
|
||||
async def consume_notifications(
|
||||
run_message: Callable[[str, list[AsyncTaskNotification]], Awaitable[None]],
|
||||
read_async_tasks_state: Callable[[], Awaitable[AsyncTasksState]],
|
||||
current_thread_id: str | None = None,
|
||||
) -> None:
|
||||
"""Drain queue, dedup, batch, and inject as a synthetic user message.
|
||||
|
||||
Args:
|
||||
run_message: async callable receiving (llm_text, notifs_list).
|
||||
``llm_text`` is the full structured message for the LLM
|
||||
(from ``format_batch_message``). ``notifs_list`` is the
|
||||
survivors list so callers can render per-task visual lines
|
||||
without re-parsing the text.
|
||||
read_async_tasks_state: async callable returning current ``async_tasks``
|
||||
from the agent's state for dedup.
|
||||
current_thread_id: the active CLI thread id. When given, only
|
||||
notifications whose ``origin_cli_thread_id`` matches (or that
|
||||
were enqueued unrouted) are drained — notifications belonging
|
||||
to other threads stay queued and naturally drain on the next
|
||||
poller tick after the user ``/resume``s back into them. When
|
||||
omitted (legacy callers / tests), every queue drains.
|
||||
"""
|
||||
notifs = drain_notifications(current_thread_id)
|
||||
if not notifs:
|
||||
return
|
||||
# Adaptive grace: if other watchers tied to THIS thread (or unrouted) are
|
||||
# still in flight, wait briefly for them to settle so co-completing tasks
|
||||
# batch into a single agent turn. Sibling-thread watchers don't count —
|
||||
# their notifications wouldn't drain on this tick anyway.
|
||||
loop = asyncio.get_running_loop()
|
||||
deadline = loop.time() + NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS
|
||||
while _has_relevant_active_watchers(current_thread_id) and loop.time() < deadline:
|
||||
await asyncio.sleep(0.2)
|
||||
notifs.extend(drain_notifications(current_thread_id))
|
||||
# Final brief grace to catch arrivals enqueued just before this tick
|
||||
await asyncio.sleep(NOTIFICATION_BATCH_GRACE_SECONDS)
|
||||
notifs.extend(drain_notifications(current_thread_id))
|
||||
|
||||
try:
|
||||
async_tasks = await read_async_tasks_state()
|
||||
except Exception:
|
||||
logger.warning("Failed to read async_tasks state for dedup", exc_info=True)
|
||||
async_tasks = {}
|
||||
|
||||
survivors = dedup_notifications(notifs, async_tasks)
|
||||
if not survivors:
|
||||
logger.info(
|
||||
"All %d notifications deduped (already known to agent)", len(notifs)
|
||||
)
|
||||
return
|
||||
|
||||
text = format_batch_message(survivors)
|
||||
await run_message(text, survivors)
|
||||
@@ -9,20 +9,26 @@ enqueues a ``ChannelMessage`` on a thread-safe ``queue.Queue`` and waits
|
||||
for the main thread to set a response via ``_set_channel_response()``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from rich.panel import Panel
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
from ..stream.display import console
|
||||
from ..commands.base import ChannelRuntime
|
||||
from ..stream.console import console
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway
|
||||
|
||||
_channel_logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -59,6 +65,10 @@ _response_lock = threading.Lock()
|
||||
_RESPONSE_TIMEOUT = 600.0
|
||||
_LATE_RESPONSE_TIMEOUT = 86400.0
|
||||
_LATE_RESPONSE_NOTICE = "Still working on it. I'll send the result when it's ready."
|
||||
_channel_request_lock = threading.Lock()
|
||||
_channel_requests: dict[str, dict[str, str]] = {}
|
||||
_session_requests: dict[str, list[str]] = {}
|
||||
_cancelled_channel_messages: set[str] = set()
|
||||
|
||||
|
||||
def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
|
||||
@@ -71,6 +81,7 @@ def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
|
||||
"loop": loop,
|
||||
"response": None,
|
||||
}
|
||||
_register_channel_request(msg)
|
||||
_message_queue.put(msg)
|
||||
return future
|
||||
|
||||
@@ -106,6 +117,336 @@ def _pop_channel_response(msg_id: str, *, cancel_pending: bool = False) -> str |
|
||||
return slot["response"]
|
||||
|
||||
|
||||
def _channel_session_key(channel_type: str, chat_id: str) -> str:
|
||||
return f"{channel_type}:{chat_id}"
|
||||
|
||||
|
||||
def _channel_message_session_key(msg: ChannelMessage) -> str:
|
||||
return _channel_session_key(msg.channel_type, msg.chat_id)
|
||||
|
||||
|
||||
def _channel_message_cancel_scope(msg: ChannelMessage) -> str:
|
||||
return f"channel:{msg.channel_type}:{msg.chat_id}:{msg.msg_id}"
|
||||
|
||||
|
||||
def _register_channel_request(msg: ChannelMessage) -> None:
|
||||
"""Track a queued channel request so `/stop` can find it later."""
|
||||
session_key = _channel_message_session_key(msg)
|
||||
with _channel_request_lock:
|
||||
_channel_requests[msg.msg_id] = {
|
||||
"session_key": session_key,
|
||||
"cancel_scope": _channel_message_cancel_scope(msg),
|
||||
"state": "queued",
|
||||
}
|
||||
_session_requests.setdefault(session_key, []).append(msg.msg_id)
|
||||
|
||||
|
||||
def _claim_channel_request(msg: ChannelMessage) -> bool:
|
||||
"""Mark a queued request active. Returns False if it was cancelled first."""
|
||||
with _channel_request_lock:
|
||||
slot = _channel_requests.get(msg.msg_id)
|
||||
if slot is None or msg.msg_id in _cancelled_channel_messages:
|
||||
return False
|
||||
slot["state"] = "active"
|
||||
return True
|
||||
|
||||
|
||||
def _claim_or_complete_channel_request(msg: ChannelMessage) -> bool:
|
||||
"""Claim a request, or clean it up if `/stop` cancelled it while queued."""
|
||||
if _claim_channel_request(msg):
|
||||
return True
|
||||
_complete_channel_request(msg.msg_id)
|
||||
return False
|
||||
|
||||
|
||||
def _channel_request_state(msg_id: str) -> str | None:
|
||||
with _channel_request_lock:
|
||||
slot = _channel_requests.get(msg_id)
|
||||
return slot.get("state") if slot is not None else None
|
||||
|
||||
|
||||
def _complete_channel_request(
|
||||
msg_id: str,
|
||||
*,
|
||||
discard_cancel_scope: bool = True,
|
||||
) -> None:
|
||||
"""Forget a request once its waiter is resolved or cancelled."""
|
||||
with _channel_request_lock:
|
||||
slot = _channel_requests.pop(msg_id, None)
|
||||
_cancelled_channel_messages.discard(msg_id)
|
||||
if slot is not None:
|
||||
request_ids = _session_requests.get(slot["session_key"])
|
||||
if request_ids:
|
||||
try:
|
||||
request_ids.remove(msg_id)
|
||||
except ValueError:
|
||||
pass
|
||||
if not request_ids:
|
||||
_session_requests.pop(slot["session_key"], None)
|
||||
|
||||
if slot is not None and discard_cancel_scope:
|
||||
from ..stream.display import discard_stream_cancel
|
||||
|
||||
discard_stream_cancel(slot["cancel_scope"])
|
||||
|
||||
|
||||
def _cancel_channel_session(channel_type: str, chat_id: str) -> tuple[int, int]:
|
||||
"""Cancel queued and active work for one channel chat session."""
|
||||
session_key = _channel_session_key(channel_type, chat_id)
|
||||
with _channel_request_lock:
|
||||
request_ids: list[str] = []
|
||||
cancelled_ids: list[str] = []
|
||||
active_scopes: list[str] = []
|
||||
with _response_lock:
|
||||
for msg_id in tuple(_session_requests.get(session_key, ())):
|
||||
request_slot = _channel_requests.get(msg_id)
|
||||
if request_slot is None:
|
||||
continue
|
||||
response_slot = _pending_responses.get(msg_id)
|
||||
response_resolved = False
|
||||
if response_slot is not None:
|
||||
future = response_slot["future"]
|
||||
# Once a response is already resolved, leave the slot alone
|
||||
# so the bus waiter can still publish it instead of falling
|
||||
# back to "No response".
|
||||
response_resolved = (
|
||||
response_slot.get("response") is not None or future.done()
|
||||
)
|
||||
if not response_resolved:
|
||||
request_ids.append(msg_id)
|
||||
|
||||
should_cancel = False
|
||||
if response_slot is None:
|
||||
should_cancel = request_slot.get("state") == "active"
|
||||
else:
|
||||
should_cancel = not response_resolved
|
||||
|
||||
if should_cancel:
|
||||
cancelled_ids.append(msg_id)
|
||||
if request_slot.get("state") == "active" and should_cancel:
|
||||
active_scopes.append(request_slot["cancel_scope"])
|
||||
_cancelled_channel_messages.update(cancelled_ids)
|
||||
|
||||
for msg_id in request_ids:
|
||||
_pop_channel_response(msg_id, cancel_pending=True)
|
||||
|
||||
if active_scopes:
|
||||
from ..stream.display import request_stream_cancel
|
||||
|
||||
for cancel_scope in active_scopes:
|
||||
request_stream_cancel(cancel_scope)
|
||||
|
||||
return len(request_ids), len(active_scopes)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Slash command dispatch for channel messages
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared by all three UI surfaces that accept inbound channel messages:
|
||||
# Rich CLI (``cli/interactive.py::_process_channel_message``), Textual
|
||||
# TUI (``cli/tui_interactive.py``'s channel handler), and headless
|
||||
# serve (``cli/commands.py::_serve_process_message``). They all route
|
||||
# ``/foo`` text through ``cmd_manager`` instead of feeding it to the
|
||||
# LLM as a plain prompt.
|
||||
|
||||
|
||||
async def dispatch_channel_slash_command(
|
||||
msg: ChannelMessage,
|
||||
*,
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
workspace_dir: str | None,
|
||||
checkpointer: Any,
|
||||
append_system: Callable[[str, str], None],
|
||||
graph_gateway: GraphGateway,
|
||||
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
|
||||
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
|
||||
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
|
||||
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
|
||||
channel_runtime: ChannelRuntime | None = None,
|
||||
) -> bool:
|
||||
"""Dispatch a slash command from a channel message.
|
||||
|
||||
Returns True if the helper handled the message (successfully or with
|
||||
an error) — the caller must then return without streaming anything
|
||||
to the agent. Returns False for non-slash content or unresolved
|
||||
slash commands, so the caller can fall through to the agent
|
||||
streaming path (matches TUI behavior).
|
||||
|
||||
Parameters
|
||||
----------
|
||||
msg:
|
||||
The inbound ``ChannelMessage`` to inspect.
|
||||
agent:
|
||||
Default agent handle for the ``CommandContext``. Commands that
|
||||
do not need the agent use this value directly.
|
||||
thread_id, workspace_dir, checkpointer:
|
||||
Populate ``CommandContext``.
|
||||
append_system:
|
||||
``(text, style)`` callback for local CLI/TUI log output. Used
|
||||
by ``ChannelCommandUI`` to surface system breadcrumbs and by
|
||||
this helper to print the "Executed command from ..." line.
|
||||
start_new_session_cb, handle_session_resume_cb:
|
||||
Optional lifecycle callbacks forwarded to ``ChannelCommandUI``.
|
||||
Headless serve passes ``None`` — ``/new`` and ``/resume`` degrade
|
||||
gracefully via the default ``ChannelCommandUI`` messages.
|
||||
graph_gateway:
|
||||
Graph gateway forwarded to slash commands and channel resume-history
|
||||
rendering.
|
||||
await_agent_ready:
|
||||
Optional async resolver that blocks until the background agent
|
||||
load finishes. Called only when ``cmd.needs_agent(args)`` is
|
||||
True. Headless serve passes ``None`` because the agent is
|
||||
loaded up-front before the bus starts.
|
||||
on_cmd_completed:
|
||||
Optional ``async (ctx, original_agent, cmd) -> None`` callback
|
||||
fired only after ``cmd_manager.execute`` returns True. The
|
||||
``original_agent`` argument is the agent handle command execution
|
||||
started against: ``agent_for_ctx`` after any ``await_agent_ready``
|
||||
resolution, or the dispatcher's input agent when no resolver is
|
||||
supplied. Callers can compare ``ctx.agent`` with
|
||||
``original_agent`` to detect command-driven swaps. Used by Rich
|
||||
CLI to (a) adopt an agent swap back into the
|
||||
running session and (b) refresh the status snapshot for
|
||||
commands that mutate session-level state (``/new``,
|
||||
``/compact``) — mirrors the REPL dispatch at
|
||||
``cli/interactive.py:1002-1030``. Headless serve passes
|
||||
``None`` since it cannot hot-swap its polling-loop agent.
|
||||
"""
|
||||
if not msg.content.strip().startswith("/"):
|
||||
return False
|
||||
|
||||
try:
|
||||
return await _dispatch_channel_slash_impl(
|
||||
msg,
|
||||
agent=agent,
|
||||
thread_id=thread_id,
|
||||
workspace_dir=workspace_dir,
|
||||
checkpointer=checkpointer,
|
||||
append_system=append_system,
|
||||
start_new_session_cb=start_new_session_cb,
|
||||
handle_session_resume_cb=handle_session_resume_cb,
|
||||
await_agent_ready=await_agent_ready,
|
||||
on_cmd_completed=on_cmd_completed,
|
||||
channel_runtime=channel_runtime,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Last-ditch safety: any uncaught exception from inside the
|
||||
# dispatch pipeline (lazy import failure, ChannelCommandUI
|
||||
# construction, terminal I/O from ``append_system``, bus
|
||||
# publish races, ...) must not take down the caller's polling
|
||||
# loop — a crashed serve / dead channel queue task is worse
|
||||
# than one failed command.
|
||||
_channel_logger.exception(
|
||||
"Unexpected slash dispatch failure for %s (msg=%s)",
|
||||
msg.channel_type,
|
||||
msg.msg_id,
|
||||
)
|
||||
try:
|
||||
_set_channel_response(msg.msg_id, f"Command error: {exc}")
|
||||
except Exception: # pragma: no cover — defensive
|
||||
pass
|
||||
# Return True so the caller treats the message as handled and
|
||||
# does not fall through to the agent streaming path.
|
||||
return True
|
||||
|
||||
|
||||
async def _dispatch_channel_slash_impl(
|
||||
msg: ChannelMessage,
|
||||
*,
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
workspace_dir: str | None,
|
||||
checkpointer: Any,
|
||||
append_system: Callable[[str, str], None],
|
||||
graph_gateway: GraphGateway,
|
||||
start_new_session_cb: Callable[[], Awaitable[None]] | None,
|
||||
handle_session_resume_cb: Callable[..., Awaitable[None]] | None,
|
||||
await_agent_ready: Callable[[], Awaitable[Any]] | None,
|
||||
on_cmd_completed: Callable[..., Awaitable[None]] | None,
|
||||
channel_runtime: ChannelRuntime | None,
|
||||
) -> bool:
|
||||
"""Inner body of ``dispatch_channel_slash_command``.
|
||||
|
||||
Split from the public wrapper so the wrapper can guard with a
|
||||
top-level try/except without visually obscuring the main flow.
|
||||
"""
|
||||
# Lazy imports: avoid coupling the channel module to ``commands`` at
|
||||
# import time (tui_interactive.py does the same).
|
||||
from ..commands.base import CommandContext
|
||||
from ..commands.channel_ui import ChannelCommandUI
|
||||
from ..commands.manager import manager as cmd_manager
|
||||
|
||||
parsed = cmd_manager.resolve(msg.content)
|
||||
if parsed is None:
|
||||
# Unknown slash command — let the agent handle it (matches TUI).
|
||||
return False
|
||||
cmd, cmd_args = parsed
|
||||
|
||||
agent_for_ctx = agent
|
||||
if cmd.needs_agent(cmd_args) and await_agent_ready is not None:
|
||||
try:
|
||||
agent_for_ctx = await await_agent_ready()
|
||||
except Exception as exc:
|
||||
_set_channel_response(msg.msg_id, f"Command error: {exc}")
|
||||
return True
|
||||
|
||||
ui = ChannelCommandUI(
|
||||
msg,
|
||||
append_system_callback=append_system,
|
||||
start_new_session_callback=start_new_session_cb,
|
||||
handle_session_resume_callback=handle_session_resume_cb,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
ctx = CommandContext(
|
||||
agent=agent_for_ctx,
|
||||
thread_id=thread_id,
|
||||
ui=ui,
|
||||
workspace_dir=workspace_dir,
|
||||
checkpointer=checkpointer,
|
||||
channel_runtime=channel_runtime,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
|
||||
try:
|
||||
cmd_executed = await cmd_manager.execute(msg.content, ctx)
|
||||
except Exception as exc:
|
||||
_channel_logger.debug(f"Channel command error: {exc}", exc_info=True)
|
||||
_set_channel_response(msg.msg_id, f"Command error: {exc}")
|
||||
return True # must return — do NOT fall through to the agent
|
||||
|
||||
if cmd_executed:
|
||||
if ctx.command_error is not None:
|
||||
details = ctx.command_error or "(no details)"
|
||||
_set_channel_response(msg.msg_id, f"Command error: {details}")
|
||||
return True
|
||||
|
||||
if on_cmd_completed is not None:
|
||||
try:
|
||||
# Command output already flushed by ``cmd_manager.execute``
|
||||
# via ``ctx.ui.flush()`` — the hook does internal state
|
||||
# sync (agent adoption, status snapshot refresh) only,
|
||||
# so swallowing its errors keeps the user-visible reply
|
||||
# intact even if the sync path is broken.
|
||||
await on_cmd_completed(ctx, agent_for_ctx, cmd)
|
||||
except Exception as exc:
|
||||
_channel_logger.debug(
|
||||
f"Channel command post-exec callback error: {exc}",
|
||||
exc_info=True,
|
||||
)
|
||||
append_system(
|
||||
f"[{msg.channel_type}: Executed command from {msg.sender}]",
|
||||
"dim",
|
||||
)
|
||||
_set_channel_response(msg.msg_id, f"Command executed: {msg.content}")
|
||||
return True
|
||||
|
||||
# ``cmd_manager.execute`` returned False (empty / unparseable input).
|
||||
# Fall through to the agent streaming path.
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HITL approval intercept: bus thread ⇄ main CLI thread
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -121,6 +462,143 @@ _HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
|
||||
_ASK_USER_TIMEOUT = (
|
||||
300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
||||
)
|
||||
_STOP_COMMANDS = frozenset(("/stop", "/cancel"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-thread channel-origin registry
|
||||
# ---------------------------------------------------------------------------
|
||||
# When a channel-originated message starts an agent turn, the turn's
|
||||
# thread_id is remembered against its (channel_type, chat_id, metadata).
|
||||
# Later, when an async sub-agent notification fires a synthetic agent turn
|
||||
# for that same thread_id, the notifier path pushes the synthesized final
|
||||
# response back to the same chat — otherwise the follow-up would only render
|
||||
# locally and the channel user would never see it. v1 forwards only the
|
||||
# final response (no mid-turn thinking/todo/media).
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ChannelOrigin:
|
||||
"""Channel destination remembered for a thread, for notifier push-back."""
|
||||
|
||||
channel_type: str
|
||||
chat_id: str
|
||||
sender: str
|
||||
metadata: dict | None = None
|
||||
|
||||
|
||||
_thread_channel_origins: dict[str, _ChannelOrigin] = {}
|
||||
_thread_channel_origins_lock = threading.Lock()
|
||||
|
||||
|
||||
def remember_channel_origin(thread_id: str | None, msg: ChannelMessage) -> None:
|
||||
"""Record that ``thread_id`` is currently bound to ``msg``'s channel chat.
|
||||
|
||||
Called on entry to each channel-triggered agent turn (Rich CLI / TUI /
|
||||
serve). The latest channel turn for a given thread wins — re-registering
|
||||
is intentional, since the user can keep talking on the same thread from
|
||||
the same channel and we always want the most recent metadata.
|
||||
"""
|
||||
if not thread_id:
|
||||
return
|
||||
with _thread_channel_origins_lock:
|
||||
_thread_channel_origins[thread_id] = _ChannelOrigin(
|
||||
channel_type=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
sender=msg.sender,
|
||||
metadata=dict(msg.metadata) if msg.metadata else None,
|
||||
)
|
||||
|
||||
|
||||
def get_channel_origin(thread_id: str | None) -> _ChannelOrigin | None:
|
||||
"""Return the channel origin remembered for ``thread_id``, or ``None``."""
|
||||
if not thread_id:
|
||||
return None
|
||||
with _thread_channel_origins_lock:
|
||||
return _thread_channel_origins.get(thread_id)
|
||||
|
||||
|
||||
def forget_channel_origin(thread_id: str | None) -> None:
|
||||
"""Drop the registry entry for ``thread_id`` (e.g. on ``/new`` rotation)."""
|
||||
if not thread_id:
|
||||
return
|
||||
with _thread_channel_origins_lock:
|
||||
_thread_channel_origins.pop(thread_id, None)
|
||||
|
||||
|
||||
def publish_to_channel_origin(thread_id: str | None, content: str) -> bool:
|
||||
"""Schedule pushing ``content`` to the channel remembered for ``thread_id``.
|
||||
|
||||
Fire-and-forget: returns ``True`` iff a publish coroutine was scheduled
|
||||
on the bus loop; returns ``False`` if no origin is registered, the bus
|
||||
isn't running, ``content`` is empty/whitespace, or scheduling itself
|
||||
fails. The publish runs asynchronously — failures inside the coroutine
|
||||
are logged via a done-callback so callers (which are often on event
|
||||
loops that must not block) don't pay any latency.
|
||||
"""
|
||||
from ..channels.bus.events import OutboundMessage
|
||||
|
||||
if not content or not content.strip():
|
||||
return False
|
||||
origin = get_channel_origin(thread_id)
|
||||
if origin is None:
|
||||
return False
|
||||
loop = _bus_loop
|
||||
manager = _manager
|
||||
if loop is None or manager is None:
|
||||
return False
|
||||
bus = getattr(manager, "bus", None)
|
||||
if bus is None:
|
||||
return False
|
||||
|
||||
async def _publish_and_record() -> None:
|
||||
await bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=origin.channel_type,
|
||||
chat_id=origin.chat_id,
|
||||
content=content,
|
||||
metadata=origin.metadata or {},
|
||||
)
|
||||
)
|
||||
# Mirror the normal channel-reply path, which records a "sent"
|
||||
# message after a successful publish so per-channel stats stay
|
||||
# accurate for forwarded notifications too.
|
||||
manager.record_message(origin.channel_type, "sent")
|
||||
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(_publish_and_record(), loop)
|
||||
except Exception as exc:
|
||||
_channel_logger.warning(
|
||||
"Async notification publish to %s:%s failed to schedule: %s",
|
||||
origin.channel_type,
|
||||
origin.chat_id,
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
def _on_publish_done(fut) -> None:
|
||||
"""Log any exception raised by the fire-and-forget publish coroutine."""
|
||||
# A cancelled future raises CancelledError from .exception() rather
|
||||
# than returning it (e.g. bus loop torn down mid-publish); treat that
|
||||
# as a benign shutdown, not a failure to log.
|
||||
if fut.cancelled():
|
||||
return
|
||||
exc = fut.exception()
|
||||
if exc is not None:
|
||||
_channel_logger.warning(
|
||||
"Async notification publish to %s:%s failed: %s",
|
||||
origin.channel_type,
|
||||
origin.chat_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
future.add_done_callback(_on_publish_done)
|
||||
return True
|
||||
|
||||
|
||||
def _is_stop_command(content: str | None) -> bool:
|
||||
"""Whether incoming content is a stop/cancel slash command."""
|
||||
return (content or "").strip().lower() in _STOP_COMMANDS
|
||||
|
||||
|
||||
def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event:
|
||||
@@ -154,7 +632,7 @@ def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
|
||||
|
||||
def channel_ask_user_prompt(
|
||||
ask_user_data: dict,
|
||||
msg: "ChannelMessage | None" = None,
|
||||
msg: ChannelMessage | None = None,
|
||||
) -> dict:
|
||||
"""Format ask_user questions and collect answers from a channel user.
|
||||
|
||||
@@ -186,7 +664,7 @@ def channel_ask_user_prompt(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
metadata=msg.metadata or {},
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
@@ -243,6 +721,8 @@ def channel_ask_user_prompt(
|
||||
return {"status": "cancelled"}
|
||||
|
||||
raw = reply_text.strip()
|
||||
if _is_stop_command(raw):
|
||||
return {"status": "cancelled"}
|
||||
if raw.lower() == "cancel":
|
||||
return {"status": "cancelled"}
|
||||
|
||||
@@ -260,6 +740,8 @@ def channel_ask_user_prompt(
|
||||
if not replied or not other_text:
|
||||
_send("\u23f0 Response timed out.")
|
||||
return {"status": "cancelled"}
|
||||
if _is_stop_command(other_text):
|
||||
return {"status": "cancelled"}
|
||||
if other_text.strip().lower() == "cancel":
|
||||
return {"status": "cancelled"}
|
||||
answers.append(other_text.strip())
|
||||
@@ -279,7 +761,7 @@ def channel_ask_user_prompt(
|
||||
|
||||
def channel_hitl_prompt(
|
||||
action_requests: list,
|
||||
msg: "ChannelMessage",
|
||||
msg: ChannelMessage,
|
||||
) -> list[dict] | None:
|
||||
"""Send HITL approval prompt to channel user and wait for reply.
|
||||
|
||||
@@ -290,6 +772,7 @@ def channel_hitl_prompt(
|
||||
"""
|
||||
from ..channels.bus.events import OutboundMessage
|
||||
from ..channels.consumer import (
|
||||
_approval_prompt_metadata,
|
||||
_format_approval_prompt,
|
||||
_parse_approval_reply,
|
||||
)
|
||||
@@ -304,7 +787,17 @@ def channel_hitl_prompt(
|
||||
_channel_logger.debug("HITL: no bus_loop or bus_ref, rejecting")
|
||||
return None
|
||||
|
||||
def _send(content: str) -> bool:
|
||||
# Look up the channel instance so we can attach buttons when the channel
|
||||
# supports `inline_buttons` (Feishu cards, QQ keyboards, …).
|
||||
channel_obj = (
|
||||
_manager.get_channel(msg.channel_type) if _manager is not None else None
|
||||
)
|
||||
has_buttons = channel_obj is not None and channel_obj.capabilities.inline_buttons
|
||||
approval_metadata = _approval_prompt_metadata(
|
||||
msg.metadata, with_buttons=has_buttons
|
||||
)
|
||||
|
||||
def _send(content: str, *, metadata: dict | None = None) -> bool:
|
||||
"""Send a message to the channel user. Returns True on success."""
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
@@ -313,7 +806,9 @@ def channel_hitl_prompt(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
metadata=metadata
|
||||
if metadata is not None
|
||||
else msg.metadata or {},
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
@@ -324,8 +819,8 @@ def channel_hitl_prompt(
|
||||
return False
|
||||
|
||||
# 1. Send approval prompt
|
||||
prompt_text = _format_approval_prompt(action_requests)
|
||||
if not _send(prompt_text):
|
||||
prompt_text = _format_approval_prompt(action_requests, with_buttons=has_buttons)
|
||||
if not _send(prompt_text, metadata=approval_metadata):
|
||||
return None
|
||||
|
||||
# 2. Wait for channel user's reply
|
||||
@@ -337,16 +832,24 @@ def channel_hitl_prompt(
|
||||
_send("\u23f0 Approval timed out. Action rejected.")
|
||||
return None
|
||||
|
||||
if _is_stop_command(reply_text):
|
||||
# `/stop` already got its own immediate ack from the bus fast-path.
|
||||
# Treat it as a pure cancel signal here so we don't send a second,
|
||||
# contradictory "Unrecognized reply" message.
|
||||
return None
|
||||
|
||||
# 3. Parse decision
|
||||
decision = _parse_approval_reply(reply_text)
|
||||
if decision == "auto":
|
||||
_hitl_auto_approve.add(session_key)
|
||||
_send("\u2705 已批准(后续自动通过)")
|
||||
return [{"type": "approve"} for _ in action_requests]
|
||||
if decision == "approve":
|
||||
_send("\u2705 已批准")
|
||||
return [{"type": "approve"} for _ in action_requests]
|
||||
|
||||
feedback = (
|
||||
"Action rejected."
|
||||
"\u274c 已拒绝"
|
||||
if decision == "reject"
|
||||
else "Unrecognized reply. Action rejected."
|
||||
)
|
||||
@@ -361,8 +864,6 @@ def channel_hitl_prompt(
|
||||
_manager: Any | None = None # ChannelManager
|
||||
_bus_loop: asyncio.AbstractEventLoop | None = None
|
||||
_bus_thread: threading.Thread | None = None
|
||||
_cli_agent: Any = None # shared agent reference (same as CLI)
|
||||
_cli_thread_id: str | None = None # shared thread_id (same conversation)
|
||||
|
||||
|
||||
def _channels_is_running(channel_type: str | None = None) -> bool:
|
||||
@@ -380,9 +881,18 @@ def _channels_running_list() -> list[str]:
|
||||
return _manager.running_channels() if _manager else []
|
||||
|
||||
|
||||
def _channels_stop(channel_type: str | None = None) -> None:
|
||||
"""Stop channel(s) and clean up module-level state."""
|
||||
global _manager, _bus_loop, _bus_thread, _cli_agent, _cli_thread_id
|
||||
def _channels_stop(
|
||||
channel_type: str | None = None,
|
||||
*,
|
||||
runtime: ChannelRuntime | None = None,
|
||||
) -> None:
|
||||
"""Stop channel(s) and clean up module-level state.
|
||||
|
||||
``runtime`` is the ``ChannelRuntime`` whose binding should be
|
||||
cleared once the channels are gone — the caller owns it (commands
|
||||
keep a reference via ``ctx.channel_runtime``).
|
||||
"""
|
||||
global _manager, _bus_loop, _bus_thread
|
||||
|
||||
if channel_type is None:
|
||||
# Stop everything
|
||||
@@ -395,15 +905,13 @@ def _channels_stop(channel_type: str | None = None) -> None:
|
||||
future.result(timeout=10)
|
||||
except Exception as e:
|
||||
_channel_logger.debug(f"Error stopping channels: {e}")
|
||||
if _manager:
|
||||
_manager.bus.stop()
|
||||
if _bus_thread:
|
||||
_bus_thread.join(timeout=5)
|
||||
_manager = None
|
||||
_bus_loop = None
|
||||
_bus_thread = None
|
||||
_cli_agent = None
|
||||
_cli_thread_id = None
|
||||
if runtime is not None:
|
||||
runtime.clear()
|
||||
return
|
||||
|
||||
# Stop a specific channel
|
||||
@@ -417,9 +925,8 @@ def _channels_stop(channel_type: str | None = None) -> None:
|
||||
except Exception as e:
|
||||
_channel_logger.debug(f"Error removing channel {channel_type}: {e}")
|
||||
|
||||
if _manager and not _manager.running_channels():
|
||||
_cli_agent = None
|
||||
_cli_thread_id = None
|
||||
if _manager and not _manager.running_channels() and runtime is not None:
|
||||
runtime.clear()
|
||||
|
||||
|
||||
def _start_channels_bus_mode(
|
||||
@@ -533,6 +1040,20 @@ async def _bus_inbound_consumer(bus, manager) -> None:
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
# /stop should preempt HITL interception so cancel works while
|
||||
# waiting for approvals/questions. If a HITL wait is pending,
|
||||
# still release it so the blocking prompt can unwind immediately.
|
||||
if _is_stop_command(msg.content):
|
||||
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
|
||||
_channel_logger.info(
|
||||
f"[bus] stop request released HITL wait for "
|
||||
f"{msg.channel}:{msg.chat_id}"
|
||||
)
|
||||
_task = asyncio.create_task(_handle_bus_message(bus, manager, msg))
|
||||
_tasks.add(_task)
|
||||
_task.add_done_callback(_tasks.discard)
|
||||
continue
|
||||
|
||||
# Check if this message is a HITL approval reply
|
||||
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
|
||||
_channel_logger.info(
|
||||
@@ -561,6 +1082,37 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
||||
)
|
||||
manager.record_message(msg.channel, "received")
|
||||
|
||||
# Fast-path: /stop intercept. Handle on the bus task itself so we
|
||||
# don't deadlock behind the main-thread stream we're trying to
|
||||
# interrupt. No typing indicator, no queue entry.
|
||||
if _is_stop_command(msg.content):
|
||||
cancelled_count, active_count = _cancel_channel_session(
|
||||
msg.channel, msg.chat_id
|
||||
)
|
||||
try:
|
||||
await bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Stopped.",
|
||||
reply_to=msg.message_id or None,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
manager.record_message(msg.channel, "sent")
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"[bus] /stop ack send error: {e}")
|
||||
else:
|
||||
if cancelled_count or active_count:
|
||||
_channel_logger.info(
|
||||
"[bus] /stop cancelled %d request(s) (%d active) for %s:%s",
|
||||
cancelled_count,
|
||||
active_count,
|
||||
msg.channel,
|
||||
msg.chat_id,
|
||||
)
|
||||
return
|
||||
|
||||
channel = manager.get_channel(msg.channel)
|
||||
typing_active = False
|
||||
if channel:
|
||||
@@ -630,6 +1182,8 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
||||
f"for {cm.msg_id}"
|
||||
)
|
||||
_pop_channel_response(cm.msg_id, cancel_pending=True)
|
||||
if _channel_request_state(cm.msg_id) != "active":
|
||||
_complete_channel_request(cm.msg_id)
|
||||
return
|
||||
|
||||
response = _pop_channel_response(cm.msg_id) or "No response"
|
||||
@@ -645,6 +1199,8 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
||||
manager.record_message(msg.channel, "sent")
|
||||
except asyncio.CancelledError:
|
||||
_pop_channel_response(cm.msg_id, cancel_pending=True)
|
||||
if _channel_request_state(cm.msg_id) != "active":
|
||||
_complete_channel_request(cm.msg_id)
|
||||
raise
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"[bus] Outbound error: {e}")
|
||||
@@ -682,131 +1238,13 @@ def _print_channel_panel(channels: list[tuple[str, bool, str]]) -> None:
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_channel(
|
||||
args: str,
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
*,
|
||||
send_thinking: bool | None = None,
|
||||
) -> None:
|
||||
"""Start a channel in background using bus mode.
|
||||
|
||||
Usage:
|
||||
/channel [telegram|discord|imessage] -- start channel (default from config)
|
||||
/channel status -- show current channel status
|
||||
/channel stop -- stop running channel
|
||||
"""
|
||||
global _cli_agent, _cli_thread_id
|
||||
|
||||
from ..config import load_config
|
||||
|
||||
app_config = load_config()
|
||||
|
||||
channel_type = args.strip().lower() if args and args.strip() else ""
|
||||
if channel_type == "status":
|
||||
running = _channels_running_list()
|
||||
if running and _manager:
|
||||
detailed = _manager.get_detailed_status()
|
||||
table = Table(title="Channel Status", show_header=True, expand=False)
|
||||
table.add_column("Channel", style="cyan")
|
||||
table.add_column("Status")
|
||||
table.add_column("Uptime", style="dim")
|
||||
table.add_column("Rx", justify="right")
|
||||
table.add_column("Tx", justify="right")
|
||||
for ch_name in running:
|
||||
info = detailed.get(ch_name, {})
|
||||
secs = info.get("uptime_seconds", 0)
|
||||
mins, s = divmod(int(secs), 60)
|
||||
hours, mins = divmod(mins, 60)
|
||||
uptime = f"{hours}h{mins:02d}m" if hours else f"{mins}m{s:02d}s"
|
||||
rx = str(info.get("received", 0))
|
||||
tx = str(info.get("sent", 0))
|
||||
table.add_row(ch_name, "[green]running[/green]", uptime, rx, tx)
|
||||
console.print(table)
|
||||
console.print()
|
||||
else:
|
||||
console.print("[dim]No channel running[/dim]\n")
|
||||
return
|
||||
|
||||
if not channel_type:
|
||||
channel_type = app_config.channel_enabled
|
||||
if not channel_type:
|
||||
console.print("[yellow]No channel configured.[/yellow]")
|
||||
console.print(
|
||||
"[dim]Run[/dim] evosci onboard [dim]or specify:[/dim] /channel telegram\n"
|
||||
)
|
||||
return
|
||||
|
||||
requested = [t.strip() for t in channel_type.split(",") if t.strip()]
|
||||
|
||||
if _channels_is_running():
|
||||
running = _channels_running_list()
|
||||
results: list[tuple[str, bool, str]] = []
|
||||
for ct in requested:
|
||||
if ct in running:
|
||||
results.append((ct, True, "already running"))
|
||||
else:
|
||||
try:
|
||||
_add_channel_to_running_bus(
|
||||
ct,
|
||||
app_config,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
results.append((ct, True, "connected (bus)"))
|
||||
except Exception as e:
|
||||
results.append((ct, False, str(e)))
|
||||
_print_channel_panel(results)
|
||||
return
|
||||
|
||||
_cli_agent = agent
|
||||
_cli_thread_id = thread_id
|
||||
|
||||
# Override channel_enabled for this invocation
|
||||
original = app_config.channel_enabled
|
||||
app_config.channel_enabled = channel_type
|
||||
try:
|
||||
_start_channels_bus_mode(
|
||||
app_config,
|
||||
agent,
|
||||
thread_id,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
results = [(ct, True, "connected (bus)") for ct in requested]
|
||||
except Exception as e:
|
||||
results = [(ct, False, str(e)) for ct in requested]
|
||||
finally:
|
||||
app_config.channel_enabled = original
|
||||
|
||||
_print_channel_panel(results)
|
||||
|
||||
|
||||
def _cmd_channel_stop(channel_type: str | None = None) -> None:
|
||||
"""Stop background channel(s).
|
||||
|
||||
Args:
|
||||
channel_type: Specific channel to stop, or None to stop all.
|
||||
"""
|
||||
if not _channels_is_running():
|
||||
console.print("[dim]No channel running[/dim]\n")
|
||||
return
|
||||
if channel_type:
|
||||
if not _channels_is_running(channel_type):
|
||||
console.print(f"[dim]{channel_type} is not running[/dim]\n")
|
||||
return
|
||||
_channels_stop(channel_type)
|
||||
console.print(f"[dim]{channel_type} stopped[/dim]\n")
|
||||
else:
|
||||
running = _channels_running_list()
|
||||
_channels_stop()
|
||||
console.print(f"[dim]{', '.join(running)} stopped[/dim]\n")
|
||||
|
||||
|
||||
def _auto_start_channel(
|
||||
agent: Any,
|
||||
thread_id: str,
|
||||
config,
|
||||
*,
|
||||
send_thinking: bool | None = None,
|
||||
runtime: ChannelRuntime | None = None,
|
||||
) -> None:
|
||||
"""Start channels automatically from config (bus mode).
|
||||
|
||||
@@ -814,21 +1252,23 @@ def _auto_start_channel(
|
||||
agent: Compiled agent graph.
|
||||
thread_id: Current thread ID.
|
||||
config: EvoScientistConfig with channel settings.
|
||||
runtime: Caller-owned ``ChannelRuntime`` to bind so commands
|
||||
running over the channels can swap the agent later. ``None``
|
||||
is accepted for callers that don't yet pass one.
|
||||
"""
|
||||
global _cli_agent, _cli_thread_id
|
||||
|
||||
if not config.channel_enabled:
|
||||
return
|
||||
|
||||
_cli_agent = agent
|
||||
_cli_thread_id = thread_id
|
||||
|
||||
_start_channels_bus_mode(
|
||||
config,
|
||||
agent,
|
||||
thread_id,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
# Bind only after startup succeeds; a failure above must not leave
|
||||
# a stale runtime binding pointing at channels that never started.
|
||||
if runtime is not None:
|
||||
runtime.bind(agent, thread_id)
|
||||
types = [t.strip() for t in config.channel_enabled.split(",") if t.strip()]
|
||||
results = [(ct, True, "connected (bus)") for ct in types]
|
||||
_print_channel_panel(results)
|
||||
|
||||
@@ -21,7 +21,7 @@ import os
|
||||
import pathlib
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from textual.app import App
|
||||
@@ -29,6 +29,15 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_PREVIEW_MAX = 40
|
||||
_pyperclip_notify_shown = False
|
||||
|
||||
|
||||
def _is_remote_session() -> bool:
|
||||
"""Return True when running over SSH without a local display."""
|
||||
if os.environ.get("SSH_CLIENT") or os.environ.get("SSH_CONNECTION"):
|
||||
if not os.environ.get("DISPLAY") and not os.environ.get("WAYLAND_DISPLAY"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# ── Platform-native clipboard read ────────────────────────────────
|
||||
@@ -132,8 +141,6 @@ def copy_selection_to_clipboard(app: App) -> None:
|
||||
if not hasattr(widget, "text_selection") or not widget.text_selection:
|
||||
continue
|
||||
selection = widget.text_selection
|
||||
if selection.end is None:
|
||||
continue
|
||||
try:
|
||||
result = widget.get_selection(selection)
|
||||
except (AttributeError, TypeError, ValueError, IndexError) as exc:
|
||||
@@ -154,37 +161,59 @@ def copy_selection_to_clipboard(app: App) -> None:
|
||||
|
||||
combined = "\n".join(selected_texts)
|
||||
|
||||
# Try methods in priority order
|
||||
copy_methods = [app.copy_to_clipboard]
|
||||
# Build method list: (fn, reliable) — reliable means we *know* the text
|
||||
# reached the system clipboard (e.g. pyperclip). OSC 52 / Textual write
|
||||
# to the terminal and succeed even when the terminal silently ignores the
|
||||
# sequence (PuTTY, older terminals).
|
||||
copy_methods: list[tuple[Any, bool]] = [
|
||||
(app.copy_to_clipboard, False),
|
||||
]
|
||||
|
||||
try:
|
||||
import pyperclip
|
||||
|
||||
copy_methods.insert(0, pyperclip.copy)
|
||||
copy_methods.insert(0, (pyperclip.copy, True))
|
||||
except ImportError:
|
||||
pass
|
||||
global _pyperclip_notify_shown
|
||||
if not _pyperclip_notify_shown:
|
||||
_pyperclip_notify_shown = True
|
||||
app.notify(
|
||||
'Failed to import "pyperclip", text copying might not work.',
|
||||
severity="information",
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
copy_methods.append(_copy_osc52)
|
||||
copy_methods.append((_copy_osc52, False))
|
||||
|
||||
for fn in copy_methods:
|
||||
remote = _is_remote_session()
|
||||
|
||||
for fn, reliable in copy_methods:
|
||||
try:
|
||||
fn(combined)
|
||||
except (OSError, RuntimeError, TypeError) as exc:
|
||||
logger.debug(
|
||||
"Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc
|
||||
)
|
||||
continue
|
||||
|
||||
if reliable or not remote:
|
||||
app.notify(
|
||||
f'"{_shorten(selected_texts)}" copied',
|
||||
severity="information",
|
||||
timeout=2,
|
||||
markup=False,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError) as exc:
|
||||
logger.debug(
|
||||
"Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc
|
||||
)
|
||||
continue
|
||||
else:
|
||||
# OSC 52 over SSH — may be silently ignored (e.g. Windows Terminal, PuTTY)
|
||||
app.notify(
|
||||
"Copied text - if paste fails, use Shift+mouse-select for native copy",
|
||||
severity="information",
|
||||
timeout=3,
|
||||
)
|
||||
return
|
||||
|
||||
app.notify(
|
||||
"Failed to copy — no clipboard method available",
|
||||
"Copy failed — use Shift+mouse-select for native terminal copy",
|
||||
severity="warning",
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
@@ -19,16 +19,35 @@ from pathlib import Path
|
||||
|
||||
_PATH_CHARS = r"A-Za-z0-9._~/\\:-"
|
||||
|
||||
FILE_MENTION_PATTERN = re.compile(r"@(?P<path>(?:\\.|[" + _PATH_CHARS + r"])+)")
|
||||
FILE_MENTION_PATTERN = re.compile(
|
||||
r"@(?:"
|
||||
r'"(?P<dquoted>[^"\n]+)"'
|
||||
r"|'(?P<squoted>[^'\n]+)'"
|
||||
r"|(?P<bare>(?:\\.|[" + _PATH_CHARS + r"])+)"
|
||||
r")"
|
||||
)
|
||||
"""Matches ``@path/to/file`` in user input.
|
||||
|
||||
Escaped spaces (``@my\\\\ folder/file``) are supported. Bare ``@`` with no
|
||||
path characters is not matched (uses ``+`` not ``*``).
|
||||
Three forms are supported, in priority order:
|
||||
|
||||
1. ``@"path with spaces.pdf"`` — explicit double-quoted path
|
||||
2. ``@'path with spaces.pdf'`` — explicit single-quoted path
|
||||
3. ``@bare/path`` — backslash-escaped spaces (``@my\\\\ folder/file``) work;
|
||||
raw unescaped spaces are handled via greedy expansion in
|
||||
:func:`parse_file_mentions`.
|
||||
|
||||
Bare ``@`` with no path characters is not matched (uses ``+`` not ``*``).
|
||||
"""
|
||||
|
||||
_EMAIL_PREFIX = re.compile(r"[a-zA-Z0-9._%+-]$")
|
||||
"""If the character immediately before ``@`` matches this, it's an email address."""
|
||||
|
||||
# Hard cap on tokens consumed during greedy expansion across whitespace.
|
||||
_GREEDY_MAX_TOKENS = 20
|
||||
|
||||
# Trailing punctuation stripped before checking if a greedy candidate exists.
|
||||
_GREEDY_TRAIL_PUNCT = ",;:!?)]}>"
|
||||
|
||||
# Files larger than this are referenced by path only (not embedded inline).
|
||||
_MAX_EMBED_BYTES = 256 * 1024 # 256 KB
|
||||
|
||||
@@ -193,6 +212,75 @@ def _read_file(path: Path) -> str:
|
||||
return f"\n### {path.name}\nPath: `{path}`\n```\n{content}\n```"
|
||||
|
||||
|
||||
def _resolve_path(raw: str, cwd: Path) -> Path | None:
|
||||
"""Resolve *raw* to an existing file path, or ``None``.
|
||||
|
||||
Honors backslash-escaped spaces and ``~`` expansion. Returns ``None``
|
||||
when the path does not exist, is not a regular file, or raises
|
||||
``OSError``/``RuntimeError`` during resolution.
|
||||
"""
|
||||
clean = raw.replace("\\ ", " ")
|
||||
try:
|
||||
p = Path(clean).expanduser()
|
||||
if not p.is_absolute():
|
||||
p = cwd / p
|
||||
resolved = p.resolve()
|
||||
except (OSError, RuntimeError):
|
||||
return None
|
||||
if resolved.is_file():
|
||||
return resolved
|
||||
return None
|
||||
|
||||
|
||||
def _greedy_extend(
|
||||
text: str,
|
||||
raw: str,
|
||||
match_end: int,
|
||||
cwd: Path,
|
||||
) -> tuple[str, Path, int] | None:
|
||||
"""Try to extend *raw* across whitespace until the path resolves.
|
||||
|
||||
Walks the text after *match_end*, capped at the next newline or the
|
||||
start of another ``@`` mention. Tries the longest plausible suffix
|
||||
first and shrinks one token at a time, stripping trailing punctuation
|
||||
that is unlikely to be part of a filename.
|
||||
|
||||
Returns ``(extended_raw, resolved_file, new_end_pos)`` on success,
|
||||
else ``None``.
|
||||
"""
|
||||
rest = text[match_end:]
|
||||
|
||||
# Hard boundaries that should never be crossed.
|
||||
boundary = len(rest)
|
||||
nl = rest.find("\n")
|
||||
if nl >= 0:
|
||||
boundary = nl
|
||||
next_at = re.search(r"\s@", rest[:boundary])
|
||||
if next_at:
|
||||
boundary = next_at.start()
|
||||
|
||||
region = rest[:boundary]
|
||||
if not region or not region[0].isspace():
|
||||
return None
|
||||
|
||||
tokens = list(re.finditer(r"\S+", region))
|
||||
if not tokens:
|
||||
return None
|
||||
|
||||
# Try longest-first so we prefer the most specific match.
|
||||
for i in range(min(len(tokens), _GREEDY_MAX_TOKENS), 0, -1):
|
||||
end = tokens[i - 1].end()
|
||||
suffix = region[:end].rstrip(_GREEDY_TRAIL_PUNCT)
|
||||
if not suffix:
|
||||
continue
|
||||
candidate = raw + suffix
|
||||
resolved = _resolve_path(candidate, cwd)
|
||||
if resolved is not None:
|
||||
return candidate, resolved, match_end + len(suffix)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def parse_file_mentions(
|
||||
text: str,
|
||||
cwd: Path | None = None,
|
||||
@@ -218,23 +306,39 @@ def parse_file_mentions(
|
||||
files: list[Path] = []
|
||||
warnings: list[str] = []
|
||||
seen: set[Path] = set()
|
||||
# finditer would normally re-scan from each match's end, but greedy
|
||||
# expansion can consume bytes past that point. Track a manual cursor
|
||||
# and skip matches that start before it.
|
||||
cursor = 0
|
||||
for match in FILE_MENTION_PATTERN.finditer(text):
|
||||
if match.start() < cursor:
|
||||
continue
|
||||
# Skip email addresses — character immediately before @ is alphanumeric
|
||||
before = text[: match.start()]
|
||||
if before and _EMAIL_PREFIX.search(before):
|
||||
continue
|
||||
|
||||
raw = match.group("path")
|
||||
clean = raw.replace("\\ ", " ")
|
||||
dquoted = match.group("dquoted")
|
||||
squoted = match.group("squoted")
|
||||
bare = match.group("bare")
|
||||
quoted_raw = dquoted if dquoted is not None else squoted
|
||||
raw = quoted_raw if quoted_raw is not None else bare
|
||||
is_quoted = quoted_raw is not None
|
||||
|
||||
try:
|
||||
p = Path(clean).expanduser()
|
||||
if not p.is_absolute():
|
||||
p = cwd / p
|
||||
resolved = p.resolve()
|
||||
if not resolved.exists() or not resolved.is_file():
|
||||
resolved = _resolve_path(raw, cwd)
|
||||
end_pos = match.end()
|
||||
|
||||
if resolved is None and not is_quoted:
|
||||
extended = _greedy_extend(text, raw, match.end(), cwd)
|
||||
if extended is not None:
|
||||
raw, resolved, end_pos = extended
|
||||
|
||||
cursor = end_pos
|
||||
|
||||
if resolved is None:
|
||||
warnings.append(f"@file not found: {raw}")
|
||||
continue
|
||||
|
||||
# Deduplicate: skip paths already seen in this message.
|
||||
if resolved in seen:
|
||||
continue
|
||||
@@ -250,8 +354,6 @@ def parse_file_mentions(
|
||||
f"@{raw} is outside the workspace "
|
||||
f"({workspace_root}) — embedding may expose sensitive files"
|
||||
)
|
||||
except (OSError, RuntimeError) as exc:
|
||||
warnings.append(f"invalid @file path {raw!r}: {exc}")
|
||||
|
||||
return files, warnings
|
||||
|
||||
@@ -302,6 +404,13 @@ def _type_hint(rel_path: str) -> str:
|
||||
return suffix or "file"
|
||||
|
||||
|
||||
def _format_mention(rel_path: str) -> str:
|
||||
"""Render *rel_path* as an ``@`` mention, quoting if it contains spaces."""
|
||||
if " " in rel_path:
|
||||
return f'@"{rel_path}"'
|
||||
return f"@{rel_path}"
|
||||
|
||||
|
||||
def complete_file_mention(
|
||||
text: str,
|
||||
workspace_dir: str | None = None,
|
||||
@@ -320,13 +429,22 @@ def complete_file_mention(
|
||||
List of ``(completion_string, type_hint)`` tuples, e.g.
|
||||
``[("@results/v2.json", "json"), ("@README.md", "md")]``.
|
||||
Directories have a trailing ``/`` and type hint ``"dir"``.
|
||||
Paths containing spaces are returned in double-quoted form,
|
||||
e.g. ``@"my docs/file.pdf"``.
|
||||
"""
|
||||
# Find the last @token
|
||||
match = re.search(r"@([^\s]*)$", text)
|
||||
# Find the last @token. Allow whitespace inside a quoted partial so
|
||||
# completion keeps working as the user types ``@"PRE`` → ``@"PREPING_ B``.
|
||||
quoted_match = re.search(r'@"([^"\n]*)$', text)
|
||||
if quoted_match:
|
||||
partial = quoted_match.group(1)
|
||||
quoted = True
|
||||
else:
|
||||
match = re.search(r"@([^\s\"']*)$", text)
|
||||
if not match:
|
||||
return []
|
||||
|
||||
partial = match.group(1).replace("\\ ", " ")
|
||||
quoted = False
|
||||
|
||||
base_str = workspace_dir or str(Path.cwd())
|
||||
base = Path(base_str)
|
||||
|
||||
@@ -344,10 +462,10 @@ def complete_file_mention(
|
||||
rel = entry.relative_to(base)
|
||||
suffix = "/" if entry.is_dir() else ""
|
||||
candidates_raw.append(rel.as_posix() + suffix)
|
||||
except OSError:
|
||||
except (OSError, ValueError):
|
||||
return []
|
||||
return [
|
||||
(f"@{r}", "dir" if r.endswith("/") else _type_hint(r))
|
||||
(_format_mention(r), "dir" if r.endswith("/") else _type_hint(r))
|
||||
for r in candidates_raw[:10]
|
||||
]
|
||||
|
||||
@@ -363,16 +481,24 @@ def complete_file_mention(
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
combined = all_files + dir_candidates
|
||||
|
||||
# Determine query: if partial has a slash, search within that subtree
|
||||
if "/" in partial:
|
||||
# Filter candidates to those starting with the directory prefix
|
||||
# Search within the subtree for the given directory prefix
|
||||
combined = all_files + dir_candidates
|
||||
dir_prefix = partial.rsplit("/", 1)[0] + "/"
|
||||
file_query = partial.rsplit("/", 1)[1]
|
||||
subtree = [c for c in combined if c.startswith(dir_prefix)]
|
||||
results = _fuzzy_search(file_query, subtree)
|
||||
else:
|
||||
results = _fuzzy_search(partial, combined)
|
||||
# Depth 1 only: top-level files and directories
|
||||
top_files = [f for f in all_files if "/" not in f]
|
||||
results = _fuzzy_search(partial, top_files + dir_candidates)
|
||||
|
||||
return [(f"@{r}", "dir" if r.endswith("/") else _type_hint(r)) for r in results]
|
||||
if quoted:
|
||||
# User opened a quoted mention — close it for them.
|
||||
return [
|
||||
(f'@"{r}"', "dir" if r.endswith("/") else _type_hint(r)) for r in results
|
||||
]
|
||||
return [
|
||||
(_format_mention(r), "dir" if r.endswith("/") else _type_hint(r))
|
||||
for r in results
|
||||
]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""History-based auto-suggest for Textual TUI Input widget.
|
||||
|
||||
Reads prompt_toolkit FileHistory format so Rich CLI and TUI share the same
|
||||
history file at ~/.config/ai4scientist/history.
|
||||
history file at ~/.evoscientist/history.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -9,7 +9,6 @@ from __future__ import annotations
|
||||
from collections import Counter
|
||||
|
||||
import questionary
|
||||
from prompt_toolkit.styles import Style as PtStyle
|
||||
from questionary import Choice
|
||||
|
||||
from ..mcp.registry import (
|
||||
@@ -21,18 +20,8 @@ from ..mcp.registry import (
|
||||
install_mcp_server,
|
||||
install_mcp_servers,
|
||||
)
|
||||
from ..stream.display import console
|
||||
|
||||
_PICKER_STYLE = PtStyle.from_dict(
|
||||
{
|
||||
"questionmark": "#888888",
|
||||
"question": "",
|
||||
"pointer": "bold",
|
||||
"highlighted": "bold",
|
||||
"text": "#888888",
|
||||
"answer": "bold",
|
||||
}
|
||||
)
|
||||
from ..stream.console import console
|
||||
from .widgets.thread_selector import PICKER_STYLE
|
||||
|
||||
_INSTALLED_INDICATOR = ("fg:#4caf50", "\u2713 ")
|
||||
|
||||
@@ -57,7 +46,7 @@ def _checkbox_ask(choices, message: str, **kwargs):
|
||||
return questionary.checkbox(
|
||||
message,
|
||||
choices=choices,
|
||||
style=_PICKER_STYLE,
|
||||
style=PICKER_STYLE,
|
||||
qmark="\u276f",
|
||||
**kwargs,
|
||||
).ask()
|
||||
@@ -100,7 +89,7 @@ def _browse_and_select(
|
||||
selected_tag = questionary.select(
|
||||
"Filter by tag:",
|
||||
choices=tag_choices,
|
||||
style=_PICKER_STYLE,
|
||||
style=PICKER_STYLE,
|
||||
qmark="\u276f",
|
||||
).ask()
|
||||
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
"""MCP server display, operations, and /mcp slash-command dispatcher."""
|
||||
"""Shared UI helpers for MCP server display and operations (used by the Typer `mcp` commands)."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ..stream.display import console
|
||||
from ..stream.console import console
|
||||
|
||||
|
||||
def _mcp_list_servers() -> None:
|
||||
@@ -181,129 +181,3 @@ def _show_mcp_config(name: str = "", *, show_blank_line: bool = True) -> str:
|
||||
if show_blank_line:
|
||||
console.print()
|
||||
return "ok"
|
||||
|
||||
|
||||
def _cmd_mcp_add(args_str: str) -> None:
|
||||
"""Handle ``/mcp add ...``."""
|
||||
import shlex
|
||||
|
||||
from ..mcp import parse_mcp_add_args
|
||||
|
||||
if not args_str.strip():
|
||||
console.print("[bold]Usage:[/bold] /mcp add <name> <command-or-url> [args...]")
|
||||
console.print()
|
||||
console.print(
|
||||
"[dim]Transport is auto-detected: URLs \u2192 http, commands \u2192 stdio[/dim]"
|
||||
)
|
||||
console.print()
|
||||
console.print("[bold]Examples:[/bold]")
|
||||
console.print(
|
||||
" /mcp add sequential-thinking npx -y @modelcontextprotocol/server-sequential-thinking"
|
||||
)
|
||||
console.print(" /mcp add docs-langchain https://docs.langchain.com/mcp")
|
||||
console.print(
|
||||
" /mcp add my-sse http://localhost:9090/sse --transport sse --expose-to research-agent"
|
||||
)
|
||||
console.print()
|
||||
console.print("[dim]Options:[/dim]")
|
||||
console.print(" --transport T Transport type (default: auto-detect)")
|
||||
console.print(
|
||||
" --tools t1,t2 Tool allowlist (supports wildcards: *_exa, read_*)"
|
||||
)
|
||||
console.print(" --expose-to a1,a2 Target agents (default: main)")
|
||||
console.print(" --header Key:Value HTTP header (repeatable)")
|
||||
console.print(" --env KEY=VALUE Env var for stdio (repeatable)")
|
||||
console.print(
|
||||
" --env-ref KEY Env var as runtime ${KEY} reference (repeatable)"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
try:
|
||||
tokens = shlex.split(args_str)
|
||||
kwargs = parse_mcp_add_args(tokens)
|
||||
_mcp_add_server_from_kwargs(kwargs, show_reload_hint=True)
|
||||
except ValueError as exc:
|
||||
console.print(f"[red]{exc}[/red]")
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_mcp_edit(args_str: str) -> None:
|
||||
"""Handle ``/mcp edit <name> --field value ...``."""
|
||||
import shlex
|
||||
|
||||
from ..mcp import parse_mcp_edit_args
|
||||
|
||||
if not args_str.strip():
|
||||
console.print("[bold]Usage:[/bold] /mcp edit <name> --<field> <value> ...")
|
||||
console.print()
|
||||
console.print(
|
||||
"[dim]Fields:[/dim] --transport, --command, --url, --args, --tools, --expose-to, --header, --env"
|
||||
)
|
||||
console.print(
|
||||
"[dim]Use[/dim] --tools none [dim]or[/dim] --expose-to none [dim]to clear a field.[/dim]"
|
||||
)
|
||||
console.print()
|
||||
console.print("[bold]Examples:[/bold]")
|
||||
console.print(" /mcp edit filesystem --expose-to main,code-agent")
|
||||
console.print(" /mcp edit filesystem --tools read_file,write_file")
|
||||
console.print(" /mcp edit my-api --url http://new-host:8080/mcp")
|
||||
console.print(" /mcp edit my-api --tools none")
|
||||
console.print()
|
||||
return
|
||||
|
||||
try:
|
||||
tokens = shlex.split(args_str)
|
||||
name, fields = parse_mcp_edit_args(tokens)
|
||||
_mcp_edit_server_fields(name, fields, show_reload_hint=True)
|
||||
except ValueError as exc:
|
||||
console.print(f"[red]{exc}[/red]")
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_mcp_remove(name: str) -> None:
|
||||
"""Handle ``/mcp remove <name>``."""
|
||||
_mcp_remove_server(name, show_reload_hint=True)
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_mcp_config(name: str) -> None:
|
||||
"""Handle ``/mcp config [name]``."""
|
||||
_show_mcp_config(name, show_blank_line=True)
|
||||
|
||||
|
||||
def _cmd_mcp(args: str) -> None:
|
||||
"""Dispatch ``/mcp`` subcommands."""
|
||||
args = args.strip()
|
||||
if not args:
|
||||
_mcp_list_servers()
|
||||
return
|
||||
|
||||
parts = args.split(maxsplit=1)
|
||||
subcmd = parts[0].lower()
|
||||
subargs = parts[1] if len(parts) > 1 else ""
|
||||
|
||||
if subcmd == "list":
|
||||
_mcp_list_servers()
|
||||
elif subcmd == "add":
|
||||
_cmd_mcp_add(subargs)
|
||||
elif subcmd == "edit":
|
||||
_cmd_mcp_edit(subargs)
|
||||
elif subcmd == "remove":
|
||||
_cmd_mcp_remove(subargs)
|
||||
elif subcmd == "config":
|
||||
_cmd_mcp_config(subargs)
|
||||
elif subcmd == "install":
|
||||
from .mcp_install_cmd import _cmd_install_mcp
|
||||
|
||||
_cmd_install_mcp(subargs)
|
||||
else:
|
||||
console.print("[bold]MCP commands:[/bold]")
|
||||
console.print(" /mcp List configured servers")
|
||||
console.print(" /mcp list List configured servers")
|
||||
console.print(" /mcp config Show detailed server config")
|
||||
console.print(" /mcp add ... Add a server")
|
||||
console.print(" /mcp edit ... Edit an existing server")
|
||||
console.print(" /mcp remove ... Remove a server")
|
||||
console.print(" /mcp install ... Browse and install servers")
|
||||
console.print()
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Helper for printing the session-exit Goodbye message and resume hint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
|
||||
|
||||
def print_resume_hint(
|
||||
thread_id: str | None,
|
||||
console: Console | None = None,
|
||||
) -> None:
|
||||
"""Print ``Goodbye!`` and, when available, a resume hint for *thread_id*."""
|
||||
out = console or Console()
|
||||
out.print("[dim]Goodbye![/dim]")
|
||||
if thread_id:
|
||||
from ..sessions import short_thread_id
|
||||
|
||||
out.print()
|
||||
out.print("[dim]Resume this session with:[/dim]")
|
||||
out.print(f"[cyan]EvoSci --resume {escape(short_thread_id(thread_id))}[/cyan]")
|
||||
@@ -0,0 +1,221 @@
|
||||
"""CommandUI Protocol adapter for the Rich CLI surface.
|
||||
|
||||
Lifecycle methods (``request_quit``, ``force_quit``, ``clear_chat``,
|
||||
``start_new_session``, ``handle_session_resume``, ``update_status_after_compact``)
|
||||
are callback-driven: when their corresponding ``on_*`` constructor kwarg
|
||||
is ``None``, the method is a silent no-op, mirroring
|
||||
``ChannelCommandUI``'s fallback pattern. Callers that need a specific
|
||||
side-effect (REPL quit flag flip, status-bar refresh, …) wire the
|
||||
callback at construction time; non-interactive surfaces (tests,
|
||||
alternate REPLs) can leave callbacks unset without crashing.
|
||||
|
||||
``wait_for_*`` methods return ``None`` on cancel / fallback and are
|
||||
always safe to ``await``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from rich.console import Console
|
||||
|
||||
from ..commands.base import CommandUI
|
||||
|
||||
|
||||
class RichCLICommandUI(CommandUI):
|
||||
"""CommandUI implementation that prints to a Rich ``Console``.
|
||||
|
||||
Commands that affect CLI-closure state (session lifecycle, exit flag,
|
||||
status-bar snapshot) go through optional callbacks wired by the REPL.
|
||||
This mirrors ``ChannelCommandUI``'s injection pattern and keeps
|
||||
``interactive.py``'s ``state`` dict as the single source of truth.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
console: Console,
|
||||
*,
|
||||
on_request_quit: Callable[[], None] | None = None,
|
||||
on_force_quit: Callable[[], None] | None = None,
|
||||
on_clear_chat: Callable[[], None] | None = None,
|
||||
on_status_after_compact: Callable[[int], None] | None = None,
|
||||
on_start_new_session: Callable[[], Awaitable[None]] | None = None,
|
||||
on_handle_session_resume: (
|
||||
Callable[[str, str | None], Awaitable[None]] | None
|
||||
) = None,
|
||||
) -> None:
|
||||
self.console = console
|
||||
self._on_request_quit = on_request_quit
|
||||
self._on_force_quit = on_force_quit
|
||||
self._on_clear_chat = on_clear_chat
|
||||
self._on_status_after_compact = on_status_after_compact
|
||||
self._on_start_new_session = on_start_new_session
|
||||
self._on_handle_session_resume = on_handle_session_resume
|
||||
# Bound ``console.status(...)`` context manager used by
|
||||
# /compact's start/stop indicator pair.
|
||||
self._compact_status_ctx: Any = None
|
||||
|
||||
# ── Core I/O ─────────────────────────────────────────────
|
||||
|
||||
@property
|
||||
def supports_interactive(self) -> bool:
|
||||
return True
|
||||
|
||||
def append_system(self, text: str, style: str = "dim") -> None:
|
||||
self.console.print(text, style=style)
|
||||
|
||||
def mount_renderable(self, renderable: Any) -> None:
|
||||
self.console.print(renderable)
|
||||
|
||||
async def flush(self) -> None:
|
||||
# Rich console flushes synchronously; nothing to await.
|
||||
return
|
||||
|
||||
# ── Interactive pickers ────────────────────────────────
|
||||
|
||||
async def wait_for_thread_pick(
|
||||
self, threads: list[dict], current_thread: str, title: str
|
||||
) -> str | None:
|
||||
"""Interactive workspace-grouped thread picker using ``questionary``.
|
||||
|
||||
Ported from the pre-migration ``_cmd_resume`` implementation.
|
||||
Returns the selected ``thread_id`` string, or ``None`` on cancel.
|
||||
Callers (``ResumeCommand``/``DeleteCommand``) pre-check for
|
||||
empty thread lists before invoking this method.
|
||||
"""
|
||||
import questionary # type: ignore[import-untyped]
|
||||
from prompt_toolkit.layout.dimension import ( # type: ignore[import-untyped]
|
||||
Dimension,
|
||||
)
|
||||
from questionary.prompts.common import ( # type: ignore[import-untyped]
|
||||
InquirerControl,
|
||||
)
|
||||
|
||||
from ..sessions import _format_relative_time
|
||||
from .widgets.thread_selector import PICKER_STYLE, _build_items
|
||||
|
||||
choices: list[Any] = []
|
||||
for item in _build_items(threads):
|
||||
if item["type"] == "header":
|
||||
choices.append(questionary.Separator(f"── \U0001f4c2 {item['label']}"))
|
||||
elif item["type"] == "subheader":
|
||||
choices.append(questionary.Separator(f" {item['label']}"))
|
||||
else:
|
||||
t = item["thread"]
|
||||
tid = t["thread_id"]
|
||||
preview = t.get("preview", "") or ""
|
||||
msgs = t.get("message_count", 0)
|
||||
model = t.get("model", "") or ""
|
||||
when = _format_relative_time(t.get("updated_at"))
|
||||
indent = " " if item.get("indented") else " "
|
||||
marker = " *" if tid == current_thread else ""
|
||||
parts = [f"{indent}{tid}{marker}"]
|
||||
if preview:
|
||||
parts.append(preview[:40] + "…" if len(preview) > 40 else preview)
|
||||
parts.append(f"({msgs} msgs)")
|
||||
if model:
|
||||
parts.append(model)
|
||||
if when:
|
||||
parts.append(when)
|
||||
label = " ".join(parts)
|
||||
choices.append(questionary.Choice(title=label, value=tid))
|
||||
|
||||
prompt = questionary.select(title, choices=choices, style=PICKER_STYLE)
|
||||
# Limit visible list to 10 rows with scrolling. Touches
|
||||
# questionary/prompt-toolkit private internals so guard against
|
||||
# library-shape changes — picker stays functional at default
|
||||
# height even if the cap fails.
|
||||
try:
|
||||
for window in prompt.application.layout.find_all_windows():
|
||||
if isinstance(window.content, InquirerControl):
|
||||
window.height = Dimension(max=10)
|
||||
break
|
||||
except Exception:
|
||||
pass
|
||||
# ``ask_async`` (questionary >= 2.0.1) avoids blocking the
|
||||
# asyncio event loop while the user interacts with the picker.
|
||||
return await prompt.ask_async()
|
||||
|
||||
# ── Lifecycle callbacks ───────────────────────────────
|
||||
|
||||
def clear_chat(self) -> None:
|
||||
if self._on_clear_chat is not None:
|
||||
self._on_clear_chat()
|
||||
else:
|
||||
self.console.clear()
|
||||
|
||||
def request_quit(self) -> None:
|
||||
if self._on_request_quit is not None:
|
||||
self._on_request_quit()
|
||||
|
||||
def force_quit(self) -> None:
|
||||
if self._on_force_quit is not None:
|
||||
self._on_force_quit()
|
||||
|
||||
async def start_new_session(self) -> None:
|
||||
if self._on_start_new_session is not None:
|
||||
await self._on_start_new_session()
|
||||
|
||||
async def handle_session_resume(
|
||||
self, thread_id: str, workspace_dir: str | None = None
|
||||
) -> None:
|
||||
if self._on_handle_session_resume is not None:
|
||||
await self._on_handle_session_resume(thread_id, workspace_dir)
|
||||
|
||||
# /compact indicator pair — duck-typed by ``CompactCommand`` via
|
||||
# ``getattr``, not declared on the ``CommandUI`` Protocol.
|
||||
def start_compacting_indicator(self) -> None:
|
||||
# Idempotent: close any lingering context before starting a new
|
||||
# one so a double-call (e.g. two overlapping /compact attempts
|
||||
# via the message queue) can't leak a Rich Live handle.
|
||||
if self._compact_status_ctx is not None:
|
||||
try:
|
||||
self._compact_status_ctx.__exit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
self._compact_status_ctx = None
|
||||
status = self.console.status("[cyan]Compacting conversation...[/cyan]")
|
||||
status.__enter__()
|
||||
self._compact_status_ctx = status
|
||||
|
||||
def stop_compacting_indicator(self) -> None:
|
||||
ctx = self._compact_status_ctx
|
||||
self._compact_status_ctx = None
|
||||
if ctx is not None:
|
||||
try:
|
||||
ctx.__exit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def update_status_after_compact(self, input_tokens: int) -> None:
|
||||
if self._on_status_after_compact is not None:
|
||||
self._on_status_after_compact(input_tokens)
|
||||
|
||||
# ── Skill / MCP browse (delegated to worker threads) ──
|
||||
|
||||
async def wait_for_skill_browse(
|
||||
self, index: list[dict], installed_names: set[str], pre_filter_tag: str
|
||||
) -> list[str] | None:
|
||||
"""Delegate to the extracted questionary picker on a worker
|
||||
thread — questionary blocks the event loop so the call must
|
||||
not happen on the main asyncio thread."""
|
||||
import asyncio
|
||||
|
||||
from .skills_cmd import _pick_skills_interactive
|
||||
|
||||
return await asyncio.to_thread(
|
||||
_pick_skills_interactive, index, installed_names, pre_filter_tag
|
||||
)
|
||||
|
||||
async def wait_for_mcp_browse(
|
||||
self, servers: list, installed_names: set[str], pre_filter_tag: str
|
||||
) -> list | None:
|
||||
"""Delegate to the MCP browse picker on a worker thread."""
|
||||
import asyncio
|
||||
|
||||
from .mcp_install_cmd import _browse_and_select
|
||||
|
||||
return await asyncio.to_thread(
|
||||
_browse_and_select, servers, installed_names, pre_filter_tag
|
||||
)
|
||||
@@ -1,171 +1,30 @@
|
||||
"""Slash commands for skill management: /skills, /install-skill, /uninstall-skill, /evoskills."""
|
||||
"""Shared UI helpers for skill-management commands (picker used by /evoskills)."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from ..stream.display import console
|
||||
from .agent import _shorten_path
|
||||
from ..stream.console import console
|
||||
|
||||
|
||||
def _cmd_list_skills() -> None:
|
||||
"""List all available skills (workspace, global, and built-in)."""
|
||||
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
|
||||
from ..tools.skills_manager import list_skills
|
||||
def _pick_skills_interactive(
|
||||
index: list[dict],
|
||||
installed_names: set[str],
|
||||
pre_filter_tag: str,
|
||||
) -> list[str] | None:
|
||||
"""Interactive questionary picker for EvoSkills browse.
|
||||
|
||||
skills = list_skills(include_system=True)
|
||||
Two-phase picker:
|
||||
1. tag filter — ``questionary.select`` (skipped if ``pre_filter_tag``)
|
||||
2. multi-select — ``questionary.checkbox`` with installed items disabled
|
||||
|
||||
if not skills:
|
||||
console.print("[dim]No skills available.[/dim]")
|
||||
console.print("[dim]Install with:[/dim] /install-skill <path-or-url>")
|
||||
console.print(
|
||||
f"[dim]Global skills:[/dim] [cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
workspace_skills = [s for s in skills if s.source == "workspace"]
|
||||
global_skills = [s for s in skills if s.source == "global"]
|
||||
builtin_skills = [s for s in skills if s.source == "builtin"]
|
||||
|
||||
sections = [
|
||||
("Workspace Skills", workspace_skills, "green"),
|
||||
("Global Skills", global_skills, "cyan"),
|
||||
("Built-in Skills", builtin_skills, "blue"),
|
||||
]
|
||||
|
||||
printed = False
|
||||
for title, group, color in sections:
|
||||
if not group:
|
||||
continue
|
||||
if printed:
|
||||
console.print()
|
||||
console.print(f"[bold]{title}[/bold] ({len(group)}):")
|
||||
for skill in group:
|
||||
tags_str = f" [dim]({', '.join(skill.tags)})[/dim]" if skill.tags else ""
|
||||
console.print(
|
||||
f" [{color}]{skill.name}[/{color}] - {skill.description}{tags_str}"
|
||||
)
|
||||
printed = True
|
||||
|
||||
console.print(
|
||||
f"\n[dim]Global skills:[/dim] [cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan]"
|
||||
)
|
||||
console.print(
|
||||
f"[dim]Workspace skills:[/dim] [green]{_shorten_path(str(USER_SKILLS_DIR))}[/green]"
|
||||
)
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_install_skill(args: str) -> None:
|
||||
"""Install a skill from local path or GitHub URL.
|
||||
|
||||
By default, installs to the global skills directory (~/.config/ai4scientist/skills/).
|
||||
Append --local to install to the current workspace instead.
|
||||
|
||||
Usage: /install-skill <path-or-url> [--local]
|
||||
"""
|
||||
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
|
||||
from ..tools.skills_manager import install_skill
|
||||
|
||||
# Parse --local flag out of the args string
|
||||
local = "--local" in args.split()
|
||||
source = args.replace("--local", "").strip()
|
||||
|
||||
if not source:
|
||||
console.print("[red]Usage:[/red] /install-skill <path-or-url> [--local]")
|
||||
console.print("[dim]Examples:[/dim]")
|
||||
console.print(" /install-skill ./my-skill")
|
||||
console.print(
|
||||
" /install-skill https://github.com/user/repo/tree/main/skill-name"
|
||||
)
|
||||
console.print(" /install-skill user/repo@skill-name")
|
||||
console.print(
|
||||
" /install-skill ./my-skill --local [dim](workspace only)[/dim]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
dest_label = (
|
||||
f"[cyan]{_shorten_path(str(USER_SKILLS_DIR))}[/cyan] [dim](workspace)[/dim]"
|
||||
if local
|
||||
else f"[cyan]{_shorten_path(str(GLOBAL_SKILLS_DIR))}[/cyan] [dim](global)[/dim]"
|
||||
)
|
||||
console.print(f"[dim]Installing skill from:[/dim] {source}")
|
||||
console.print(f"[dim]Destination:[/dim] {dest_label}")
|
||||
|
||||
result = install_skill(source, global_install=not local)
|
||||
|
||||
if result.get("batch"):
|
||||
# Batch install — multiple skills
|
||||
for item in result.get("installed", []):
|
||||
console.print(f"[green]Installed:[/green] {item['name']}")
|
||||
console.print(
|
||||
f" [dim]Description:[/dim] {item.get('description', '(none)')}"
|
||||
)
|
||||
console.print(
|
||||
f" [dim]Path:[/dim] [cyan]{_shorten_path(item['path'])}[/cyan]"
|
||||
)
|
||||
for item in result.get("failed", []):
|
||||
console.print(f"[red]Failed:[/red] {item['name']} — {item['error']}")
|
||||
installed_count = len(result.get("installed", []))
|
||||
if installed_count:
|
||||
console.print(f"\n[green]{installed_count} skill(s) installed.[/green]")
|
||||
console.print("[dim]Reload with /new to apply.[/dim]")
|
||||
elif result["success"]:
|
||||
console.print(f"[green]Installed:[/green] {result['name']}")
|
||||
console.print(f"[dim]Description:[/dim] {result.get('description', '(none)')}")
|
||||
console.print(f"[dim]Path:[/dim] [cyan]{_shorten_path(result['path'])}[/cyan]")
|
||||
console.print()
|
||||
console.print("[dim]Reload with /new to apply.[/dim]")
|
||||
else:
|
||||
console.print(f"[red]Failed:[/red] {result['error']}")
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_uninstall_skill(name: str) -> None:
|
||||
"""Uninstall a user-installed skill."""
|
||||
from ..tools.skills_manager import uninstall_skill
|
||||
|
||||
if not name:
|
||||
console.print("[red]Usage:[/red] /uninstall-skill <skill-name>")
|
||||
console.print("[dim]Use /skills to see installed skills.[/dim]")
|
||||
console.print()
|
||||
return
|
||||
|
||||
result = uninstall_skill(name)
|
||||
|
||||
if result["success"]:
|
||||
console.print(f"[green]Uninstalled:[/green] {name}")
|
||||
console.print("[dim]Reload with /new to apply.[/dim]")
|
||||
else:
|
||||
console.print(f"[red]Failed:[/red] {result['error']}")
|
||||
console.print()
|
||||
|
||||
|
||||
def _cmd_install_skills(args: str = "") -> None:
|
||||
"""Browse and install skills from the EvoSkills repository.
|
||||
|
||||
Args:
|
||||
args: Optional tag name to pre-filter (e.g. "core").
|
||||
Returns:
|
||||
list of ``install_source`` strings selected by the user,
|
||||
``None`` if the user cancelled at either phase, or
|
||||
``[]`` if nothing was selectable / all-installed in the filter.
|
||||
"""
|
||||
from collections import Counter
|
||||
|
||||
import questionary
|
||||
from prompt_toolkit.styles import Style as PtStyle
|
||||
from questionary import Choice
|
||||
|
||||
from ..paths import GLOBAL_SKILLS_DIR, USER_SKILLS_DIR
|
||||
from ..tools.skills_manager import fetch_remote_skill_index, install_skill
|
||||
|
||||
_PICKER_STYLE = PtStyle.from_dict(
|
||||
{
|
||||
"questionmark": "#888888",
|
||||
"question": "",
|
||||
"pointer": "bold",
|
||||
"highlighted": "bold",
|
||||
"text": "#888888",
|
||||
"answer": "bold",
|
||||
}
|
||||
)
|
||||
from .widgets.thread_selector import PICKER_STYLE
|
||||
|
||||
# Installed-item indicator style for disabled checkbox choices.
|
||||
_INSTALLED_INDICATOR = ("fg:#4caf50", "✓ ")
|
||||
@@ -190,46 +49,24 @@ def _cmd_install_skills(args: str = "") -> None:
|
||||
return questionary.checkbox(
|
||||
message,
|
||||
choices=choices,
|
||||
style=_PICKER_STYLE,
|
||||
style=PICKER_STYLE,
|
||||
qmark="❯",
|
||||
**kwargs,
|
||||
).ask()
|
||||
finally:
|
||||
InquirerControl._get_choice_tokens = original
|
||||
|
||||
# Step 1: Fetch remote index
|
||||
console.print("[dim]Fetching skill index...[/dim]")
|
||||
try:
|
||||
index = fetch_remote_skill_index()
|
||||
except Exception as e:
|
||||
console.print(f"[red]Failed to fetch skill index: {e}[/red]")
|
||||
console.print(
|
||||
"[dim]Try installing directly: /install-skill EvoScientist/EvoSkills@skills[/dim]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
pre_filter_tag = (pre_filter_tag or "").strip().lower()
|
||||
|
||||
if not index:
|
||||
console.print("[yellow]No skills found in the repository.[/yellow]")
|
||||
console.print()
|
||||
return
|
||||
|
||||
# Detect already-installed skills (both global and workspace tiers)
|
||||
installed_names: set[str] = set()
|
||||
for skills_dir in (Path(GLOBAL_SKILLS_DIR), Path(USER_SKILLS_DIR)):
|
||||
if skills_dir.exists():
|
||||
installed_names.update(e.name for e in skills_dir.iterdir() if e.is_dir())
|
||||
|
||||
pre_filter_tag = args.strip().lower() if args else ""
|
||||
|
||||
# Step 2: Tag filter (skip if pre-filtered via args)
|
||||
# Phase 1: tag filter (skip if pre-filtered via args)
|
||||
if pre_filter_tag:
|
||||
filtered = [
|
||||
s for s in index if pre_filter_tag in [t.lower() for t in s.get("tags", [])]
|
||||
]
|
||||
if not filtered:
|
||||
console.print(f"[yellow]No skills found with tag: {args.strip()}[/yellow]")
|
||||
# Show available tags
|
||||
console.print(
|
||||
f"[yellow]No skills found with tag: {pre_filter_tag}[/yellow]"
|
||||
)
|
||||
tag_counter: Counter[str] = Counter()
|
||||
for s in index:
|
||||
for t in s.get("tags", []):
|
||||
@@ -238,10 +75,8 @@ def _cmd_install_skills(args: str = "") -> None:
|
||||
sorted_tags = sorted(tag_counter.items(), key=lambda x: (-x[1], x[0]))
|
||||
tags_str = ", ".join(f"{tag} ({count})" for tag, count in sorted_tags)
|
||||
console.print(f"[dim]Available tags: {tags_str}[/dim]")
|
||||
console.print()
|
||||
return
|
||||
return []
|
||||
else:
|
||||
# Build tag choices for interactive picker
|
||||
tag_counter = Counter()
|
||||
for s in index:
|
||||
for t in s.get("tags", []):
|
||||
@@ -255,13 +90,12 @@ def _cmd_install_skills(args: str = "") -> None:
|
||||
selected_tag = questionary.select(
|
||||
"Filter by tag:",
|
||||
choices=tag_choices,
|
||||
style=_PICKER_STYLE,
|
||||
style=PICKER_STYLE,
|
||||
qmark="❯",
|
||||
).ask()
|
||||
|
||||
if selected_tag is None:
|
||||
console.print()
|
||||
return
|
||||
return None
|
||||
|
||||
if selected_tag == "__all__":
|
||||
filtered = index
|
||||
@@ -272,14 +106,12 @@ def _cmd_install_skills(args: str = "") -> None:
|
||||
if selected_tag in [t.lower() for t in s.get("tags", [])]
|
||||
]
|
||||
|
||||
# Step 3: Skill selection checkbox
|
||||
all_installed = all(s["name"] in installed_names for s in filtered)
|
||||
if all_installed:
|
||||
# Phase 2: skill selection checkbox
|
||||
if all(s["name"] in installed_names for s in filtered):
|
||||
console.print(
|
||||
"[green]All skills in this category are already installed.[/green]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
return []
|
||||
|
||||
choices = []
|
||||
for s in filtered:
|
||||
@@ -305,31 +137,5 @@ def _cmd_install_skills(args: str = "") -> None:
|
||||
selected = _checkbox_ask(choices, "Select skills to install:")
|
||||
|
||||
if selected is None:
|
||||
console.print()
|
||||
return
|
||||
|
||||
if not selected:
|
||||
console.print("[dim]No skills selected.[/dim]")
|
||||
console.print()
|
||||
return
|
||||
|
||||
# Step 4: Install selected skills (default: global)
|
||||
installed_count = 0
|
||||
for source in selected:
|
||||
result = install_skill(source, global_install=True)
|
||||
if result.get("batch"):
|
||||
for item in result.get("installed", []):
|
||||
console.print(f"[green]Installed:[/green] {item['name']}")
|
||||
installed_count += 1
|
||||
for item in result.get("failed", []):
|
||||
console.print(f"[red]Failed:[/red] {item['name']} — {item['error']}")
|
||||
elif result.get("success"):
|
||||
console.print(f"[green]Installed:[/green] {result['name']}")
|
||||
installed_count += 1
|
||||
else:
|
||||
console.print(f"[red]Failed:[/red] {result.get('error', 'unknown')}")
|
||||
|
||||
if installed_count:
|
||||
console.print(f"\n[green]{installed_count} skill(s) installed.[/green]")
|
||||
console.print("[dim]Reload with /new to apply.[/dim]")
|
||||
console.print()
|
||||
return None
|
||||
return list(selected)
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langchain_core.messages.utils import count_tokens_approximately
|
||||
@@ -13,7 +13,15 @@ from ..llm.context_window import (
|
||||
DEFAULT_CONTEXT_WINDOW_FALLBACK,
|
||||
resolve_context_window,
|
||||
)
|
||||
from ..sessions import get_thread_messages
|
||||
from ..memory.worker_activity import (
|
||||
MemoryWorkerStatusSnapshot,
|
||||
ObservationLinkerStatusSnapshot,
|
||||
memory_worker_status,
|
||||
observation_linker_status,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway
|
||||
|
||||
_FALLBACK_CONTEXT_WINDOW = DEFAULT_CONTEXT_WINDOW_FALLBACK
|
||||
STATUS_BAR_BG = "#171a20"
|
||||
@@ -26,6 +34,11 @@ STATUS_BAD = "#d08c61"
|
||||
STATUS_CRITICAL = "#d86f6f"
|
||||
STATUS_HINT_IDLE = "#8b9bb0"
|
||||
STATUS_HINT_BUSY = "#f0c36a"
|
||||
STATUS_HINT_WRITING = "#7eb8e0"
|
||||
|
||||
# Braille spinner frames used by the CLI bottom toolbar and TUI status bar
|
||||
# to animate the "Loading MCP tools" indicator.
|
||||
SPINNER_FRAMES = "\u280b\u2819\u2839\u2838\u283c\u2834\u2826\u2827\u2807\u280f"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
@@ -86,15 +99,15 @@ def format_token_count_compact(value: int) -> str:
|
||||
"""Format large token counts into a compact human-readable form."""
|
||||
abs_value = abs(int(value))
|
||||
if abs_value >= 1_000_000:
|
||||
num = value / 1_000_000
|
||||
num = float(value) / 1_000_000
|
||||
suffix = "M"
|
||||
elif abs_value >= 1_000:
|
||||
num = value / 1_000
|
||||
num = float(value) / 1_000
|
||||
suffix = "K"
|
||||
else:
|
||||
return str(value)
|
||||
|
||||
if num.is_integer():
|
||||
if num == int(num):
|
||||
return f"{int(num)}{suffix}"
|
||||
return f"{num:.1f}{suffix}"
|
||||
|
||||
@@ -155,15 +168,10 @@ def trim_status_text(text: str, max_width: int) -> str:
|
||||
if max_width <= ellipsis_width:
|
||||
return ellipsis[:max_width]
|
||||
|
||||
try:
|
||||
from prompt_toolkit.utils import get_cwidth
|
||||
except Exception:
|
||||
get_cwidth = None
|
||||
|
||||
out: list[str] = []
|
||||
width = 0
|
||||
for ch in text:
|
||||
ch_width = get_cwidth(ch) if get_cwidth else len(ch)
|
||||
ch_width = _display_width(ch)
|
||||
if width + ch_width + ellipsis_width > max_width:
|
||||
break
|
||||
out.append(ch)
|
||||
@@ -171,13 +179,94 @@ def trim_status_text(text: str, max_width: int) -> str:
|
||||
return "".join(out).rstrip() + ellipsis
|
||||
|
||||
|
||||
def get_memory_worker_status() -> MemoryWorkerStatusSnapshot | None:
|
||||
"""Read completed EvoMemory save counts without making rendering fail."""
|
||||
try:
|
||||
return memory_worker_status()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def get_observation_linker_status() -> ObservationLinkerStatusSnapshot | None:
|
||||
"""Read active observation-linker status without making rendering fail."""
|
||||
try:
|
||||
return observation_linker_status()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _plural(count: int, singular: str, plural: str | None = None) -> str:
|
||||
word = singular if count == 1 else (plural or f"{singular}s")
|
||||
return f"{count} {word}"
|
||||
|
||||
|
||||
def _memory_activity_label(
|
||||
*,
|
||||
worker_status: MemoryWorkerStatusSnapshot | None,
|
||||
linker_status: ObservationLinkerStatusSnapshot | None,
|
||||
) -> str:
|
||||
parts: list[str] = []
|
||||
if worker_status is not None and worker_status.is_running:
|
||||
parts.append("🧠")
|
||||
if linker_status is not None and linker_status.is_running:
|
||||
parts.append("🔗")
|
||||
|
||||
saved: list[str] = []
|
||||
if worker_status is not None:
|
||||
if worker_status.profile_updates:
|
||||
saved.append(_plural(worker_status.profile_updates, "profile edit"))
|
||||
if worker_status.observations_recorded:
|
||||
saved.append(_plural(worker_status.observations_recorded, "observation"))
|
||||
if saved:
|
||||
parts.append(f"Saved {', '.join(saved)}")
|
||||
|
||||
if linker_status is not None and linker_status.relations_linked:
|
||||
parts.append(
|
||||
f"Created {_plural(linker_status.relations_linked, 'memory link')}"
|
||||
)
|
||||
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _append_memory_indicator(
|
||||
frags: list[tuple[str, str]],
|
||||
*,
|
||||
worker_status: MemoryWorkerStatusSnapshot | None,
|
||||
linker_status: ObservationLinkerStatusSnapshot | None,
|
||||
width: int,
|
||||
) -> None:
|
||||
if worker_status is None and linker_status is None:
|
||||
return
|
||||
|
||||
label = _memory_activity_label(
|
||||
worker_status=worker_status,
|
||||
linker_status=linker_status,
|
||||
)
|
||||
if not label:
|
||||
return
|
||||
|
||||
tail: list[tuple[str, str]] = []
|
||||
if frags and frags[-1] == ("class:status-bar", " "):
|
||||
tail.append(frags.pop())
|
||||
|
||||
separator = " │ " if width >= 76 else " · "
|
||||
frags.extend(
|
||||
[
|
||||
("class:status-bar-dim", separator),
|
||||
("class:status-bar-warn", label),
|
||||
]
|
||||
)
|
||||
frags.extend(tail)
|
||||
|
||||
|
||||
def build_status_fragments(
|
||||
snapshot: SessionStatusSnapshot,
|
||||
started_at: datetime,
|
||||
width: int,
|
||||
) -> list[tuple[str, str]]:
|
||||
"""Build prompt_toolkit formatted-text fragments for the status bar."""
|
||||
duration_label = format_duration_compact(started_at)
|
||||
now = datetime.now()
|
||||
duration_label = format_duration_compact(started_at, now=now)
|
||||
percent = snapshot.context_percent
|
||||
percent_label = f"{percent}%"
|
||||
if width < 52:
|
||||
@@ -215,6 +304,13 @@ def build_status_fragments(
|
||||
("class:status-bar", " "),
|
||||
]
|
||||
|
||||
_append_memory_indicator(
|
||||
frags,
|
||||
worker_status=get_memory_worker_status(),
|
||||
linker_status=get_observation_linker_status(),
|
||||
width=width,
|
||||
)
|
||||
|
||||
total_width = sum(_display_width(text) for _, text in frags)
|
||||
if total_width > width:
|
||||
plain_text = "".join(text for _, text in frags)
|
||||
@@ -345,11 +441,12 @@ async def build_session_status_snapshot(
|
||||
model_name: str | None = None,
|
||||
model_obj: Any | None = None,
|
||||
pending_user_text: str | None = None,
|
||||
graph_gateway: GraphGateway,
|
||||
) -> SessionStatusSnapshot:
|
||||
"""Count current thread context and return a display snapshot."""
|
||||
resolved_name = _resolve_model_name(model_name, model_obj)
|
||||
window = _resolve_context_window(model_obj)
|
||||
messages = list(await get_thread_messages(thread_id))
|
||||
messages = list(await graph_gateway.get_thread_messages(thread_id))
|
||||
|
||||
pending = (pending_user_text or "").strip()
|
||||
if pending:
|
||||
|
||||
@@ -6,6 +6,7 @@ from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from ..gateway import GraphGateway
|
||||
from ..stream.display import _run_streaming
|
||||
|
||||
|
||||
@@ -30,6 +31,8 @@ class StreamingTUIBackend(Protocol):
|
||||
metadata: dict | None = None,
|
||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||
cancel_scope: str | None = None,
|
||||
gateway: GraphGateway,
|
||||
) -> str:
|
||||
"""Run streaming and return final response text."""
|
||||
|
||||
@@ -56,6 +59,8 @@ class RichStreamingBackend:
|
||||
metadata: dict | None = None,
|
||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||
cancel_scope: str | None = None,
|
||||
gateway: GraphGateway,
|
||||
) -> str:
|
||||
return _run_streaming(
|
||||
agent=agent,
|
||||
@@ -71,4 +76,6 @@ class RichStreamingBackend:
|
||||
metadata=metadata,
|
||||
hitl_prompt_fn=hitl_prompt_fn,
|
||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||
cancel_scope=cancel_scope,
|
||||
gateway=gateway,
|
||||
)
|
||||
|
||||
@@ -5,11 +5,16 @@ from __future__ import annotations
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from ..stream.display import console
|
||||
from ..gateway import GraphGateway
|
||||
from ..stream.console import console
|
||||
from .tui_backends import RichStreamingBackend, StreamingTUIBackend
|
||||
|
||||
DEFAULT_UI_BACKEND = "cli"
|
||||
SUPPORTED_UI_BACKENDS = ("cli", "tui")
|
||||
# "webui" launches the browser front-end instead of an in-terminal UI; it is
|
||||
# intercepted earlier (cli/commands.py:_main_callback) and never reaches the
|
||||
# streaming backends, but is listed here so normalize/resolve preserve it
|
||||
# rather than falling back to "cli".
|
||||
SUPPORTED_UI_BACKENDS = ("cli", "tui", "webui")
|
||||
_LEGACY_BACKEND_MAP = {"textual": "tui", "rich": "cli"}
|
||||
|
||||
|
||||
@@ -74,6 +79,8 @@ def run_streaming(
|
||||
metadata: dict | None = None,
|
||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||
cancel_scope: str | None = None,
|
||||
gateway: GraphGateway,
|
||||
) -> str:
|
||||
"""Run streaming with the selected backend."""
|
||||
backend = get_backend(ui_backend, warn_fallback=True)
|
||||
@@ -92,6 +99,8 @@ def run_streaming(
|
||||
metadata=metadata,
|
||||
hitl_prompt_fn=hitl_prompt_fn,
|
||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||
cancel_scope=cancel_scope,
|
||||
gateway=gateway,
|
||||
)
|
||||
except RuntimeError:
|
||||
requested = normalize_ui_backend(ui_backend)
|
||||
@@ -113,5 +122,7 @@ def run_streaming(
|
||||
metadata=metadata,
|
||||
hitl_prompt_fn=hitl_prompt_fn,
|
||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||
cancel_scope=cancel_scope,
|
||||
gateway=gateway,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -6,6 +6,7 @@ from .assistant_message import AssistantMessage
|
||||
from .compact_summary_widget import CompactSummaryWidget
|
||||
from .compacting_widget import CompactingWidget
|
||||
from .loading_widget import LoadingWidget
|
||||
from .mcp_loader_widget import MCPLoaderWidget
|
||||
from .subagent_widget import SubAgentWidget
|
||||
from .summarization_widget import SummarizationWidget
|
||||
from .system_message import SystemMessage
|
||||
@@ -23,6 +24,7 @@ __all__ = [
|
||||
"CompactSummaryWidget",
|
||||
"CompactingWidget",
|
||||
"LoadingWidget",
|
||||
"MCPLoaderWidget",
|
||||
"SubAgentWidget",
|
||||
"SummarizationWidget",
|
||||
"SystemMessage",
|
||||
|
||||
@@ -105,11 +105,7 @@ class ApprovalWidget(Widget):
|
||||
self._option_widgets = []
|
||||
count = len(self._action_requests)
|
||||
if count == 1:
|
||||
name = (
|
||||
self._action_requests[0].get("name", "")
|
||||
if isinstance(self._action_requests[0], dict)
|
||||
else getattr(self._action_requests[0], "name", "")
|
||||
)
|
||||
name = self._action_requests[0].get("name", "")
|
||||
title = f">>> {name} Requires Approval <<<"
|
||||
else:
|
||||
title = f">>> {count} Tool Calls Require Approval <<<"
|
||||
@@ -117,16 +113,8 @@ class ApprovalWidget(Widget):
|
||||
|
||||
# Show each action request as a compact line
|
||||
for req in self._action_requests:
|
||||
name = (
|
||||
req.get("name", "")
|
||||
if isinstance(req, dict)
|
||||
else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {})
|
||||
if isinstance(req, dict)
|
||||
else getattr(req, "args", {})
|
||||
)
|
||||
name = req.get("name", "")
|
||||
args = req.get("args", {})
|
||||
if isinstance(args, dict):
|
||||
command = args.get("command", args.get("path", ""))
|
||||
else:
|
||||
|
||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
from textual.containers import Vertical
|
||||
from textual.widgets import Markdown
|
||||
|
||||
from ...stream.display import _fix_markdown_heading_spacing
|
||||
from .timestamp_mixin import TimestampClickMixin
|
||||
|
||||
|
||||
@@ -39,8 +40,11 @@ class AssistantMessage(TimestampClickMixin, Vertical):
|
||||
yield Markdown("")
|
||||
|
||||
def on_mount(self) -> None:
|
||||
"""Render ``initial_content`` once the widget enters the DOM."""
|
||||
if self._content:
|
||||
self.query_one(Markdown).update(self._content)
|
||||
self.query_one(Markdown).update(
|
||||
_fix_markdown_heading_spacing(self._content)
|
||||
)
|
||||
|
||||
async def append_content(self, text: str) -> None:
|
||||
"""Append text and schedule a debounced Markdown re-render."""
|
||||
@@ -50,12 +54,14 @@ class AssistantMessage(TimestampClickMixin, Vertical):
|
||||
self.set_timer(0.1, self._flush_markdown)
|
||||
|
||||
def _flush_markdown(self) -> None:
|
||||
"""Flush accumulated content to the Markdown widget."""
|
||||
"""Flush accumulated content to the Markdown widget on a display copy."""
|
||||
self._flush_pending = False
|
||||
self.query_one(Markdown).update(self._content)
|
||||
self.query_one(Markdown).update(_fix_markdown_heading_spacing(self._content))
|
||||
|
||||
async def stop_stream(self) -> None:
|
||||
"""Finalize the stream — ensure final content is rendered."""
|
||||
self._flush_pending = False
|
||||
if self._content:
|
||||
self.query_one(Markdown).update(self._content)
|
||||
self.query_one(Markdown).update(
|
||||
_fix_markdown_heading_spacing(self._content)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,191 @@
|
||||
"""Live-updating widget that shows per-server MCP load progress.
|
||||
|
||||
Mounted above the chat input while MCP tools are being fetched in the
|
||||
background. Re-renders on a 100 ms tick so the spinner animates and the
|
||||
per-server states transition smoothly from pending → ok/error.
|
||||
|
||||
When the load finishes:
|
||||
- All-success runs auto-dismiss after a short grace period so the chat
|
||||
area isn't permanently crowded.
|
||||
- Failures stick around longer so the user has time to read the error
|
||||
detail, then auto-dismiss — otherwise the widget pins itself above
|
||||
the input forever.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from rich.text import Text
|
||||
from textual.widgets import Static
|
||||
|
||||
from ..status_bar import SPINNER_FRAMES
|
||||
|
||||
_DIM = "#7c8594"
|
||||
_STRONG = "#e5e7eb"
|
||||
_GOOD = "#5fcf8b"
|
||||
_WARN = "#d7b45a"
|
||||
_BAD = "#d86f6f"
|
||||
|
||||
# How long to wait after an all-success load before auto-dismissing.
|
||||
_AUTO_DISMISS_SECONDS = 2.5
|
||||
# Longer grace on failure so the user has time to read error detail.
|
||||
_AUTO_DISMISS_ON_ERROR_SECONDS = 12.0
|
||||
|
||||
|
||||
class MCPLoaderWidget(Static):
|
||||
"""Shows a header line + one line per MCP server with its live status."""
|
||||
|
||||
DEFAULT_CSS = """
|
||||
MCPLoaderWidget {
|
||||
height: auto;
|
||||
padding: 0 1;
|
||||
margin: 0 0 1 0;
|
||||
}
|
||||
"""
|
||||
|
||||
TICK_SECONDS = 0.1
|
||||
|
||||
def __init__(self, servers: list[str]) -> None:
|
||||
# server_name -> (state, detail); state ∈ {"pending","ok","error"}.
|
||||
self._progress: dict[str, tuple[str, str]] = dict.fromkeys(
|
||||
servers, ("pending", "")
|
||||
)
|
||||
self._frame = 0
|
||||
self._tick_handle = None
|
||||
self._finished = False
|
||||
self._dismissed = False
|
||||
self._auto_dismiss_at: float | None = None
|
||||
# Seed with real content so Textual can measure us before the
|
||||
# first tick; ``self.update()`` during ``__init__`` is unsafe
|
||||
# (widget isn't attached yet), but we can pass the renderable
|
||||
# straight into ``Static.__init__``.
|
||||
super().__init__(self._build_renderable())
|
||||
|
||||
def on_mount(self) -> None:
|
||||
self._tick_handle = self.set_interval(self.TICK_SECONDS, self._tick)
|
||||
|
||||
def on_unmount(self) -> None:
|
||||
if self._tick_handle is not None:
|
||||
self._tick_handle.stop()
|
||||
self._tick_handle = None
|
||||
|
||||
# ── Public API ───────────────────────────────────────────────────
|
||||
|
||||
@property
|
||||
def dismissed(self) -> bool:
|
||||
"""Whether the widget has already removed itself from the DOM."""
|
||||
return self._dismissed
|
||||
|
||||
def update_server(self, name: str, state: str, detail: str = "") -> None:
|
||||
"""Record a progress event for one server and re-render."""
|
||||
if self._dismissed or state not in ("pending", "ok", "error"):
|
||||
return
|
||||
# First-time-seen servers (e.g., ones missing from the initial
|
||||
# prime set because the config file changed mid-load) just get
|
||||
# appended — order stays stable for already-known entries.
|
||||
self._progress[name] = (state, detail)
|
||||
self._refresh_content()
|
||||
|
||||
def mark_finished(self) -> None:
|
||||
"""Call once the background load task resolves (success or error).
|
||||
|
||||
If nothing ever progressed past ``pending``, the load was served
|
||||
from cache (no events emitted) — drop the widget immediately
|
||||
instead of flashing a misleading "0/N loaded" header.
|
||||
|
||||
Otherwise schedule an auto-dismiss: short on full success so the
|
||||
chat area isn't cluttered, longer on failure so the user has
|
||||
time to read the error detail before it goes away.
|
||||
"""
|
||||
if self._finished:
|
||||
return
|
||||
self._finished = True
|
||||
progressed = any(state != "pending" for state, _ in self._progress.values())
|
||||
if not progressed:
|
||||
self._dismiss()
|
||||
return
|
||||
has_errors = any(state == "error" for state, _ in self._progress.values())
|
||||
delay = _AUTO_DISMISS_ON_ERROR_SECONDS if has_errors else _AUTO_DISMISS_SECONDS
|
||||
self._auto_dismiss_at = time.monotonic() + delay
|
||||
self._refresh_content()
|
||||
|
||||
# ── Internal ─────────────────────────────────────────────────────
|
||||
|
||||
def _tick(self) -> None:
|
||||
self._frame = (self._frame + 1) % len(SPINNER_FRAMES)
|
||||
if (
|
||||
self._auto_dismiss_at is not None
|
||||
and time.monotonic() >= self._auto_dismiss_at
|
||||
):
|
||||
self._auto_dismiss_at = None
|
||||
self._dismiss()
|
||||
return
|
||||
if not self._finished:
|
||||
self._refresh_content()
|
||||
|
||||
def _dismiss(self) -> None:
|
||||
"""Stop the tick timer and detach from the DOM.
|
||||
|
||||
Sets :attr:`dismissed` so the app can clear its widget reference
|
||||
and late progress events become no-ops.
|
||||
"""
|
||||
if self._dismissed:
|
||||
return
|
||||
self._dismissed = True
|
||||
if self._tick_handle is not None:
|
||||
self._tick_handle.stop()
|
||||
self._tick_handle = None
|
||||
# Fire-and-forget remove — nothing awaits us.
|
||||
self.remove()
|
||||
|
||||
def _build_renderable(self) -> Text:
|
||||
spinner = SPINNER_FRAMES[self._frame]
|
||||
pending = sum(1 for state, _ in self._progress.values() if state == "pending")
|
||||
total = len(self._progress)
|
||||
done = total - pending
|
||||
|
||||
header = Text()
|
||||
if self._finished:
|
||||
errors = sum(1 for state, _ in self._progress.values() if state == "error")
|
||||
if errors:
|
||||
header.append("✗ MCP ", style=f"{_BAD} bold")
|
||||
header.append(
|
||||
f"{done - errors}/{total} loaded, {errors} failed",
|
||||
style=_STRONG,
|
||||
)
|
||||
else:
|
||||
header.append("✓ MCP ", style=f"{_GOOD} bold")
|
||||
header.append(f"{done}/{total} servers loaded", style=_STRONG)
|
||||
else:
|
||||
header.append(f"{spinner} ", style=f"{_WARN} bold")
|
||||
header.append("Loading MCP tools ", style=_STRONG)
|
||||
header.append(f"{done}/{total}", style=_DIM)
|
||||
|
||||
lines: list[Text] = [header]
|
||||
for name, (state, detail) in self._progress.items():
|
||||
line = Text(" ")
|
||||
if state == "pending":
|
||||
line.append(f"{spinner} ", style=_WARN)
|
||||
line.append(name, style=_DIM)
|
||||
elif state == "ok":
|
||||
line.append("✓ ", style=_GOOD)
|
||||
line.append(name, style=_STRONG)
|
||||
if detail:
|
||||
line.append(f" {detail} tools", style=_DIM)
|
||||
else: # error
|
||||
line.append("✗ ", style=_BAD)
|
||||
line.append(name, style=_STRONG)
|
||||
if detail:
|
||||
summary = detail if len(detail) <= 80 else detail[:77] + "…"
|
||||
line.append(f" {summary}", style=_BAD)
|
||||
lines.append(line)
|
||||
|
||||
return Text("\n").join(lines)
|
||||
|
||||
def _refresh_content(self) -> None:
|
||||
# NB: don't name this ``_render`` — that shadows Textual's internal
|
||||
# ``Widget._render`` which must return a ``Visual``. Silently
|
||||
# breaking that contract triggers ``'NoneType' object has no
|
||||
# attribute 'get_height'`` during layout.
|
||||
self.update(self._build_renderable())
|
||||
@@ -22,7 +22,7 @@ class SubAgentWidget(Vertical):
|
||||
|
||||
┌ ▶ Cooking with research-agent — Search literature ─┐
|
||||
│ ✓ 8 completed │
|
||||
│ ● web_search query="LLM attention" │
|
||||
│ ● tavily_search query="LLM attention" │
|
||||
│ ✓ 3 results │
|
||||
└─────────────────────────────────────────────────────┘
|
||||
|
||||
@@ -208,10 +208,11 @@ class SubAgentWidget(Vertical):
|
||||
widget.set_success(content)
|
||||
else:
|
||||
widget.set_error(content)
|
||||
# Move from running to completed
|
||||
# Move from running to completed (dedup guards against repeat
|
||||
# deliveries of the same tool result inflating the collapse summary).
|
||||
if matched_key and matched_key in self._running_ids:
|
||||
self._running_ids.remove(matched_key)
|
||||
if matched_key:
|
||||
if matched_key and matched_key not in self._completed_ids:
|
||||
self._completed_ids.append(matched_key)
|
||||
self._update_visibility()
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from prompt_toolkit.styles import Style as PtStyle # type: ignore[import-untyped]
|
||||
from rich.text import Text
|
||||
from textual.binding import Binding, BindingType
|
||||
from textual.containers import Container
|
||||
@@ -30,6 +31,22 @@ if TYPE_CHECKING:
|
||||
from textual.app import ComposeResult
|
||||
|
||||
|
||||
# Style for questionary pickers used by the Rich CLI ``/resume`` and
|
||||
# ``/delete`` interactive selectors. Matches the slash-completion menu's
|
||||
# visual language: gray (#888888) for non-selected, bold for selected,
|
||||
# no background changes.
|
||||
PICKER_STYLE = PtStyle.from_dict(
|
||||
{
|
||||
"questionmark": "#888888",
|
||||
"question": "",
|
||||
"pointer": "bold",
|
||||
"highlighted": "bold",
|
||||
"text": "#888888",
|
||||
"answer": "bold",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Path helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -185,9 +202,9 @@ def build_row_text(
|
||||
indented: bool = False,
|
||||
) -> Text:
|
||||
"""Thread row. *indented* adds extra leading space for L2-grouped rows."""
|
||||
from ...sessions import _format_relative_time
|
||||
from ...sessions import _format_relative_time, short_thread_id
|
||||
|
||||
tid = thread["thread_id"]
|
||||
tid = short_thread_id(thread["thread_id"])
|
||||
preview = thread.get("preview", "") or ""
|
||||
msgs = thread.get("message_count", 0)
|
||||
model = thread.get("model", "") or ""
|
||||
|
||||
@@ -16,7 +16,12 @@ class UsageWidget(Static):
|
||||
}
|
||||
"""
|
||||
|
||||
def __init__(self, input_tokens: int, output_tokens: int) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
elapsed: str | None = None,
|
||||
) -> None:
|
||||
stats = Text(justify="right")
|
||||
stats.append("[", style="dim italic")
|
||||
stats.append("Usage: ", style="dim italic")
|
||||
@@ -24,5 +29,8 @@ class UsageWidget(Static):
|
||||
stats.append(" in · ", style="dim italic")
|
||||
stats.append(f"{output_tokens:,}", style="green italic")
|
||||
stats.append(" out", style="dim italic")
|
||||
if elapsed:
|
||||
stats.append(" · ", style="dim italic")
|
||||
stats.append(f"Elapsed: {elapsed}", style="dim italic")
|
||||
stats.append("]", style="dim italic")
|
||||
super().__init__(stats)
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Transient widget shown while a /resume restarts the langgraph dev subprocess.
|
||||
|
||||
Mirrors ``CompactingWidget`` — a timer-backed status line that ticks elapsed
|
||||
seconds so the user has live feedback during the up-to-60s langgraph dev
|
||||
workspace sync (subprocess stop + restart so deployed sub-agents see the
|
||||
resumed thread's workspace).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from .timed_status_widget import TimedStatusWidget
|
||||
|
||||
|
||||
class WorkspaceSyncWidget(TimedStatusWidget):
|
||||
"""Timer-backed status line for an in-progress workspace sync."""
|
||||
|
||||
DEFAULT_CSS = """
|
||||
WorkspaceSyncWidget {
|
||||
height: auto;
|
||||
color: #94a3b8;
|
||||
padding: 0 0;
|
||||
margin: 0 0 1 0;
|
||||
}
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def _refresh_display(self) -> None:
|
||||
self.update(
|
||||
f"Syncing async sub-agent server to resumed workspace... "
|
||||
f"({self.elapsed_seconds}s)"
|
||||
)
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Stop timer and remove from DOM."""
|
||||
self._stop_timer()
|
||||
if self.is_mounted:
|
||||
await self.remove()
|
||||
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from . import implementation
|
||||
from .base import Argument, Command, CommandContext, CommandUI
|
||||
from .base import Argument, Command, CommandContext, CommandUI, SubCommand
|
||||
from .channel_ui import ChannelCommandUI
|
||||
from .manager import CommandManager, manager
|
||||
|
||||
@@ -12,6 +12,7 @@ __all__ = [
|
||||
"CommandContext",
|
||||
"CommandManager",
|
||||
"CommandUI",
|
||||
"SubCommand",
|
||||
"implementation",
|
||||
"manager",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class CompletionKind(StrEnum):
|
||||
"""Discriminator for the kind of completion result."""
|
||||
|
||||
COMMANDS = "commands"
|
||||
SUBCOMMANDS = "subcommands"
|
||||
EMPTY = "empty"
|
||||
|
||||
|
||||
_CATEGORY_ORDER = ["Session", "Skills", "MCP", "Channels", "General"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CompletionCandidate:
|
||||
"""A single completion suggestion with its replacement range."""
|
||||
|
||||
text: str
|
||||
description: str
|
||||
replace_start: int
|
||||
replace_end: int
|
||||
category: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CompletionResult:
|
||||
"""The result of parsing a slash command input for completions."""
|
||||
|
||||
kind: CompletionKind
|
||||
candidates: list[CompletionCandidate]
|
||||
|
||||
|
||||
def compute_completions(text: str, cursor_pos: int) -> CompletionResult:
|
||||
"""Parse *text* up to *cursor_pos* and return completion candidates.
|
||||
|
||||
This is the shared engine used by both the Rich CLI
|
||||
(``SlashCommandCompleter``) and the TUI (``on_text_area_changed``).
|
||||
Both thin adapters only need to translate the returned candidates
|
||||
into their respective render/apply primitives.
|
||||
"""
|
||||
from .manager import manager as cmd_manager
|
||||
|
||||
before = text[:cursor_pos]
|
||||
|
||||
if not before.startswith("/"):
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
parts = before.split()
|
||||
if not parts:
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
cmd_name = parts[0].lower()
|
||||
has_trailing_space = before.endswith(" ")
|
||||
|
||||
# --- Top-level command completion ---
|
||||
if len(parts) == 1:
|
||||
prefix = before.lower().rstrip()
|
||||
|
||||
# Match commands by canonical name AND aliases
|
||||
by_cat: dict[str, list[tuple[str, str]]] = {}
|
||||
for cmd in cmd_manager.get_all_commands():
|
||||
all_names = [cmd.name.lower()] + [
|
||||
a.lower() if a.startswith("/") else f"/{a.lower()}" for a in cmd.alias
|
||||
]
|
||||
if any(n.startswith(prefix) for n in all_names):
|
||||
by_cat.setdefault(cmd.category, []).append((cmd.name, cmd.description))
|
||||
|
||||
# Whether the typed prefix resolves to an exact command/alias
|
||||
exact_cmd = cmd_manager.get_command(prefix)
|
||||
|
||||
if exact_cmd and not has_trailing_space:
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
if exact_cmd and has_trailing_space:
|
||||
completions = exact_cmd.get_completions([""])
|
||||
if completions:
|
||||
insert_pos = len(before)
|
||||
return CompletionResult(
|
||||
CompletionKind.SUBCOMMANDS,
|
||||
[
|
||||
CompletionCandidate(
|
||||
text=name,
|
||||
description=desc,
|
||||
replace_start=insert_pos,
|
||||
replace_end=insert_pos,
|
||||
)
|
||||
for name, desc in completions
|
||||
],
|
||||
)
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
all_matches = [v for vs in by_cat.values() for v in vs]
|
||||
if not all_matches:
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
# Build candidates ordered by category
|
||||
candidates: list[CompletionCandidate] = []
|
||||
for cat in _CATEGORY_ORDER:
|
||||
for cmd_text, desc in by_cat.get(cat, []):
|
||||
candidates.append(
|
||||
CompletionCandidate(
|
||||
text=cmd_text,
|
||||
description=desc,
|
||||
replace_start=0,
|
||||
replace_end=len(before),
|
||||
category=cat,
|
||||
)
|
||||
)
|
||||
for cat, items in by_cat.items():
|
||||
if cat not in _CATEGORY_ORDER:
|
||||
for cmd_text, desc in items:
|
||||
candidates.append(
|
||||
CompletionCandidate(
|
||||
text=cmd_text,
|
||||
description=desc,
|
||||
replace_start=0,
|
||||
replace_end=len(before),
|
||||
category=cat,
|
||||
)
|
||||
)
|
||||
|
||||
return CompletionResult(CompletionKind.COMMANDS, candidates)
|
||||
|
||||
# --- Subcommand / argument completion (len(parts) >= 2) ---
|
||||
cmd = cmd_manager.get_command(cmd_name)
|
||||
if cmd is None:
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
# Delegate to Command.get_completions for all depths
|
||||
tokens = parts[1:]
|
||||
if has_trailing_space:
|
||||
tokens.append("")
|
||||
completions = cmd.get_completions(tokens)
|
||||
if not completions:
|
||||
return CompletionResult(CompletionKind.EMPTY, [])
|
||||
|
||||
# Compute replacement range.
|
||||
if tokens[-1]:
|
||||
# User is typing a partial — replace it
|
||||
sub_start = before.rfind(tokens[-1])
|
||||
if sub_start < 0:
|
||||
sub_start = len(before)
|
||||
replace_end = len(before)
|
||||
elif len(tokens) >= 2 and tokens[-2]:
|
||||
# Trailing space after a token. Check if the previous token is
|
||||
# a known subcommand name — if so, the completion is for the
|
||||
# NEXT argument (insert at cursor). If not, the completions
|
||||
# refine the partial (replace it).
|
||||
prev = tokens[-2]
|
||||
is_known_sub = any(sc.name == prev for sc in cmd.subcommands)
|
||||
if not is_known_sub:
|
||||
sub_start = before.rfind(prev)
|
||||
if sub_start < 0:
|
||||
sub_start = len(before)
|
||||
else:
|
||||
sub_start = len(before)
|
||||
replace_end = len(before)
|
||||
else:
|
||||
sub_start = len(before)
|
||||
replace_end = len(before)
|
||||
|
||||
return CompletionResult(
|
||||
CompletionKind.SUBCOMMANDS,
|
||||
[
|
||||
CompletionCandidate(
|
||||
text=name,
|
||||
description=desc,
|
||||
replace_start=sub_start,
|
||||
replace_end=replace_end,
|
||||
)
|
||||
for name, desc in completions
|
||||
],
|
||||
)
|
||||
@@ -1,8 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, ClassVar, Protocol, runtime_checkable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -15,6 +18,15 @@ class Argument:
|
||||
required: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class SubCommand:
|
||||
"""A subcommand of a parent slash command."""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
arguments: list[Argument] = field(default_factory=list)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CommandUI(Protocol):
|
||||
"""Protocol for UI operations that commands can perform."""
|
||||
@@ -38,13 +50,29 @@ class CommandUI(Protocol):
|
||||
def clear_chat(self) -> None: ...
|
||||
def request_quit(self) -> None: ...
|
||||
def force_quit(self) -> None: ...
|
||||
def start_new_session(self) -> None: ...
|
||||
async def start_new_session(self) -> None: ...
|
||||
async def handle_session_resume(
|
||||
self, thread_id: str, workspace_dir: str | None = None
|
||||
) -> None: ...
|
||||
async def flush(self) -> None: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelRuntime:
|
||||
"""Mutable handle to the agent + thread bound to running channels."""
|
||||
|
||||
agent: Any = None
|
||||
thread_id: str | None = None
|
||||
|
||||
def bind(self, agent: Any, thread_id: str) -> None:
|
||||
self.agent = agent
|
||||
self.thread_id = thread_id
|
||||
|
||||
def clear(self) -> None:
|
||||
self.agent = None
|
||||
self.thread_id = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class CommandContext:
|
||||
"""Context passed to commands during execution."""
|
||||
@@ -55,6 +83,9 @@ class CommandContext:
|
||||
workspace_dir: str | None = None
|
||||
checkpointer: Any = None
|
||||
config: Any = None
|
||||
channel_runtime: ChannelRuntime | None = None
|
||||
graph_gateway: GraphGateway | None = None
|
||||
command_error: str | None = None
|
||||
# Real LLM input token count from last usage_metadata (includes system
|
||||
# prompt + tool schemas). Used by /compact for accurate display.
|
||||
input_tokens_hint: int | None = None
|
||||
@@ -67,6 +98,53 @@ class Command(ABC):
|
||||
alias: ClassVar[list[str]] = []
|
||||
description: str
|
||||
arguments: ClassVar[list[Argument]] = []
|
||||
category: ClassVar[str] = "General"
|
||||
subcommands: ClassVar[list[SubCommand]] = []
|
||||
# When False, callers may dispatch this command without waiting for
|
||||
# the background agent load to finish — important so recovery
|
||||
# commands like ``/mcp add`` can run even when the MCP load is
|
||||
# failing and ``_await_agent_ready`` would hang.
|
||||
requires_agent: ClassVar[bool] = False
|
||||
|
||||
def needs_agent(self, args: list[str]) -> bool:
|
||||
"""Whether this specific invocation needs the agent.
|
||||
|
||||
Default returns :attr:`requires_agent`. Override when a command
|
||||
has a mix of agent-using and agent-free subcommands (e.g.
|
||||
``/channel start`` vs ``/channel status``).
|
||||
"""
|
||||
return self.requires_agent
|
||||
|
||||
def get_completions(self, tokens: list[str]) -> list[tuple[str, str]]:
|
||||
"""Return completions for args typed after the command name.
|
||||
|
||||
Default walks :attr:`subcommands` for the first positional token
|
||||
only. Override for deeper levels (e.g. server names, thread IDs).
|
||||
"""
|
||||
if not self.subcommands:
|
||||
return []
|
||||
if len(tokens) <= 1:
|
||||
prefix = tokens[0].lower() if tokens else ""
|
||||
matches = [
|
||||
(sc.name, sc.description)
|
||||
for sc in self.subcommands
|
||||
if sc.name.startswith(prefix)
|
||||
]
|
||||
# Exact match — subcommand already complete, hide popup
|
||||
if len(matches) == 1 and matches[0][0] == prefix:
|
||||
return []
|
||||
return matches
|
||||
# partial + trailing space: /mcp a → still show "add"
|
||||
if len(tokens) == 2 and tokens[1] == "":
|
||||
prefix = tokens[0].lower()
|
||||
if any(sc.name == prefix for sc in self.subcommands):
|
||||
return []
|
||||
return [
|
||||
(sc.name, sc.description)
|
||||
for sc in self.subcommands
|
||||
if sc.name.startswith(prefix)
|
||||
]
|
||||
return []
|
||||
|
||||
@abstractmethod
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
|
||||
@@ -1,14 +1,23 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from .base import CommandUI
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ChannelCommandUI(CommandUI):
|
||||
"""CommandUI implementation for messaging channels with output buffering."""
|
||||
|
||||
_TEXT_CHUNK_LIMIT = 3500
|
||||
|
||||
@property
|
||||
def supports_interactive(self) -> bool:
|
||||
return False
|
||||
@@ -16,24 +25,67 @@ class ChannelCommandUI(CommandUI):
|
||||
def __init__(
|
||||
self,
|
||||
channel_msg: Any,
|
||||
*,
|
||||
graph_gateway: GraphGateway,
|
||||
append_system_callback: Any = None,
|
||||
start_new_session_callback: Any = None,
|
||||
start_new_session_callback: Callable[[], Awaitable[None]] | None = None,
|
||||
handle_session_resume_callback: Any = None,
|
||||
):
|
||||
self.msg = channel_msg
|
||||
self.append_system_callback = append_system_callback
|
||||
self.start_new_session_callback = start_new_session_callback
|
||||
self.handle_session_resume_callback = handle_session_resume_callback
|
||||
self.graph_gateway = graph_gateway
|
||||
self._system_buffer: list[str] = []
|
||||
|
||||
def append_system(self, text: str, style: str = "dim") -> None:
|
||||
if self.append_system_callback:
|
||||
def _queue_system(
|
||||
self,
|
||||
text: str,
|
||||
style: str = "dim",
|
||||
*,
|
||||
mirror_local: bool = True,
|
||||
) -> None:
|
||||
if mirror_local and self.append_system_callback:
|
||||
self.append_system_callback(text, style)
|
||||
|
||||
# Buffer the text for grouped delivery to the channel
|
||||
# We ignore style for grouping but keep it for individual lines if needed
|
||||
self._system_buffer.append(text)
|
||||
|
||||
def append_system(self, text: str, style: str = "dim") -> None:
|
||||
self._queue_system(text, style)
|
||||
|
||||
@staticmethod
|
||||
def _extract_message_text(message: Any) -> str:
|
||||
content = getattr(message, "content", "") or ""
|
||||
if isinstance(content, list):
|
||||
parts = [
|
||||
block.get("text", "")
|
||||
for block in content
|
||||
if isinstance(block, dict) and block.get("type") == "text"
|
||||
]
|
||||
content = " ".join(parts) if parts else ""
|
||||
return str(content).strip()
|
||||
|
||||
async def _send_text_chunks(self, text: str, *, mirror_local: bool = True) -> None:
|
||||
"""Flush long plain-text payloads in channel-safe chunks."""
|
||||
text = (text or "").strip()
|
||||
if not text:
|
||||
return
|
||||
|
||||
pending = text
|
||||
while pending:
|
||||
chunk = pending[: self._TEXT_CHUNK_LIMIT]
|
||||
if len(pending) > self._TEXT_CHUNK_LIMIT:
|
||||
split_at = chunk.rfind("\n")
|
||||
if split_at > 0:
|
||||
chunk = chunk[:split_at]
|
||||
chunk = chunk.rstrip()
|
||||
if not chunk:
|
||||
chunk = pending[: self._TEXT_CHUNK_LIMIT]
|
||||
self._queue_system(chunk, mirror_local=mirror_local)
|
||||
await self.flush()
|
||||
pending = pending[len(chunk) :].lstrip("\n")
|
||||
|
||||
async def flush(self) -> None:
|
||||
"""Send all buffered system messages as a single grouped message."""
|
||||
if not self._system_buffer:
|
||||
@@ -129,9 +181,9 @@ class ChannelCommandUI(CommandUI):
|
||||
def force_quit(self) -> None:
|
||||
self.request_quit()
|
||||
|
||||
def start_new_session(self) -> None:
|
||||
async def start_new_session(self) -> None:
|
||||
if self.start_new_session_callback:
|
||||
self.start_new_session_callback()
|
||||
await self.start_new_session_callback()
|
||||
else:
|
||||
self.append_system(
|
||||
"New session requested. Please restart the channel link or use /new if supported."
|
||||
@@ -140,5 +192,45 @@ class ChannelCommandUI(CommandUI):
|
||||
async def handle_session_resume(
|
||||
self, thread_id: str, workspace_dir: str | None = None
|
||||
) -> None:
|
||||
mirror_local = self.handle_session_resume_callback is None
|
||||
if self.handle_session_resume_callback:
|
||||
await self.handle_session_resume_callback(thread_id, workspace_dir)
|
||||
lines = [f"Resumed session: {thread_id}"]
|
||||
try:
|
||||
messages = await self.graph_gateway.get_thread_messages(thread_id)
|
||||
except Exception as exc:
|
||||
_logger.exception(
|
||||
"Failed to load saved history for resumed thread %s",
|
||||
thread_id,
|
||||
)
|
||||
lines.append(f"(history unavailable: {exc})")
|
||||
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
|
||||
return
|
||||
|
||||
display = [m for m in messages if getattr(m, "type", None) in ("human", "ai")]
|
||||
|
||||
if not display:
|
||||
if messages:
|
||||
lines.append("No displayable messages in this session.")
|
||||
else:
|
||||
lines.append("No saved messages in this session.")
|
||||
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
|
||||
return
|
||||
|
||||
HISTORY_WINDOW = 10
|
||||
if len(display) > HISTORY_WINDOW:
|
||||
display = display[-HISTORY_WINDOW:]
|
||||
lines.append(f"Conversation history (last {HISTORY_WINDOW} messages):")
|
||||
else:
|
||||
lines.append("Conversation history:")
|
||||
|
||||
for message in display:
|
||||
text = self._extract_message_text(message)
|
||||
if not text:
|
||||
continue
|
||||
if getattr(message, "type", None) == "human":
|
||||
lines.append(f"User: {text}")
|
||||
else:
|
||||
lines.append(f"EvoScientist: {text}")
|
||||
|
||||
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
|
||||
|
||||
@@ -1,5 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from . import channel, general, mcp, session, skills
|
||||
from . import (
|
||||
autoskills,
|
||||
channel,
|
||||
general,
|
||||
mcp,
|
||||
schedule,
|
||||
session,
|
||||
skills,
|
||||
)
|
||||
|
||||
__all__ = ["channel", "general", "mcp", "session", "skills"]
|
||||
__all__ = [
|
||||
"autoskills",
|
||||
"channel",
|
||||
"general",
|
||||
"mcp",
|
||||
"schedule",
|
||||
"session",
|
||||
"skills",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,435 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from enum import Enum
|
||||
from typing import ClassVar
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ..base import Command, CommandContext, SubCommand
|
||||
from ..manager import manager
|
||||
|
||||
AUTOSKILLS_COMMAND = "/autoskills"
|
||||
_PROPOSAL_STATUSES = {
|
||||
"review": "pending",
|
||||
"approved": "approved",
|
||||
"rejected": "rejected",
|
||||
}
|
||||
|
||||
|
||||
class AutoSkillsCommand(Command):
|
||||
"""Manage EvoMemory AutoSkills proposals."""
|
||||
|
||||
name = AUTOSKILLS_COMMAND
|
||||
alias: ClassVar[list[str]] = ["/skills-review"]
|
||||
description = "Review EvoMemory autoskill proposals"
|
||||
subcommands: ClassVar[list[SubCommand]] = [
|
||||
SubCommand("status", "Show AutoSkills config and proposals for review"),
|
||||
SubCommand("help", "Show AutoSkills command examples"),
|
||||
SubCommand("list", "List autoskill proposals, optionally filtered by status"),
|
||||
SubCommand("review", "Review autoskill proposals awaiting a decision"),
|
||||
SubCommand("approve", "Approve an autoskill proposal by id"),
|
||||
SubCommand("reject", "Reject an autoskill proposal by id"),
|
||||
SubCommand("run", "Run AutoSkills once now"),
|
||||
SubCommand("on", "Enable periodic AutoSkills"),
|
||||
SubCommand("off", "Disable periodic AutoSkills"),
|
||||
SubCommand("mode", "Set review or auto approval mode"),
|
||||
SubCommand("cadence", "Set nightly, weekly, or monthly cadence"),
|
||||
SubCommand("time", "Set local run time as HH:MM"),
|
||||
]
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
sub = args[0].lower() if args else "help"
|
||||
rest = args[1:]
|
||||
if sub in {"help", "-h", "--help", "?"}:
|
||||
self._show_help(ctx)
|
||||
elif sub in {"status", "show"}:
|
||||
await self._status(ctx)
|
||||
elif sub in {"list", "ls", "proposals"}:
|
||||
await self._list_command(ctx, rest)
|
||||
elif sub == "review":
|
||||
await self._list(ctx, status="pending")
|
||||
elif sub in {"approve", "accept"}:
|
||||
await self._approve(ctx, self._first_arg(rest))
|
||||
elif sub in {"reject", "deny", "decline"}:
|
||||
await self._reject(ctx, self._first_arg(rest))
|
||||
elif sub in {"run", "now"}:
|
||||
await self._run(ctx)
|
||||
elif sub in {"on", "enable"}:
|
||||
await self._set_config(ctx, "memory_skill_synthesis_enabled", "true")
|
||||
elif sub in {"off", "disable"}:
|
||||
await self._set_config(ctx, "memory_skill_synthesis_enabled", "false")
|
||||
elif sub == "mode":
|
||||
await self._set_config(
|
||||
ctx,
|
||||
"memory_skill_synthesis_mode",
|
||||
self._first_arg(rest),
|
||||
)
|
||||
elif sub in {"auto", "automatic"}:
|
||||
await self._set_config(ctx, "memory_skill_synthesis_mode", "auto")
|
||||
elif sub == "manual":
|
||||
await self._set_config(ctx, "memory_skill_synthesis_mode", "review")
|
||||
elif sub == "cadence":
|
||||
await self._set_config(
|
||||
ctx,
|
||||
"memory_skill_synthesis_cadence",
|
||||
self._first_arg(rest),
|
||||
)
|
||||
elif sub in {"nightly", "weekly", "monthly"}:
|
||||
await self._set_config(ctx, "memory_skill_synthesis_cadence", sub)
|
||||
elif sub == "time":
|
||||
await self._set_config(
|
||||
ctx,
|
||||
"memory_skill_synthesis_time",
|
||||
self._first_arg(rest),
|
||||
)
|
||||
else:
|
||||
self._show_help(ctx, prefix=f"Unknown AutoSkills command: {sub}")
|
||||
|
||||
async def _status(self, ctx: CommandContext) -> None:
|
||||
from ... import paths
|
||||
from ...config import get_effective_config
|
||||
from ...memory.autoskills.proposals import list_skill_proposals
|
||||
from ...memory.autoskills.schedule import alist_autoskill_schedules
|
||||
|
||||
cfg = get_effective_config()
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
pending = list_skill_proposals(
|
||||
paths.MEMORIES_DIR,
|
||||
status="pending",
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
(
|
||||
"AutoSkills: "
|
||||
f"{'on' if cfg.memory_skill_synthesis_enabled else 'off'} | "
|
||||
f"mode={cfg.memory_skill_synthesis_mode.value} | "
|
||||
f"cadence={cfg.memory_skill_synthesis_cadence.value} | "
|
||||
f"time={cfg.memory_skill_synthesis_time}"
|
||||
),
|
||||
style="dim",
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
f"AutoSkill proposal(s) ready for review: {len(pending)}",
|
||||
style="yellow" if pending else "dim",
|
||||
)
|
||||
if pending:
|
||||
ctx.ui.append_system(
|
||||
(
|
||||
f"Next: {AUTOSKILLS_COMMAND} review, then "
|
||||
f"{AUTOSKILLS_COMMAND} approve <id> or "
|
||||
f"{AUTOSKILLS_COMMAND} reject <id>."
|
||||
),
|
||||
style="dim",
|
||||
)
|
||||
elif cfg.memory_skill_synthesis_enabled:
|
||||
ctx.ui.append_system(
|
||||
f"Next: {AUTOSKILLS_COMMAND} run to search now, or "
|
||||
f"{AUTOSKILLS_COMMAND} help for commands.",
|
||||
style="dim",
|
||||
)
|
||||
else:
|
||||
ctx.ui.append_system(
|
||||
f"Next: {AUTOSKILLS_COMMAND} run to search once, or "
|
||||
f"{AUTOSKILLS_COMMAND} on to enable scheduled runs.",
|
||||
style="dim",
|
||||
)
|
||||
if cfg.memory_skill_synthesis_enabled:
|
||||
try:
|
||||
rows = await alist_autoskill_schedules(cfg, limit=1)
|
||||
except Exception:
|
||||
rows = []
|
||||
if rows:
|
||||
ctx.ui.append_system(
|
||||
f"Background schedule id: {str(rows[0].get('cron_id', ''))[:8]}",
|
||||
style="dim",
|
||||
)
|
||||
|
||||
async def _list_command(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
if not args or args[0].lower() == "all":
|
||||
await self._list(ctx)
|
||||
return
|
||||
status = _PROPOSAL_STATUSES.get(args[0].lower())
|
||||
if status is None:
|
||||
ctx.ui.append_system(
|
||||
f"Usage: {AUTOSKILLS_COMMAND} list [review|approved|rejected|all]",
|
||||
style="yellow",
|
||||
)
|
||||
return
|
||||
await self._list(ctx, status=status)
|
||||
|
||||
async def _list(self, ctx: CommandContext, *, status: str | None = None) -> None:
|
||||
from ... import paths
|
||||
from ...memory.autoskills.proposals import list_skill_proposals
|
||||
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
proposals = list_skill_proposals(
|
||||
paths.MEMORIES_DIR,
|
||||
status=status,
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
if not proposals:
|
||||
if status:
|
||||
label = self._status_label(status)
|
||||
ctx.ui.append_system(
|
||||
f"No autoskill proposals {label}.",
|
||||
style="dim",
|
||||
)
|
||||
else:
|
||||
ctx.ui.append_system("No autoskill proposals.", style="dim")
|
||||
return
|
||||
|
||||
title = "EvoMemory AutoSkill Proposals"
|
||||
if status:
|
||||
title = (
|
||||
f"EvoMemory AutoSkill Proposals {self._status_label(status).title()}"
|
||||
)
|
||||
table = Table(title=title, show_header=True)
|
||||
table.add_column("ID", style="cyan")
|
||||
table.add_column("Action", style="magenta")
|
||||
table.add_column("AutoSkill", style="green")
|
||||
table.add_column("Status", style="yellow")
|
||||
table.add_column("Observations", justify="right")
|
||||
table.add_column("Description", style="dim")
|
||||
for proposal in proposals:
|
||||
table.add_row(
|
||||
proposal.proposal_id,
|
||||
proposal.operation,
|
||||
proposal.skill_name,
|
||||
proposal.status,
|
||||
str(len(proposal.source_observation_ids)),
|
||||
proposal.description,
|
||||
)
|
||||
ctx.ui.mount_renderable(table)
|
||||
ctx.ui.append_system(
|
||||
f"Use {AUTOSKILLS_COMMAND} approve <id> or "
|
||||
f"{AUTOSKILLS_COMMAND} reject <id>.",
|
||||
style="dim",
|
||||
)
|
||||
|
||||
async def _approve(self, ctx: CommandContext, proposal_id: str | None) -> None:
|
||||
from ... import paths
|
||||
from ...memory.autoskills.proposals import approve_skill_proposal
|
||||
|
||||
if not proposal_id:
|
||||
ctx.ui.append_system(
|
||||
f"Usage: {AUTOSKILLS_COMMAND} approve <id>",
|
||||
style="yellow",
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
f"Run {AUTOSKILLS_COMMAND} review to copy a proposal ID.",
|
||||
style="dim",
|
||||
)
|
||||
return
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
result = await asyncio.to_thread(
|
||||
approve_skill_proposal,
|
||||
paths.MEMORIES_DIR,
|
||||
proposal_id,
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
if result.get("approved"):
|
||||
verb = "Updated" if result.get("operation") == "update" else "Approved"
|
||||
ctx.ui.append_system(
|
||||
f"{verb} autoskill: {result['skill_name']} ({result['path']})",
|
||||
style="green",
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
"Reload with /new to apply the new skill.", style="dim"
|
||||
)
|
||||
else:
|
||||
ctx.ui.append_system(f"Approval failed: {result.get('error')}", style="red")
|
||||
|
||||
async def _reject(self, ctx: CommandContext, proposal_id: str | None) -> None:
|
||||
from ... import paths
|
||||
from ...memory.autoskills.proposals import reject_skill_proposal
|
||||
|
||||
if not proposal_id:
|
||||
ctx.ui.append_system(
|
||||
f"Usage: {AUTOSKILLS_COMMAND} reject <id>",
|
||||
style="yellow",
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
f"Run {AUTOSKILLS_COMMAND} review to copy a proposal ID.",
|
||||
style="dim",
|
||||
)
|
||||
return
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
result = await asyncio.to_thread(
|
||||
reject_skill_proposal,
|
||||
paths.MEMORIES_DIR,
|
||||
proposal_id,
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
if result.get("rejected"):
|
||||
ctx.ui.append_system(
|
||||
f"Rejected proposal: {result['proposal_id']}",
|
||||
style="green",
|
||||
)
|
||||
else:
|
||||
ctx.ui.append_system(f"Reject failed: {result.get('error')}", style="red")
|
||||
|
||||
async def _run(self, ctx: CommandContext) -> None:
|
||||
from ...config import get_effective_config
|
||||
from ...memory.autoskills.schedule import arun_autoskill_now
|
||||
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
try:
|
||||
result = await arun_autoskill_now(
|
||||
get_effective_config(),
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Failed to start AutoSkills: {exc}", style="red")
|
||||
return
|
||||
ctx.ui.append_system(
|
||||
f"Started AutoSkills run {result['run_id']}.",
|
||||
style="green",
|
||||
)
|
||||
|
||||
async def _set_config(
|
||||
self,
|
||||
ctx: CommandContext,
|
||||
key: str,
|
||||
value: str | None,
|
||||
) -> None:
|
||||
from ...config import get_effective_config, set_config_value
|
||||
from ...memory.autoskills.schedule import reconcile_autoskill_schedule
|
||||
|
||||
workspace_dir = self._workspace_dir(ctx)
|
||||
if not value:
|
||||
cfg = get_effective_config()
|
||||
current = self._display_value(getattr(cfg, key))
|
||||
ctx.ui.append_system(
|
||||
f"Current {self._config_label(key)}: {current}",
|
||||
style="dim",
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
f"Usage: {self._config_usage(key)}",
|
||||
style="yellow",
|
||||
)
|
||||
return
|
||||
if not await asyncio.to_thread(set_config_value, key, value):
|
||||
valid = self._config_values(key)
|
||||
suffix = f" Valid values: {valid}." if valid else ""
|
||||
ctx.ui.append_system(
|
||||
f"Invalid value for {self._config_label(key)}: {value}.{suffix}",
|
||||
style="red",
|
||||
)
|
||||
return
|
||||
cfg = get_effective_config()
|
||||
if ctx.config is not None and hasattr(ctx.config, key):
|
||||
setattr(ctx.config, key, getattr(cfg, key))
|
||||
await asyncio.to_thread(
|
||||
reconcile_autoskill_schedule,
|
||||
cfg,
|
||||
workspace_dir=workspace_dir,
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
f"Updated {self._config_label(key)} = {self._display_value(getattr(cfg, key))}",
|
||||
style="green",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _config_label(key: str) -> str:
|
||||
labels = {
|
||||
"memory_skill_synthesis_enabled": "AutoSkills",
|
||||
"memory_skill_synthesis_mode": "AutoSkills mode",
|
||||
"memory_skill_synthesis_cadence": "AutoSkills cadence",
|
||||
"memory_skill_synthesis_time": "AutoSkills time",
|
||||
}
|
||||
return labels.get(key, key)
|
||||
|
||||
@staticmethod
|
||||
def _display_value(value: object) -> object:
|
||||
return getattr(value, "value", value)
|
||||
|
||||
@staticmethod
|
||||
def _enum_values(enum_type: type[Enum], *, separator: str = ", ") -> str:
|
||||
return separator.join(str(member.value) for member in enum_type)
|
||||
|
||||
@classmethod
|
||||
def _config_usage(cls, key: str) -> str:
|
||||
from ...config import MemorySkillSynthesisCadence, MemorySkillSynthesisMode
|
||||
|
||||
if key == "memory_skill_synthesis_mode":
|
||||
values = cls._enum_values(MemorySkillSynthesisMode, separator="|")
|
||||
return f"{AUTOSKILLS_COMMAND} mode {values}"
|
||||
if key == "memory_skill_synthesis_cadence":
|
||||
values = cls._enum_values(MemorySkillSynthesisCadence, separator="|")
|
||||
return f"{AUTOSKILLS_COMMAND} cadence {values}"
|
||||
if key == "memory_skill_synthesis_time":
|
||||
return f"{AUTOSKILLS_COMMAND} time HH:MM"
|
||||
return f"{AUTOSKILLS_COMMAND} <value>"
|
||||
|
||||
@classmethod
|
||||
def _config_values(cls, key: str) -> str | None:
|
||||
from ...config import MemorySkillSynthesisCadence, MemorySkillSynthesisMode
|
||||
|
||||
if key == "memory_skill_synthesis_mode":
|
||||
return cls._enum_values(MemorySkillSynthesisMode)
|
||||
if key == "memory_skill_synthesis_cadence":
|
||||
return cls._enum_values(MemorySkillSynthesisCadence)
|
||||
if key == "memory_skill_synthesis_time":
|
||||
return "24-hour local time, for example 03:00"
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _status_label(status: str) -> str:
|
||||
if status == "pending":
|
||||
return "ready for review"
|
||||
return status
|
||||
|
||||
@staticmethod
|
||||
def _show_help(ctx: CommandContext, *, prefix: str | None = None) -> None:
|
||||
if prefix:
|
||||
ctx.ui.append_system(prefix, style="yellow")
|
||||
ctx.ui.append_system(
|
||||
(
|
||||
f"Usage: {AUTOSKILLS_COMMAND} "
|
||||
"[status|review|approve|reject|run|on|off|mode|cadence|time]"
|
||||
),
|
||||
style="bold",
|
||||
)
|
||||
table = Table(title="AutoSkills Commands", show_header=True)
|
||||
table.add_column("Command", style="cyan")
|
||||
table.add_column("Use when", style="dim")
|
||||
rows = [
|
||||
(AUTOSKILLS_COMMAND, "Show this command reference"),
|
||||
(f"{AUTOSKILLS_COMMAND} status", "Show config and the next useful action"),
|
||||
(f"{AUTOSKILLS_COMMAND} review", "Review proposals waiting for a decision"),
|
||||
(f"{AUTOSKILLS_COMMAND} approve <id>", "Install a reviewed autoskill"),
|
||||
(f"{AUTOSKILLS_COMMAND} reject <id>", "Dismiss a reviewed proposal"),
|
||||
(f"{AUTOSKILLS_COMMAND} run", "Start a one-off background autoskill run"),
|
||||
(f"{AUTOSKILLS_COMMAND} on|off", "Enable or disable scheduled runs"),
|
||||
(f"{AUTOSKILLS_COMMAND} auto|manual", "Switch approval behavior"),
|
||||
(
|
||||
f"{AUTOSKILLS_COMMAND} nightly|weekly|monthly",
|
||||
"Set the built-in schedule cadence",
|
||||
),
|
||||
(f"{AUTOSKILLS_COMMAND} time 03:00", "Set the local schedule time"),
|
||||
(
|
||||
f"{AUTOSKILLS_COMMAND} list [status]",
|
||||
"List all proposals or filter by review, approved, or rejected",
|
||||
),
|
||||
]
|
||||
for command, description in rows:
|
||||
table.add_row(command, description)
|
||||
ctx.ui.mount_renderable(table)
|
||||
ctx.ui.append_system(
|
||||
"Aliases: /skills-review, ls, proposals, accept, deny, enable, disable, now.",
|
||||
style="dim",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _workspace_dir(ctx: CommandContext) -> str:
|
||||
from ... import paths
|
||||
|
||||
return str(ctx.workspace_dir or paths.WORKSPACE_ROOT)
|
||||
|
||||
@staticmethod
|
||||
def _first_arg(args: list[str]) -> str | None:
|
||||
return args[0] if args else None
|
||||
|
||||
|
||||
manager.register(AutoSkillsCommand())
|
||||
@@ -1,10 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import ClassVar
|
||||
|
||||
from rich.panel import Panel
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
from ..base import Command, CommandContext
|
||||
from ..base import Command, CommandContext, SubCommand
|
||||
from ..manager import manager
|
||||
|
||||
|
||||
@@ -13,6 +15,27 @@ class ChannelCommand(Command):
|
||||
|
||||
name = "/channel"
|
||||
description = "Configure messaging channels"
|
||||
category = "Channels"
|
||||
subcommands: ClassVar[list[SubCommand]] = [
|
||||
SubCommand("status", "Show channel status"),
|
||||
SubCommand("stop", "Stop running channels"),
|
||||
SubCommand("telegram", "Start Telegram channel"),
|
||||
SubCommand("discord", "Start Discord channel"),
|
||||
SubCommand("slack", "Start Slack channel"),
|
||||
SubCommand("feishu", "Start Feishu channel"),
|
||||
SubCommand("dingtalk", "Start DingTalk channel"),
|
||||
SubCommand("wechat", "Start WeChat channel"),
|
||||
SubCommand("email", "Start Email channel"),
|
||||
SubCommand("imessage", "Start iMessage channel"),
|
||||
]
|
||||
|
||||
def needs_agent(self, args: list[str]) -> bool:
|
||||
# ``status`` and ``stop`` are introspection / teardown; they
|
||||
# must work even when the agent load is still in flight or has
|
||||
# failed. Only start/add flows feed ``ctx.agent`` into
|
||||
# ``_start_channels_bus_mode``.
|
||||
subcmd = args[0].lower() if args else ""
|
||||
return subcmd not in {"status", "stop"}
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
import EvoScientist.cli.channel as _ch_mod
|
||||
@@ -63,10 +86,10 @@ class ChannelCommand(Command):
|
||||
ctx.ui.append_system("No channels are running.", style="dim")
|
||||
else:
|
||||
if target:
|
||||
_channels_stop(target)
|
||||
_channels_stop(target, runtime=ctx.channel_runtime)
|
||||
ctx.ui.append_system(f"Channel '{target}' stopped.", style="green")
|
||||
else:
|
||||
_channels_stop()
|
||||
_channels_stop(runtime=ctx.channel_runtime)
|
||||
ctx.ui.append_system("All channels stopped.", style="green")
|
||||
return
|
||||
|
||||
@@ -90,6 +113,11 @@ class ChannelCommand(Command):
|
||||
ctx.ui.append_system(f"Adding channel(s): {', '.join(requested)}...")
|
||||
from ...cli.channel import _add_channel_to_running_bus
|
||||
|
||||
# Bind the runtime up-front so partial-success states (one
|
||||
# channel attached, next one raises) still leave the bus
|
||||
# observing the latest agent/thread refs.
|
||||
if ctx.channel_runtime is not None:
|
||||
ctx.channel_runtime.bind(ctx.agent, ctx.thread_id)
|
||||
try:
|
||||
for ct in requested:
|
||||
_add_channel_to_running_bus(ct, config, send_thinking=send_thinking)
|
||||
@@ -123,6 +151,8 @@ class ChannelCommand(Command):
|
||||
ctx.thread_id,
|
||||
send_thinking=send_thinking,
|
||||
)
|
||||
if ctx.channel_runtime is not None:
|
||||
ctx.channel_runtime.bind(ctx.agent, ctx.thread_id)
|
||||
|
||||
# Show status panel
|
||||
if _ch_mod._manager:
|
||||
|
||||
@@ -29,6 +29,9 @@ class HelpCommand(Command):
|
||||
if cmd.alias:
|
||||
desc += f" (aliases: {', '.join(cmd.alias)})"
|
||||
help_text.append(f"{desc}\n", style="dim")
|
||||
if cmd.subcommands:
|
||||
names = ", ".join(sc.name for sc in cmd.subcommands)
|
||||
help_text.append(f" subcommands: {names}\n", style="dim italic")
|
||||
ctx.ui.mount_renderable(help_text)
|
||||
|
||||
|
||||
@@ -49,17 +52,13 @@ class CurrentCommand(Command):
|
||||
f"Workspace: {_shorten_path(ctx.workspace_dir)}",
|
||||
style="dim",
|
||||
)
|
||||
memory_path = paths.MEMORY_DIR
|
||||
memory_path = paths.MEMORIES_DIR
|
||||
if memory_path:
|
||||
from ...cli.agent import _shorten_path
|
||||
|
||||
ctx.ui.append_system(
|
||||
f"Memory dir: {_shorten_path(str(memory_path))}", style="dim"
|
||||
)
|
||||
# How to determine UI type here?
|
||||
# Maybe ctx.ui has a name or we pass it in ctx.
|
||||
# For now, let's keep it simple.
|
||||
ctx.ui.append_system("UI: auto", style="dim")
|
||||
|
||||
|
||||
# Register commands
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import ClassVar
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ..base import Command, CommandContext
|
||||
from ..base import Command, CommandContext, SubCommand
|
||||
from ..manager import manager
|
||||
|
||||
|
||||
@@ -11,8 +13,46 @@ class MCPCommand(Command):
|
||||
|
||||
name = "/mcp"
|
||||
description = "Manage MCP servers"
|
||||
category = "MCP"
|
||||
subcommands: ClassVar[list[SubCommand]] = [
|
||||
SubCommand("list", "List configured MCP servers"),
|
||||
SubCommand("config", "Show server configuration details"),
|
||||
SubCommand("add", "Add a new MCP server"),
|
||||
SubCommand("edit", "Edit an MCP server configuration"),
|
||||
SubCommand("remove", "Remove an MCP server"),
|
||||
SubCommand("install", "Browse and install MCP servers"),
|
||||
]
|
||||
|
||||
_server_names_cache: list[str] | None = None
|
||||
|
||||
def _get_server_names(self) -> list[str]:
|
||||
if self._server_names_cache is None:
|
||||
try:
|
||||
from ...mcp import load_mcp_config
|
||||
|
||||
self._server_names_cache = list(load_mcp_config().keys())
|
||||
except Exception:
|
||||
return []
|
||||
return self._server_names_cache
|
||||
|
||||
def _invalidate_server_cache(self) -> None:
|
||||
self._server_names_cache = None
|
||||
|
||||
def get_completions(self, tokens: list[str]) -> list[tuple[str, str]]:
|
||||
if len(tokens) <= 1:
|
||||
return super().get_completions(tokens)
|
||||
subcmd = tokens[0].lower()
|
||||
if subcmd in ("config", "remove", "edit") and len(tokens) == 2:
|
||||
prefix = tokens[1].lower()
|
||||
return [
|
||||
(name, "")
|
||||
for name in self._get_server_names()
|
||||
if name.lower().startswith(prefix)
|
||||
]
|
||||
return super().get_completions(tokens)
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
"""Dispatch to the appropriate MCP subcommand."""
|
||||
if not args or args[0] == "list":
|
||||
await self._mcp_list(ctx)
|
||||
return
|
||||
@@ -24,35 +64,26 @@ class MCPCommand(Command):
|
||||
await self._mcp_config(ctx, subargs[0] if subargs else "")
|
||||
elif subcmd == "add":
|
||||
await self._mcp_add(ctx, subargs)
|
||||
self._invalidate_server_cache()
|
||||
elif subcmd == "edit":
|
||||
await self._mcp_edit(ctx, subargs)
|
||||
self._invalidate_server_cache()
|
||||
elif subcmd == "remove":
|
||||
await self._mcp_remove(ctx, subargs[0] if subargs else "")
|
||||
self._invalidate_server_cache()
|
||||
elif subcmd == "install":
|
||||
from .mcp_install import InstallMCPCommand
|
||||
|
||||
await InstallMCPCommand().execute(ctx, subargs)
|
||||
else:
|
||||
ctx.ui.append_system("MCP commands:", style="bold")
|
||||
for sub in self.subcommands:
|
||||
ctx.ui.append_system(
|
||||
" /mcp List configured servers", style="dim"
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
" /mcp list List configured servers", style="dim"
|
||||
)
|
||||
ctx.ui.append_system(
|
||||
" /mcp config Show detailed server config", style="dim"
|
||||
)
|
||||
ctx.ui.append_system(" /mcp add ... Add a server", style="dim")
|
||||
ctx.ui.append_system(
|
||||
" /mcp edit ... Edit an existing server", style="dim"
|
||||
)
|
||||
ctx.ui.append_system(" /mcp remove ... Remove a server", style="dim")
|
||||
ctx.ui.append_system(
|
||||
" /mcp install ... Browse and install servers", style="dim"
|
||||
f" /mcp {sub.name:<12} {sub.description}", style="dim"
|
||||
)
|
||||
|
||||
async def _mcp_list(self, ctx: CommandContext) -> None:
|
||||
"""Display a table of all configured MCP servers."""
|
||||
from ...mcp import load_mcp_config
|
||||
from ...mcp.client import USER_MCP_CONFIG
|
||||
|
||||
@@ -85,6 +116,7 @@ class MCPCommand(Command):
|
||||
ctx.ui.append_system(f"Config file: {USER_MCP_CONFIG}", style="dim")
|
||||
|
||||
async def _mcp_config(self, ctx: CommandContext, name: str) -> None:
|
||||
"""Show detailed configuration for one or all MCP servers."""
|
||||
from ...mcp import load_mcp_config
|
||||
from ...mcp.client import USER_MCP_CONFIG
|
||||
|
||||
@@ -133,6 +165,7 @@ class MCPCommand(Command):
|
||||
ctx.ui.append_system(f"Config file: {USER_MCP_CONFIG}", style="dim")
|
||||
|
||||
async def _mcp_add(self, ctx: CommandContext, tokens: list[str]) -> None:
|
||||
"""Add a new MCP server from parsed arguments."""
|
||||
from ...mcp import add_mcp_server, parse_mcp_add_args
|
||||
|
||||
if not tokens:
|
||||
@@ -153,6 +186,7 @@ class MCPCommand(Command):
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
|
||||
async def _mcp_edit(self, ctx: CommandContext, tokens: list[str]) -> None:
|
||||
"""Edit fields of an existing MCP server configuration."""
|
||||
from ...mcp import edit_mcp_server, parse_mcp_edit_args
|
||||
|
||||
if not tokens:
|
||||
@@ -170,6 +204,7 @@ class MCPCommand(Command):
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
|
||||
async def _mcp_remove(self, ctx: CommandContext, name: str) -> None:
|
||||
"""Remove an MCP server by name."""
|
||||
from ...mcp import remove_mcp_server
|
||||
|
||||
if not name:
|
||||
|
||||
@@ -10,6 +10,7 @@ class InstallMCPCommand(Command):
|
||||
|
||||
name = "/install-mcp"
|
||||
description = "Browse and install MCP servers"
|
||||
category = "MCP"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="source",
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
from typing import ClassVar
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ..base import Command, CommandContext, SubCommand
|
||||
from ..manager import manager
|
||||
|
||||
|
||||
class ScheduleCommand(Command):
|
||||
"""Manage scheduled (cron) tasks."""
|
||||
|
||||
name = "/schedule"
|
||||
description = "Manage scheduled (cron) tasks"
|
||||
subcommands: ClassVar[list[SubCommand]] = [
|
||||
SubCommand("add", 'Add: /schedule add <m h dom mon dow> "<prompt>"'),
|
||||
SubCommand("list", "List scheduled tasks"),
|
||||
SubCommand("remove", "Remove a schedule by id"),
|
||||
SubCommand("run", "Run a schedule's prompt once now (test)"),
|
||||
SubCommand("pause", "Disable a schedule by id"),
|
||||
SubCommand("resume", "Enable a schedule by id"),
|
||||
]
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
"""Dispatch to the appropriate /schedule subcommand."""
|
||||
cfg = getattr(ctx, "config", None)
|
||||
if cfg is not None and not getattr(cfg, "enable_scheduler", True):
|
||||
ctx.ui.append_system(
|
||||
"Scheduled tasks are disabled (`enable_scheduler` is off).",
|
||||
style="yellow",
|
||||
)
|
||||
return
|
||||
|
||||
from ...cron import schedule as crons
|
||||
|
||||
# Cron SDK calls are sync HTTP; offload to a thread so backend latency
|
||||
# can never freeze the interactive event loop.
|
||||
if not await asyncio.to_thread(crons.is_available):
|
||||
ctx.ui.append_system(
|
||||
"Scheduler unavailable: the langgraph dev backend is not running.",
|
||||
style="yellow",
|
||||
)
|
||||
return
|
||||
|
||||
if not args or args[0].lower() == "list":
|
||||
await self._list(ctx, crons)
|
||||
return
|
||||
|
||||
sub = args[0].lower()
|
||||
rest = args[1:]
|
||||
if sub == "add":
|
||||
await self._add(ctx, crons, rest)
|
||||
elif sub == "remove":
|
||||
await self._remove(ctx, crons, rest[0] if rest else "")
|
||||
elif sub == "run":
|
||||
await self._run(ctx, crons, rest[0] if rest else "")
|
||||
elif sub in ("pause", "resume"):
|
||||
await self._set_enabled(
|
||||
ctx, crons, rest[0] if rest else "", sub == "resume"
|
||||
)
|
||||
else:
|
||||
ctx.ui.append_system("Schedule commands:", style="bold")
|
||||
for s in self.subcommands:
|
||||
ctx.ui.append_system(
|
||||
f" /schedule {s.name:<8} {s.description}", style="dim"
|
||||
)
|
||||
|
||||
async def _add(self, ctx: CommandContext, crons, rest: list[str]) -> None:
|
||||
# Cron may arrive as 5 separate tokens (unquoted) or 1 token (shlex-quoted).
|
||||
# Split on any whitespace so extra spaces don't break detection; the
|
||||
# backend rejects genuinely malformed expressions.
|
||||
if rest and len(rest[0].split()) == 5: # quoted 5-field cron
|
||||
schedule, prompt_tokens = " ".join(rest[0].split()), rest[1:]
|
||||
elif len(rest) >= 5: # 5 separate cron fields
|
||||
schedule, prompt_tokens = " ".join(rest[:5]), rest[5:]
|
||||
else:
|
||||
ctx.ui.append_system(
|
||||
'Usage: /schedule add "<m h dom mon dow>" "<prompt>"', style="yellow"
|
||||
)
|
||||
return
|
||||
prompt = " ".join(prompt_tokens).strip().strip('"').strip("'")
|
||||
if not prompt:
|
||||
ctx.ui.append_system("A task prompt is required.", style="yellow")
|
||||
return
|
||||
# B3: strip unsafe chars; keep only alphanumerics + hyphens (kebab-case).
|
||||
raw = prompt[:48].lower()
|
||||
name = re.sub(r"[^a-z0-9]+", "-", raw).strip("-")[:32] or "task"
|
||||
try:
|
||||
rec = await asyncio.to_thread(
|
||||
crons.create_schedule, name=name, schedule=schedule, prompt=prompt
|
||||
)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
return
|
||||
ctx.ui.append_system(
|
||||
f"Scheduled '{name}' [{schedule}] — id {rec.get('cron_id')}. "
|
||||
"Runs unattended in the background.",
|
||||
style="green",
|
||||
)
|
||||
|
||||
async def _list(self, ctx: CommandContext, crons) -> None:
|
||||
# B1: guard SDK call — backend may die after the is_available() check.
|
||||
try:
|
||||
rows = await asyncio.to_thread(crons.list_schedules)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
return
|
||||
if not rows:
|
||||
ctx.ui.append_system(
|
||||
"No scheduled tasks. Add one: /schedule add ...", style="dim"
|
||||
)
|
||||
return
|
||||
table = Table(title="Scheduled Tasks", show_header=True)
|
||||
table.add_column("ID", style="cyan")
|
||||
table.add_column("Name", style="magenta")
|
||||
table.add_column("Schedule", style="green")
|
||||
table.add_column("Enabled", style="yellow")
|
||||
table.add_column("Next run (UTC)", style="white")
|
||||
for r in rows:
|
||||
meta = r.get("metadata") or {}
|
||||
table.add_row(
|
||||
str(r.get("cron_id", ""))[:8],
|
||||
str(meta.get("name", "")),
|
||||
str(r.get("schedule", "")),
|
||||
"yes" if r.get("enabled", True) else "no",
|
||||
str(r.get("next_run_date", "")),
|
||||
)
|
||||
ctx.ui.mount_renderable(table)
|
||||
|
||||
_AMBIGUOUS = object() # B2: sentinel returned when multiple crons match a prefix
|
||||
_BACKEND_ERROR = object() # sentinel returned when list_schedules() raises
|
||||
|
||||
async def _resolve(self, crons, prefix: str):
|
||||
"""Return the unique matching record, _AMBIGUOUS if >1 match, _BACKEND_ERROR on error, or None."""
|
||||
# B1: guard SDK call — backend may die after is_available() check.
|
||||
try:
|
||||
all_rows = await asyncio.to_thread(crons.list_schedules)
|
||||
except Exception as exc:
|
||||
# Store the exception text so _resolve_or_report can surface it.
|
||||
self._last_backend_exc = exc
|
||||
return self._BACKEND_ERROR
|
||||
# B2: collect ALL matches; ambiguous prefix → sentinel so callers can warn.
|
||||
matches = [r for r in all_rows if str(r.get("cron_id", "")).startswith(prefix)]
|
||||
if len(matches) > 1:
|
||||
return self._AMBIGUOUS
|
||||
return matches[0] if matches else None
|
||||
|
||||
async def _resolve_or_report(self, ctx: CommandContext, crons, prefix: str):
|
||||
"""Resolve prefix → record, emit UI error on ambiguity/miss/error, return None on failure."""
|
||||
match = await self._resolve(crons, prefix)
|
||||
if match is self._BACKEND_ERROR:
|
||||
exc = getattr(self, "_last_backend_exc", None)
|
||||
ctx.ui.append_system(
|
||||
f"Error: scheduler backend unavailable ({exc})",
|
||||
style="red",
|
||||
)
|
||||
return None
|
||||
if match is self._AMBIGUOUS:
|
||||
ctx.ui.append_system(
|
||||
f"Multiple schedules match '{prefix}' — use a longer id.",
|
||||
style="yellow",
|
||||
)
|
||||
return None
|
||||
if not match:
|
||||
ctx.ui.append_system(f"No schedule matching {prefix}.", style="yellow")
|
||||
return None
|
||||
return match
|
||||
|
||||
async def _remove(self, ctx: CommandContext, crons, prefix: str) -> None:
|
||||
if not prefix:
|
||||
ctx.ui.append_system("Usage: /schedule remove <id>", style="yellow")
|
||||
return
|
||||
match = await self._resolve_or_report(ctx, crons, prefix)
|
||||
if match is None:
|
||||
return
|
||||
cron_id = str(match.get("cron_id", ""))
|
||||
try:
|
||||
await asyncio.to_thread(crons.delete_schedule, cron_id)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
return
|
||||
ctx.ui.append_system(f"Removed schedule {cron_id}.", style="green")
|
||||
|
||||
async def _run(self, ctx: CommandContext, crons, prefix: str) -> None:
|
||||
if not prefix:
|
||||
ctx.ui.append_system("Usage: /schedule run <id>", style="yellow")
|
||||
return
|
||||
match = await self._resolve_or_report(ctx, crons, prefix)
|
||||
if match is None:
|
||||
return
|
||||
prompt = (match.get("metadata") or {}).get("prompt", "")
|
||||
if not str(prompt).strip():
|
||||
ctx.ui.append_system(
|
||||
f"Schedule {prefix} has no stored prompt — cannot run it.",
|
||||
style="yellow",
|
||||
)
|
||||
return
|
||||
try:
|
||||
rec = await asyncio.to_thread(crons.run_now, prompt)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
return
|
||||
# Don't promise a location; the task's own prompt decides where output goes.
|
||||
ctx.ui.append_system(
|
||||
f"Fired schedule {prefix} once now (run {rec.get('run_id')}). "
|
||||
"Any output goes wherever the task's instruction specifies.",
|
||||
style="green",
|
||||
)
|
||||
|
||||
async def _set_enabled(
|
||||
self, ctx: CommandContext, crons, prefix: str, enabled: bool
|
||||
) -> None:
|
||||
if not prefix:
|
||||
ctx.ui.append_system("Usage: /schedule pause|resume <id>", style="yellow")
|
||||
return
|
||||
match = await self._resolve_or_report(ctx, crons, prefix)
|
||||
if match is None:
|
||||
return
|
||||
cron_id = str(match.get("cron_id", ""))
|
||||
try:
|
||||
await asyncio.to_thread(crons.set_enabled, cron_id, enabled)
|
||||
except Exception as exc:
|
||||
ctx.ui.append_system(f"Error: {exc}", style="red")
|
||||
return
|
||||
ctx.ui.append_system(
|
||||
f"{'Resumed' if enabled else 'Paused'} schedule {cron_id}.", style="green"
|
||||
)
|
||||
|
||||
|
||||
# Register schedule command
|
||||
manager.register(ScheduleCommand())
|
||||
@@ -5,15 +5,24 @@ from typing import ClassVar
|
||||
|
||||
from rich.table import Table
|
||||
|
||||
from ...gateway import GraphGateway, GraphTarget
|
||||
from ..base import Argument, Command, CommandContext
|
||||
from ..manager import manager
|
||||
|
||||
|
||||
def _graph_gateway(ctx: CommandContext) -> GraphGateway:
|
||||
if ctx.graph_gateway is None:
|
||||
raise RuntimeError("Session commands require a graph_gateway")
|
||||
return ctx.graph_gateway
|
||||
|
||||
|
||||
class CompactCommand(Command):
|
||||
"""Compact conversation to free context."""
|
||||
|
||||
name = "/compact"
|
||||
description = "Compact conversation to free context"
|
||||
requires_agent = True
|
||||
category = "Session"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
from ...cli.commands import (
|
||||
@@ -35,8 +44,12 @@ class CompactCommand(Command):
|
||||
|
||||
try:
|
||||
result = await compact_conversation(
|
||||
agent=ctx.agent,
|
||||
graph_gateway=_graph_gateway(ctx),
|
||||
thread_id=ctx.thread_id,
|
||||
target=GraphTarget(
|
||||
local_graph=ctx.agent,
|
||||
workspace_dir=ctx.workspace_dir,
|
||||
),
|
||||
input_tokens_hint=ctx.input_tokens_hint,
|
||||
)
|
||||
finally:
|
||||
@@ -70,11 +83,13 @@ class ThreadsCommand(Command):
|
||||
|
||||
name = "/threads"
|
||||
description = "List recent sessions"
|
||||
category = "Session"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
from ...sessions import _format_relative_time, list_threads
|
||||
from ...sessions import _format_relative_time, short_thread_id
|
||||
|
||||
threads = await list_threads(
|
||||
gateway = _graph_gateway(ctx)
|
||||
threads = await gateway.list_threads(
|
||||
limit=0,
|
||||
include_message_count=True,
|
||||
include_preview=True,
|
||||
@@ -102,7 +117,7 @@ class ThreadsCommand(Command):
|
||||
marker = " *" if thread_id_value == ctx.thread_id else ""
|
||||
|
||||
row = [
|
||||
f"{thread_id_value}{marker}",
|
||||
f"{short_thread_id(thread_id_value)}{marker}",
|
||||
thread.get("preview", "") or "",
|
||||
str(thread.get("message_count", 0)),
|
||||
]
|
||||
@@ -112,6 +127,11 @@ class ThreadsCommand(Command):
|
||||
|
||||
table.add_row(*row)
|
||||
ctx.ui.mount_renderable(table)
|
||||
if not is_channel:
|
||||
ctx.ui.append_system(
|
||||
" /resume to continue a session "
|
||||
"/delete <id> to remove /new to start fresh",
|
||||
)
|
||||
|
||||
|
||||
class ResumeCommand(Command):
|
||||
@@ -119,6 +139,7 @@ class ResumeCommand(Command):
|
||||
|
||||
name = "/resume"
|
||||
description = "Resume a previous session"
|
||||
category = "Session"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="thread_id",
|
||||
@@ -129,14 +150,10 @@ class ResumeCommand(Command):
|
||||
]
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
from ...sessions import (
|
||||
get_thread_metadata,
|
||||
list_threads,
|
||||
)
|
||||
|
||||
gateway = _graph_gateway(ctx)
|
||||
arg = args[0] if args else ""
|
||||
if not arg:
|
||||
threads = await list_threads(
|
||||
threads = await gateway.list_threads(
|
||||
limit=0,
|
||||
include_message_count=True,
|
||||
include_preview=True,
|
||||
@@ -160,7 +177,7 @@ class ResumeCommand(Command):
|
||||
if not resolved:
|
||||
return
|
||||
|
||||
metadata = await get_thread_metadata(resolved)
|
||||
metadata = await gateway.get_thread_metadata(resolved)
|
||||
restored_workspace = (metadata or {}).get("workspace_dir", "")
|
||||
if restored_workspace:
|
||||
ctx.workspace_dir = restored_workspace
|
||||
@@ -172,21 +189,16 @@ class ResumeCommand(Command):
|
||||
await ctx.ui.handle_session_resume(resolved, restored_workspace)
|
||||
|
||||
async def _resolve_thread_id(self, prefix: str, ctx: CommandContext) -> str | None:
|
||||
from ...sessions import find_similar_threads, thread_exists
|
||||
resolution = await _graph_gateway(ctx).resolve_thread(prefix)
|
||||
if resolution.thread_id:
|
||||
return resolution.thread_id
|
||||
|
||||
if await thread_exists(prefix):
|
||||
return prefix
|
||||
|
||||
similar = await find_similar_threads(prefix)
|
||||
if len(similar) == 1:
|
||||
return similar[0]
|
||||
|
||||
if len(similar) > 1:
|
||||
if resolution.matches:
|
||||
ctx.ui.append_system(
|
||||
f"Ambiguous thread ID '{prefix}'. Use a longer prefix.",
|
||||
style="yellow",
|
||||
)
|
||||
for thread in similar:
|
||||
for thread in resolution.matches:
|
||||
ctx.ui.append_system(f" - {thread}", style="dim")
|
||||
return None
|
||||
|
||||
@@ -199,9 +211,10 @@ class NewCommand(Command):
|
||||
|
||||
name = "/new"
|
||||
description = "Start a new session"
|
||||
category = "Session"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
ctx.ui.start_new_session()
|
||||
await ctx.ui.start_new_session()
|
||||
|
||||
|
||||
class ClearCommand(Command):
|
||||
@@ -209,6 +222,7 @@ class ClearCommand(Command):
|
||||
|
||||
name = "/clear"
|
||||
description = "Clear chat history"
|
||||
category = "Session"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
ctx.ui.clear_chat()
|
||||
@@ -219,6 +233,7 @@ class DeleteCommand(Command):
|
||||
|
||||
name = "/delete"
|
||||
description = "Delete a saved session"
|
||||
category = "Session"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="thread_id",
|
||||
@@ -229,16 +244,10 @@ class DeleteCommand(Command):
|
||||
]
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
from ...sessions import (
|
||||
delete_thread,
|
||||
find_similar_threads,
|
||||
list_threads,
|
||||
thread_exists,
|
||||
)
|
||||
|
||||
gateway = _graph_gateway(ctx)
|
||||
arg = args[0] if args else ""
|
||||
if not arg:
|
||||
threads = await list_threads(
|
||||
threads = await gateway.list_threads(
|
||||
limit=0,
|
||||
include_message_count=True,
|
||||
include_preview=True,
|
||||
@@ -258,22 +267,17 @@ class DeleteCommand(Command):
|
||||
arg = selected
|
||||
|
||||
# Resolve thread_id
|
||||
resolved = None
|
||||
if await thread_exists(arg):
|
||||
resolved = arg
|
||||
else:
|
||||
similar = await find_similar_threads(arg)
|
||||
if len(similar) == 1:
|
||||
resolved = similar[0]
|
||||
elif len(similar) > 1:
|
||||
resolution = await gateway.resolve_thread(arg)
|
||||
if resolution.matches:
|
||||
ctx.ui.append_system(
|
||||
f"Ambiguous thread ID '{arg}'. Use a longer prefix.",
|
||||
style="yellow",
|
||||
)
|
||||
for thread in similar:
|
||||
for thread in resolution.matches:
|
||||
ctx.ui.append_system(f" - {thread}", style="dim")
|
||||
return
|
||||
|
||||
resolved = resolution.thread_id
|
||||
if not resolved:
|
||||
ctx.ui.append_system(f"Session '{arg}' not found.", style="red")
|
||||
return
|
||||
@@ -285,7 +289,7 @@ class DeleteCommand(Command):
|
||||
)
|
||||
return
|
||||
|
||||
deleted = await delete_thread(resolved)
|
||||
deleted = await gateway.delete_thread(resolved)
|
||||
if deleted:
|
||||
ctx.ui.append_system(f"Deleted session {resolved}.", style="green")
|
||||
else:
|
||||
@@ -298,6 +302,7 @@ class ExitCommand(Command):
|
||||
name = "/exit"
|
||||
alias: ClassVar[list[str]] = ["/quit", "/q"]
|
||||
description = "Quit EvoScientist"
|
||||
category = "Session"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
ctx.ui.force_quit()
|
||||
|
||||
@@ -13,6 +13,7 @@ class SkillsCommand(Command):
|
||||
|
||||
name = "/skills"
|
||||
description = "List installed skills"
|
||||
category = "Skills"
|
||||
|
||||
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
|
||||
from ...cli.agent import _shorten_path
|
||||
@@ -65,6 +66,7 @@ class InstallSkill(Command):
|
||||
|
||||
name: ClassVar[str] = "/install-skill"
|
||||
description: ClassVar[str] = "Add a skill from path or GitHub"
|
||||
category: ClassVar[str] = "Skills"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="source",
|
||||
@@ -142,6 +144,7 @@ class InstallSkills(Command):
|
||||
description: ClassVar[str] = (
|
||||
"Browse and install EvoSkills (optional: /evoskills <tag>)"
|
||||
)
|
||||
category: ClassVar[str] = "Skills"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="tag", type=str, description="Tag to filter skills by", required=False
|
||||
@@ -213,11 +216,18 @@ class InstallSkills(Command):
|
||||
pre_filter_tag=tag,
|
||||
)
|
||||
|
||||
if not selected_sources:
|
||||
# ``None`` means user cancelled (Esc / Ctrl-C). An empty list means
|
||||
# the picker handled a "nothing to do" state (all-installed / no
|
||||
# tag matches) and already printed its own specific message; the
|
||||
# outer layer should stay silent rather than claim a cancel.
|
||||
if selected_sources is None:
|
||||
if not is_channel:
|
||||
ctx.ui.append_system("Browse cancelled.", style="dim")
|
||||
return
|
||||
|
||||
if not selected_sources:
|
||||
return
|
||||
|
||||
# Install selected skills
|
||||
installed_count = 0
|
||||
for source in selected_sources:
|
||||
@@ -248,6 +258,7 @@ class UninstallSkill(Command):
|
||||
|
||||
name: ClassVar[str] = "/uninstall-skill"
|
||||
description: ClassVar[str] = "Remove an installed skill"
|
||||
category: ClassVar[str] = "Skills"
|
||||
arguments: ClassVar[list[Argument]] = [
|
||||
Argument(
|
||||
name="name",
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
import logging
|
||||
import shlex
|
||||
|
||||
from .base import Command, CommandContext
|
||||
from .base import Command, CommandContext, SubCommand
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -27,6 +27,27 @@ class CommandManager:
|
||||
"""Lookup a command by name."""
|
||||
return self._commands.get(name.lower())
|
||||
|
||||
def resolve(self, command_str: str) -> tuple[Command, list[str]] | None:
|
||||
"""Return ``(command, args)`` for the dispatch of ``command_str``.
|
||||
|
||||
Uses the same parsing as :meth:`execute` so callers can inspect
|
||||
metadata (e.g. call :meth:`Command.needs_agent`) without
|
||||
re-implementing ``shlex`` quirks.
|
||||
"""
|
||||
command_str = command_str.strip()
|
||||
if not command_str:
|
||||
return None
|
||||
try:
|
||||
parts = shlex.split(command_str)
|
||||
except ValueError:
|
||||
parts = command_str.split()
|
||||
if not parts:
|
||||
return None
|
||||
cmd = self.get_command(parts[0])
|
||||
if cmd is None:
|
||||
return None
|
||||
return cmd, parts[1:]
|
||||
|
||||
def list_commands(self) -> list[tuple[str, str]]:
|
||||
"""List all registered command names and descriptions."""
|
||||
seen = set()
|
||||
@@ -37,6 +58,17 @@ class CommandManager:
|
||||
seen.add(cmd)
|
||||
return results
|
||||
|
||||
def get_subcommands(self, command_name: str) -> list[SubCommand]:
|
||||
"""Return subcommands declared by *command_name*, or empty list."""
|
||||
cmd = self.get_command(command_name)
|
||||
if cmd is None:
|
||||
return []
|
||||
return cmd.subcommands
|
||||
|
||||
def list_subcommands(self, command_name: str) -> list[tuple[str, str]]:
|
||||
"""Return ``(name, description)`` pairs for completion rendering."""
|
||||
return [(sc.name, sc.description) for sc in self.get_subcommands(command_name)]
|
||||
|
||||
def get_all_commands(self) -> list[Command]:
|
||||
"""Return all registered command instances."""
|
||||
seen = set()
|
||||
@@ -72,12 +104,14 @@ class CommandManager:
|
||||
if not cmd:
|
||||
return False
|
||||
|
||||
ctx.command_error = None
|
||||
try:
|
||||
await cmd.execute(ctx, args)
|
||||
await ctx.ui.flush()
|
||||
return True
|
||||
except Exception as e:
|
||||
_logger.exception(f"Error executing command {cmd_name}: {e}")
|
||||
ctx.command_error = str(e)
|
||||
ctx.ui.append_system(f"Error executing {cmd_name}: {e}", style="red")
|
||||
await ctx.ui.flush()
|
||||
return True
|
||||
|
||||
@@ -10,11 +10,18 @@ The onboard module is loaded lazily because it pulls in heavy dependencies
|
||||
|
||||
from .settings import (
|
||||
EvoScientistConfig,
|
||||
MemoryControls,
|
||||
MemoryObservationTarget,
|
||||
MemoryObservationWriter,
|
||||
MemorySkillSynthesisCadence,
|
||||
MemorySkillSynthesisMode,
|
||||
apply_config_to_env,
|
||||
get_config_dir,
|
||||
get_config_path,
|
||||
get_config_value,
|
||||
get_default_workspace_dir,
|
||||
get_effective_config,
|
||||
is_config_applied_env,
|
||||
list_config,
|
||||
load_config,
|
||||
reset_config,
|
||||
@@ -24,12 +31,19 @@ from .settings import (
|
||||
|
||||
__all__ = [
|
||||
"EvoScientistConfig",
|
||||
"MemoryControls",
|
||||
"MemoryObservationTarget",
|
||||
"MemoryObservationWriter",
|
||||
"MemorySkillSynthesisCadence",
|
||||
"MemorySkillSynthesisMode",
|
||||
"apply_config_to_env",
|
||||
# settings
|
||||
"get_config_dir",
|
||||
"get_config_path",
|
||||
"get_config_value",
|
||||
"get_default_workspace_dir",
|
||||
"get_effective_config",
|
||||
"is_config_applied_env",
|
||||
"list_config",
|
||||
"load_config",
|
||||
"reset_config",
|
||||
|
||||
@@ -1,316 +0,0 @@
|
||||
"""Configuration helpers for dedicated image-generation models.
|
||||
|
||||
Image generation models are service/tool models, not chat models. Keeping
|
||||
them in a separate config section prevents image-only models such as
|
||||
``gpt-image-2`` from being offered in the normal chat model selector.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
|
||||
IMAGE_GENERATION_SECTION = "image_generation"
|
||||
DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS = 120.0
|
||||
IMAGE_GENERATION_USAGE_NOTES = [
|
||||
"Normal chat model lists filter out image-only models such as gpt-image-* and dall-e-*.",
|
||||
"If an image-only model is accidentally saved in LLM settings, it is moved into image_generation.",
|
||||
"If the frontend sends an image-only model as the chat model, the backend returns IMAGE_MODEL_NOT_CHAT_MODEL.",
|
||||
]
|
||||
_IMAGE_MODEL_PREFIXES = (
|
||||
"gpt-image",
|
||||
"chatgpt-image",
|
||||
"dall-e",
|
||||
"dalle",
|
||||
)
|
||||
_IMAGE_MODEL_MARKERS = (
|
||||
"wanx",
|
||||
"seedream",
|
||||
)
|
||||
|
||||
|
||||
class ImageModelEntry(BaseModel):
|
||||
"""A single dedicated image-generation model."""
|
||||
|
||||
id: str = Field(..., description="Model ID sent to the image API")
|
||||
name: str = Field("", description="Display name or short alias")
|
||||
provider: str = Field("openai-compatible", description="Provider label")
|
||||
api_key: str = Field("", description="API key or ${ENV_VAR} reference")
|
||||
base_url: str = Field("", description="API base URL or ${ENV_VAR} reference")
|
||||
supports_generation: bool = True
|
||||
supports_edit: bool = True
|
||||
default_size: str = "1024x1024"
|
||||
default_quality: str = "auto"
|
||||
params: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@field_validator("id")
|
||||
@classmethod
|
||||
def id_not_empty(cls, value: str) -> str:
|
||||
value = value.strip()
|
||||
if not value:
|
||||
raise ValueError("image model id must not be empty")
|
||||
return value
|
||||
|
||||
def resolved_api_key(self) -> str:
|
||||
return _resolve_env_ref(self.api_key)
|
||||
|
||||
def resolved_base_url(self) -> str:
|
||||
return _resolve_env_ref(self.base_url)
|
||||
|
||||
def display_name(self) -> str:
|
||||
return self.name or self.id
|
||||
|
||||
|
||||
class ImageGenerationSettings(BaseModel):
|
||||
"""Dedicated image-generation model settings."""
|
||||
|
||||
default_model: str = ""
|
||||
timeout_seconds: float = Field(
|
||||
DEFAULT_IMAGE_GENERATION_TIMEOUT_SECONDS,
|
||||
description="Per-request timeout for image provider HTTP calls",
|
||||
)
|
||||
models: list[ImageModelEntry] = Field(default_factory=list)
|
||||
|
||||
@field_validator("timeout_seconds")
|
||||
@classmethod
|
||||
def timeout_must_be_positive(cls, value: float) -> float:
|
||||
if value <= 0:
|
||||
raise ValueError("image generation timeout_seconds must be greater than 0")
|
||||
return value
|
||||
|
||||
|
||||
def is_image_generation_model(model_ref: str | None) -> bool:
|
||||
"""Return True when a model ID is known to be image-generation only."""
|
||||
if not model_ref:
|
||||
return False
|
||||
value = str(model_ref).strip().lower()
|
||||
if "/" in value:
|
||||
value = value.rsplit("/", 1)[1]
|
||||
return value.startswith(_IMAGE_MODEL_PREFIXES) or any(
|
||||
marker in value for marker in _IMAGE_MODEL_MARKERS
|
||||
)
|
||||
|
||||
|
||||
def load_image_generation_settings(
|
||||
*,
|
||||
include_legacy: bool = True,
|
||||
resolve_env: bool = True,
|
||||
) -> ImageGenerationSettings:
|
||||
"""Load dedicated image-generation settings with legacy fallback."""
|
||||
raw = _load_settings_yaml()
|
||||
legacy_present = _legacy_config_present(raw)
|
||||
section = raw.get(IMAGE_GENERATION_SECTION)
|
||||
if isinstance(section, dict):
|
||||
source = _resolve_nested_env(section) if resolve_env else section
|
||||
settings = ImageGenerationSettings.model_validate(source)
|
||||
else:
|
||||
settings = ImageGenerationSettings()
|
||||
|
||||
legacy = _legacy_model_entry(raw)
|
||||
if not settings.models:
|
||||
if include_legacy or legacy_present:
|
||||
settings.models = [legacy]
|
||||
elif include_legacy and _legacy_env_overrides_present():
|
||||
_merge_entry(settings, legacy)
|
||||
|
||||
if not settings.default_model:
|
||||
if settings.models:
|
||||
settings.default_model = settings.models[0].id
|
||||
elif include_legacy or legacy_present:
|
||||
settings.default_model = legacy.id
|
||||
return settings
|
||||
|
||||
|
||||
def save_image_generation_settings(settings: ImageGenerationSettings) -> ImageGenerationSettings:
|
||||
"""Save image-generation settings into ``settings.yaml``."""
|
||||
validated = ImageGenerationSettings.model_validate(settings.model_dump(mode="python"))
|
||||
raw = _load_settings_yaml()
|
||||
raw[IMAGE_GENERATION_SECTION] = validated.model_dump(mode="python")
|
||||
_atomic_write_settings_yaml(raw)
|
||||
return validated
|
||||
|
||||
|
||||
def merge_image_generation_models(
|
||||
entries: list[ImageModelEntry],
|
||||
*,
|
||||
default_model: str | None = None,
|
||||
) -> ImageGenerationSettings:
|
||||
"""Merge image model entries into the dedicated image model list."""
|
||||
raw = _load_settings_yaml()
|
||||
settings = load_image_generation_settings(include_legacy=_legacy_config_present(raw))
|
||||
for entry in entries:
|
||||
_merge_entry(settings, entry)
|
||||
if default_model:
|
||||
settings.default_model = default_model
|
||||
elif entries and not settings.default_model:
|
||||
settings.default_model = entries[0].id
|
||||
return save_image_generation_settings(settings)
|
||||
|
||||
|
||||
def resolve_image_generation_model(model_ref: str | None = None) -> ImageModelEntry:
|
||||
"""Resolve a model ID/name to a configured image model entry."""
|
||||
settings = load_image_generation_settings()
|
||||
requested = (model_ref or settings.default_model or "").strip()
|
||||
if not requested and settings.models:
|
||||
requested = settings.models[0].id
|
||||
|
||||
for entry in settings.models:
|
||||
if requested in {entry.id, entry.name}:
|
||||
return _with_legacy_fallbacks(entry)
|
||||
|
||||
if requested:
|
||||
legacy = _legacy_model_entry(_load_settings_yaml())
|
||||
if requested == legacy.id:
|
||||
return legacy
|
||||
raise ValueError(f"Unknown image generation model: {requested}")
|
||||
raise ValueError("No image generation model configured")
|
||||
|
||||
|
||||
def list_image_generation_models(*, include_sensitive: bool = False) -> dict[str, Any]:
|
||||
"""Return image-generation model settings for API/tool display."""
|
||||
settings = load_image_generation_settings()
|
||||
models = []
|
||||
for entry in settings.models:
|
||||
item = entry.model_dump(mode="python")
|
||||
item["name"] = entry.display_name()
|
||||
if include_sensitive:
|
||||
item["api_key"] = entry.resolved_api_key()
|
||||
item["base_url"] = entry.resolved_base_url()
|
||||
else:
|
||||
item["api_key"] = _mask_secret(entry.resolved_api_key())
|
||||
item["base_url"] = entry.resolved_base_url()
|
||||
models.append(item)
|
||||
return {
|
||||
"default_model": settings.default_model,
|
||||
"timeout_seconds": settings.timeout_seconds,
|
||||
"models": models,
|
||||
"usage_notes": IMAGE_GENERATION_USAGE_NOTES,
|
||||
}
|
||||
|
||||
|
||||
def _merge_entry(settings: ImageGenerationSettings, entry: ImageModelEntry) -> None:
|
||||
for idx, existing in enumerate(settings.models):
|
||||
if existing.id == entry.id:
|
||||
data = existing.model_dump(mode="python")
|
||||
update = entry.model_dump(mode="python")
|
||||
for key, value in update.items():
|
||||
if value not in ("", None, {}, []):
|
||||
data[key] = value
|
||||
settings.models[idx] = ImageModelEntry.model_validate(data)
|
||||
return
|
||||
settings.models.append(entry)
|
||||
|
||||
|
||||
def _with_legacy_fallbacks(entry: ImageModelEntry) -> ImageModelEntry:
|
||||
legacy = _legacy_model_entry(_load_settings_yaml())
|
||||
data = entry.model_dump(mode="python")
|
||||
if not data.get("api_key") or _is_masked_secret(str(data.get("api_key") or "")):
|
||||
data["api_key"] = legacy.api_key
|
||||
if not data.get("base_url"):
|
||||
data["base_url"] = legacy.base_url
|
||||
return ImageModelEntry.model_validate(data)
|
||||
|
||||
|
||||
def _legacy_env_overrides_present() -> bool:
|
||||
return any(os.environ.get(key) for key in ("IMAGE_GEN_MODEL", "IMAGE_GEN_API_KEY", "IMAGE_GEN_BASE_URL"))
|
||||
|
||||
|
||||
def _legacy_config_present(raw: dict[str, Any]) -> bool:
|
||||
return _legacy_env_overrides_present() or any(
|
||||
str(raw.get(key) or "").strip()
|
||||
for key in ("image_gen_model", "image_gen_api_key", "image_gen_base_url")
|
||||
)
|
||||
|
||||
|
||||
def _legacy_model_entry(raw: dict[str, Any]) -> ImageModelEntry:
|
||||
model = (
|
||||
os.environ.get("IMAGE_GEN_MODEL")
|
||||
or str(raw.get("image_gen_model") or "").strip()
|
||||
or "dall-e-3"
|
||||
)
|
||||
api_key = (
|
||||
os.environ.get("IMAGE_GEN_API_KEY")
|
||||
or str(raw.get("image_gen_api_key") or "").strip()
|
||||
or os.environ.get("OPENAI_API_KEY", "")
|
||||
)
|
||||
base_url = (
|
||||
os.environ.get("IMAGE_GEN_BASE_URL")
|
||||
or str(raw.get("image_gen_base_url") or "").strip()
|
||||
or os.environ.get("OPENAI_BASE_URL", "")
|
||||
)
|
||||
return ImageModelEntry(
|
||||
id=model,
|
||||
name=model,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
provider="openai-compatible",
|
||||
)
|
||||
|
||||
|
||||
def _load_settings_yaml() -> dict[str, Any]:
|
||||
from EvoScientist.config.settings import get_config_path
|
||||
|
||||
path = get_config_path()
|
||||
if not path.exists():
|
||||
return {}
|
||||
try:
|
||||
with path.open(encoding="utf-8") as fh:
|
||||
data = yaml.safe_load(fh) or {}
|
||||
except Exception:
|
||||
return {}
|
||||
return data if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _atomic_write_settings_yaml(data: dict[str, Any]) -> None:
|
||||
from EvoScientist.config.settings import get_config_path
|
||||
|
||||
path = get_config_path()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fd, tmp_path = tempfile.mkstemp(
|
||||
dir=str(path.parent),
|
||||
prefix=".settings_",
|
||||
suffix=".yaml.tmp",
|
||||
text=True,
|
||||
)
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as fh:
|
||||
yaml.safe_dump(data, fh, default_flow_style=False, sort_keys=False)
|
||||
fh.flush()
|
||||
os.fsync(fh.fileno())
|
||||
os.replace(tmp_path, path)
|
||||
finally:
|
||||
if os.path.exists(tmp_path):
|
||||
os.unlink(tmp_path)
|
||||
|
||||
|
||||
def _resolve_nested_env(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {k: _resolve_nested_env(v) for k, v in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [_resolve_nested_env(v) for v in value]
|
||||
if isinstance(value, str):
|
||||
return _resolve_env_ref(value)
|
||||
return value
|
||||
|
||||
|
||||
def _resolve_env_ref(value: str) -> str:
|
||||
if isinstance(value, str) and value.startswith("${") and value.endswith("}"):
|
||||
return os.environ.get(value[2:-1], "")
|
||||
return value or ""
|
||||
|
||||
|
||||
def _mask_secret(value: str) -> str:
|
||||
if not value:
|
||||
return ""
|
||||
if len(value) <= 8:
|
||||
return "***"
|
||||
return f"{value[:3]}...{value[-4:]}"
|
||||
|
||||
|
||||
def _is_masked_secret(value: str) -> bool:
|
||||
return bool(value and (value == "***" or value == "********" or "..." in value))
|
||||
@@ -0,0 +1,132 @@
|
||||
"""Startup detection of pre-Registry legacy model configuration artifacts.
|
||||
|
||||
Design doc section 10 step 4: this project is in development and keeps no
|
||||
historical compatibility. When the config service or the CLI finds legacy
|
||||
artifacts — an old ``providers.yaml``, an old
|
||||
``run-runtime-snapshots.sqlite3``, or LLM fields left behind in
|
||||
``config.yaml`` — it must refuse to start and log an explicit reset guide
|
||||
instead of partially reading them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
|
||||
from .settings import get_config_dir
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
#: LLM configuration keys that ``config.yaml`` must no longer carry
|
||||
#: (section 10 step 2). Platform configuration (workspace, MCP, ports,
|
||||
#: storage, scheduling, security) stays; everything model/provider/credential
|
||||
#: related lives in the Model Registry (``model-runtime.sqlite3``) only.
|
||||
LEGACY_CONFIG_YAML_KEYS = frozenset(
|
||||
{
|
||||
"provider",
|
||||
"model",
|
||||
"model_catalog",
|
||||
"model_fallbacks",
|
||||
"auxiliary_provider",
|
||||
"auxiliary_model",
|
||||
"anthropic_api_key",
|
||||
"anthropic_base_url",
|
||||
"anthropic_auth_mode",
|
||||
"openai_api_key",
|
||||
"openai_auth_mode",
|
||||
"nvidia_api_key",
|
||||
"google_api_key",
|
||||
"minimax_api_key",
|
||||
"minimax_base_url",
|
||||
"siliconflow_api_key",
|
||||
"openrouter_api_key",
|
||||
"deepseek_api_key",
|
||||
"zhipu_api_key",
|
||||
"volcengine_api_key",
|
||||
"dashscope_api_key",
|
||||
"moonshot_api_key",
|
||||
"kimi_api_key",
|
||||
"custom_openai_api_key",
|
||||
"custom_openai_base_url",
|
||||
"custom_anthropic_api_key",
|
||||
"custom_anthropic_base_url",
|
||||
"ollama_base_url",
|
||||
"use_responses_api",
|
||||
"openrouter_anthropic_prompt_cache",
|
||||
}
|
||||
)
|
||||
|
||||
_LEGACY_PROVIDERS_FILE = "providers.yaml"
|
||||
_LEGACY_SNAPSHOTS_DB = "run-runtime-snapshots.sqlite3"
|
||||
|
||||
_RESET_GUIDANCE = """\
|
||||
EvoScientist no longer reads legacy model configuration (design doc §10).
|
||||
To reset the development environment:
|
||||
1. Delete {config_dir}/providers.yaml (Provider Profiles are superseded
|
||||
by the Model Registry).
|
||||
2. Delete {config_dir}/run-runtime-snapshots.sqlite3 (the old snapshot
|
||||
store; run snapshots now live in model-runtime.sqlite3).
|
||||
3. Remove the leftover LLM fields listed above from {config_path} —
|
||||
platform fields (workspace, MCP, ports, scheduling, security) stay.
|
||||
4. Configure providers/models through the Model Registry (WebUI
|
||||
configuration page or the model-registry API), run the provider test,
|
||||
then enable the models you need.
|
||||
Startup is refused — no legacy artifact is read, even partially.\
|
||||
"""
|
||||
|
||||
|
||||
class LegacyArtifactsError(RuntimeError):
|
||||
"""Raised at startup when pre-Registry configuration artifacts remain."""
|
||||
|
||||
|
||||
def find_legacy_artifacts(config_dir: Path | None = None) -> list[str]:
|
||||
"""Return human-readable descriptions of every legacy artifact found."""
|
||||
config_dir = config_dir if config_dir is not None else get_config_dir()
|
||||
found: list[str] = []
|
||||
|
||||
providers_yaml = config_dir / _LEGACY_PROVIDERS_FILE
|
||||
if providers_yaml.exists():
|
||||
found.append(f"legacy Provider Profiles file: {providers_yaml}")
|
||||
|
||||
snapshots_db = config_dir / _LEGACY_SNAPSHOTS_DB
|
||||
if snapshots_db.exists():
|
||||
found.append(f"legacy run snapshot database: {snapshots_db}")
|
||||
|
||||
config_path = config_dir / "config.yaml"
|
||||
if config_path.exists():
|
||||
try:
|
||||
with open(config_path, encoding="utf-8") as handle:
|
||||
data = yaml.safe_load(handle) or {}
|
||||
except yaml.YAMLError:
|
||||
data = {}
|
||||
if isinstance(data, dict):
|
||||
leftover = sorted(LEGACY_CONFIG_YAML_KEYS & data.keys())
|
||||
if leftover:
|
||||
found.append(
|
||||
f"leftover LLM fields in {config_path}: {', '.join(leftover)}"
|
||||
)
|
||||
return found
|
||||
|
||||
|
||||
def assert_no_legacy_artifacts(config_dir: Path | None = None) -> None:
|
||||
"""Refuse startup when any legacy model configuration artifact remains.
|
||||
|
||||
Logs the findings plus the explicit reset guide (section 10 step 4) and
|
||||
raises :class:`LegacyArtifactsError`. Nothing is read partially: the
|
||||
caller must not catch-and-continue.
|
||||
"""
|
||||
found = find_legacy_artifacts(config_dir)
|
||||
if not found:
|
||||
return
|
||||
config_dir = config_dir if config_dir is not None else get_config_dir()
|
||||
guidance = _RESET_GUIDANCE.format(
|
||||
config_dir=config_dir,
|
||||
config_path=config_dir / "config.yaml",
|
||||
)
|
||||
message = "Legacy model configuration artifacts detected:\n" + "\n".join(
|
||||
f" - {item}" for item in found
|
||||
)
|
||||
logger.error("%s\n%s", message, guidance)
|
||||
raise LegacyArtifactsError(f"{message}\n\n{guidance}")
|
||||
@@ -1,947 +0,0 @@
|
||||
"""Structured LLM provider/model configuration.
|
||||
|
||||
This module defines the YAML-driven configuration for providers, models,
|
||||
and their parameters. It supports:
|
||||
|
||||
- Multiple providers with api_key, base_url, and protocol
|
||||
- Multiple models per provider with alias, capabilities, and params
|
||||
- Three-level parameter inheritance: defaults ← model ← runtime
|
||||
- Environment variable references in api_key fields
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
|
||||
# ─── Environment variable reference pattern ──────────────────────────────
|
||||
_ENV_REF_RE = re.compile(r"^\$\{(\w+)\}$")
|
||||
|
||||
_FLAT_LLM_CONFIG_KEYS = {
|
||||
"provider_routes",
|
||||
"provider",
|
||||
"model",
|
||||
"reasoning_effort",
|
||||
"anthropic_api_key",
|
||||
"anthropic_base_url",
|
||||
"openai_api_key",
|
||||
"nvidia_api_key",
|
||||
"google_api_key",
|
||||
"minimax_api_key",
|
||||
"siliconflow_api_key",
|
||||
"openrouter_api_key",
|
||||
"deepseek_api_key",
|
||||
"zhipu_api_key",
|
||||
"volcengine_api_key",
|
||||
"dashscope_api_key",
|
||||
"moonshot_api_key",
|
||||
"kimi_api_key",
|
||||
"custom_openai_api_key",
|
||||
"custom_openai_base_url",
|
||||
"custom_anthropic_api_key",
|
||||
"custom_anthropic_base_url",
|
||||
"ollama_base_url",
|
||||
}
|
||||
|
||||
|
||||
def _resolve_env_ref(value: str) -> str:
|
||||
"""Resolve ${ENV_VAR} references in string values.
|
||||
|
||||
If the value matches the pattern ${ENV_VAR}, return the environment
|
||||
variable's value. Otherwise return the string as-is.
|
||||
"""
|
||||
m = _ENV_REF_RE.match(value.strip())
|
||||
if m:
|
||||
return os.environ.get(m.group(1), "")
|
||||
return value
|
||||
|
||||
|
||||
# ─── Weighted round-robin load balancer ───────────────────────────────────
|
||||
|
||||
# Per-provider counter for round-robin position
|
||||
_round_robin_counters: dict[str, int] = {}
|
||||
|
||||
|
||||
# ─── Endpoint call statistics ──────────────────────────────────────────
|
||||
|
||||
class EndpointStats:
|
||||
"""Thread-safe in-memory endpoint call and token usage statistics.
|
||||
|
||||
Tracks per-endpoint:
|
||||
- Call counts (from resolve_model)
|
||||
- Token usage (input_tokens, output_tokens)
|
||||
- Last call timestamp
|
||||
- Per-model call distribution
|
||||
|
||||
Stats are kept in memory. Call ``snapshot()`` for a point-in-time
|
||||
copy, ``reset()`` to clear counters, or ``summary()`` for a
|
||||
human-readable report.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
# key = (provider_name, endpoint_name)
|
||||
# value = dict with keys: calls, last_call_ts, models, input_tokens, output_tokens
|
||||
self._stats: dict[tuple[str, str], dict[str, Any]] = {}
|
||||
# Track last endpoint selected per model for token attribution
|
||||
# key = model_id, value = (provider, endpoint_name)
|
||||
self._last_endpoint: dict[str, tuple[str, str]] = {}
|
||||
# Track most recently used endpoint (for fallback when model_id is empty)
|
||||
self._last_recorded: tuple[str, str] | None = None
|
||||
# Stack of pending endpoint attributions for sequential matching.
|
||||
# Each call to record() pushes, each record_tokens_for_model() pops.
|
||||
# This ensures tokens are attributed to the correct endpoint in tool-call loops
|
||||
# where resolve_model() is called multiple times before usage_metadata arrives.
|
||||
self._attribution_stack: list[tuple[str, str]] = []
|
||||
|
||||
def _ensure_entry(self, key: tuple[str, str]) -> dict[str, Any]:
|
||||
"""Get or create a stats entry for the given key."""
|
||||
return self._stats.setdefault(key, {
|
||||
"calls": 0,
|
||||
"last_call_ts": 0.0,
|
||||
"models": {},
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
})
|
||||
|
||||
def record(self, provider: str, endpoint: str, model_id: str) -> None:
|
||||
"""Record one endpoint selection (called from resolve_model)."""
|
||||
with self._lock:
|
||||
entry = self._ensure_entry((provider, endpoint))
|
||||
entry["calls"] += 1
|
||||
entry["last_call_ts"] = time.time()
|
||||
entry["models"][model_id] = entry["models"].get(model_id, 0) + 1
|
||||
# Remember last endpoint for this model (for token attribution)
|
||||
self._last_endpoint[model_id] = (provider, endpoint)
|
||||
# Track most recently used endpoint overall
|
||||
self._last_recorded = (provider, endpoint)
|
||||
# Push to attribution stack for sequential token matching
|
||||
self._attribution_stack.append((provider, endpoint))
|
||||
|
||||
def record_tokens(
|
||||
self,
|
||||
provider: str,
|
||||
endpoint: str,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
) -> None:
|
||||
"""Record token usage for an endpoint.
|
||||
|
||||
Can be called after an API response is received to accumulate
|
||||
token counts. The endpoint entry is created if it doesn't exist
|
||||
(e.g. for single-endpoint providers that aren't tracked by
|
||||
``record()``).
|
||||
"""
|
||||
if not input_tokens and not output_tokens:
|
||||
return
|
||||
with self._lock:
|
||||
entry = self._ensure_entry((provider, endpoint))
|
||||
entry["input_tokens"] += input_tokens
|
||||
entry["output_tokens"] += output_tokens
|
||||
|
||||
def record_tokens_for_model(
|
||||
self,
|
||||
model_id: str,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
) -> None:
|
||||
"""Record token usage, auto-routing to the correct endpoint.
|
||||
|
||||
Resolution order:
|
||||
1. Attribution stack — pops the oldest pending endpoint (FIFO matching
|
||||
with record() calls, handles tool-call loops correctly).
|
||||
2. ``_last_endpoint[model_id]`` — direct model→endpoint mapping.
|
||||
3. Scan ``_stats`` for any endpoint that has this model registered.
|
||||
4. ``_last_recorded`` — fallback for empty model_id.
|
||||
5. ``("unknown", "unknown")`` — last resort (debug level).
|
||||
"""
|
||||
if not input_tokens and not output_tokens:
|
||||
return
|
||||
with self._lock:
|
||||
provider, endpoint = None, None
|
||||
|
||||
# Strategy 1: Pop from attribution stack (matches record() calls in order)
|
||||
if self._attribution_stack:
|
||||
provider, endpoint = self._attribution_stack.pop(0)
|
||||
|
||||
# Strategy 2: Direct lookup by model_id
|
||||
if provider is None and model_id:
|
||||
p, e = self._last_endpoint.get(model_id, (None, None))
|
||||
if p is not None:
|
||||
provider, endpoint = p, e
|
||||
|
||||
# Strategy 3: Scan _stats for this model
|
||||
if provider is None and model_id:
|
||||
for (p, e), data in self._stats.items():
|
||||
if model_id in data.get("models", {}):
|
||||
provider, endpoint = p, e
|
||||
self._last_endpoint[model_id] = (p, e)
|
||||
break
|
||||
|
||||
# Strategy 4: Use most recently recorded endpoint
|
||||
if provider is None and self._last_recorded:
|
||||
provider, endpoint = self._last_recorded
|
||||
|
||||
# Strategy 5: Unknown — no resolve_model() was called beforehand
|
||||
if provider is None:
|
||||
# When model_id is empty and all state is empty, this is a
|
||||
# known race condition (TOCTOU between events.py guard and
|
||||
# this method's lock). Silently skip — no useful attribution
|
||||
# is possible and logging it just creates noise.
|
||||
if not model_id and not self._last_recorded:
|
||||
return
|
||||
provider, endpoint = "unknown", "unknown"
|
||||
# Common when LLM calls bypass resolve_model() (e.g. LangChain
|
||||
# internal bindings, sub-agents). Debug level to avoid log spam.
|
||||
logger.debug(
|
||||
"EndpointStats: model %r not found, stack=%d _last_endpoint=%s "
|
||||
"_last_recorded=%s",
|
||||
model_id,
|
||||
len(self._attribution_stack),
|
||||
list(self._last_endpoint.keys()),
|
||||
self._last_recorded,
|
||||
)
|
||||
entry = self._ensure_entry((provider, endpoint))
|
||||
entry["input_tokens"] += input_tokens
|
||||
entry["output_tokens"] += output_tokens
|
||||
|
||||
def snapshot(self) -> dict[str, Any]:
|
||||
"""Return a deep copy of current stats."""
|
||||
with self._lock:
|
||||
import copy
|
||||
return copy.deepcopy(self._stats)
|
||||
|
||||
async def restore_today_from_db(self) -> int:
|
||||
"""Restore today's endpoint stats from endpoint_usage_daily.
|
||||
|
||||
Called once at gateway startup so that the in-memory EndpointStats
|
||||
reflects today's accumulated usage after a process restart.
|
||||
|
||||
Returns the number of rows restored.
|
||||
"""
|
||||
try:
|
||||
from EvoScientist.runtime_integrations import current_date, get_app_connection
|
||||
|
||||
db = await get_app_connection()
|
||||
today = current_date()
|
||||
rows = await db.execute_fetchall(
|
||||
"""
|
||||
SELECT provider, endpoint, model, calls,
|
||||
input_tokens, output_tokens
|
||||
FROM endpoint_usage_daily
|
||||
WHERE date = $1
|
||||
""",
|
||||
(today,),
|
||||
)
|
||||
with self._lock:
|
||||
for r in rows:
|
||||
key = (r["provider"], r["endpoint"])
|
||||
entry = self._ensure_entry(key)
|
||||
entry["calls"] += r["calls"]
|
||||
entry["input_tokens"] += r["input_tokens"]
|
||||
entry["output_tokens"] += r["output_tokens"]
|
||||
if r["model"]:
|
||||
entry["models"][r["model"]] = (
|
||||
entry["models"].get(r["model"], 0) + r["calls"]
|
||||
)
|
||||
self._last_endpoint[r["model"]] = key
|
||||
restored = len(rows)
|
||||
if restored:
|
||||
logger.info(
|
||||
"EndpointStats: restored %d endpoint entries from DB for today",
|
||||
restored,
|
||||
)
|
||||
return restored
|
||||
except Exception:
|
||||
logger.warning("EndpointStats: failed to restore from DB", exc_info=True)
|
||||
return 0
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Clear all counters."""
|
||||
with self._lock:
|
||||
self._stats.clear()
|
||||
|
||||
@staticmethod
|
||||
def _fmt_tokens(n: int) -> str:
|
||||
"""Format token count compactly."""
|
||||
if n >= 1_000_000:
|
||||
return f"{n / 1_000_000:.1f}M"
|
||||
if n >= 1_000:
|
||||
return f"{n / 1_000:.1f}K"
|
||||
return str(n)
|
||||
|
||||
def summary(self) -> str:
|
||||
"""Human-readable summary for CLI / logging."""
|
||||
snap = self.snapshot()
|
||||
if not snap:
|
||||
return "No endpoint calls recorded."
|
||||
|
||||
total_calls = sum(d["calls"] for d in snap.values())
|
||||
total_input = sum(d["input_tokens"] for d in snap.values())
|
||||
total_output = sum(d["output_tokens"] for d in snap.values())
|
||||
|
||||
if total_calls == 0:
|
||||
return "No endpoint calls recorded."
|
||||
|
||||
lines: list[str] = []
|
||||
for (provider, endpoint), data in sorted(snap.items()):
|
||||
calls = data["calls"]
|
||||
pct = calls / total_calls * 100
|
||||
elapsed = time.time() - data["last_call_ts"] if data["last_call_ts"] else 0
|
||||
if elapsed < 60:
|
||||
ago = f"{elapsed:.0f}s ago"
|
||||
elif elapsed < 3600:
|
||||
ago = f"{elapsed / 60:.0f}m ago"
|
||||
else:
|
||||
ago = f"{elapsed / 3600:.1f}h ago"
|
||||
|
||||
bar = "█" * int(pct / 5) + "░" * (20 - int(pct / 5))
|
||||
inp = self._fmt_tokens(data["input_tokens"])
|
||||
out = self._fmt_tokens(data["output_tokens"])
|
||||
models = ", ".join(f"{m}({c})" for m, c in sorted(data["models"].items()))
|
||||
|
||||
token_str = ""
|
||||
if data["input_tokens"] or data["output_tokens"]:
|
||||
token_str = f" in={inp} out={out}"
|
||||
|
||||
lines.append(
|
||||
f" {provider}/{endpoint} {bar} {calls} ({pct:.0f}%) last {ago}{token_str}\n"
|
||||
f" models: {models}"
|
||||
)
|
||||
|
||||
header = f"Endpoint usage — {total_calls} calls"
|
||||
if total_input or total_output:
|
||||
header = (
|
||||
f"Endpoint usage — {total_calls} calls, "
|
||||
f"{self._fmt_tokens(total_input)} in / {self._fmt_tokens(total_output)} out tokens"
|
||||
)
|
||||
return header + ":\n" + "\n".join(lines)
|
||||
|
||||
|
||||
# Global singleton
|
||||
_endpoint_stats = EndpointStats()
|
||||
|
||||
|
||||
def get_endpoint_stats() -> EndpointStats:
|
||||
"""Return the global endpoint call statistics instance."""
|
||||
return _endpoint_stats
|
||||
|
||||
|
||||
def _select_endpoint(
|
||||
endpoints: list,
|
||||
provider_name: str,
|
||||
preferred_name: str = "",
|
||||
) -> tuple:
|
||||
"""Select an endpoint using weighted round-robin or explicit pinning.
|
||||
|
||||
Args:
|
||||
endpoints: List of EndpointConfig objects with valid credentials.
|
||||
provider_name: Provider name for round-robin counter.
|
||||
preferred_name: If set, try to find an endpoint with this name first.
|
||||
|
||||
Returns:
|
||||
The selected EndpointConfig.
|
||||
"""
|
||||
if not endpoints:
|
||||
raise ValueError(f"No available endpoints for provider '{provider_name}'")
|
||||
|
||||
# 1. If model pins to a specific endpoint, find it
|
||||
if preferred_name:
|
||||
for ep in endpoints:
|
||||
if ep.name == preferred_name and ep.resolved_api_key():
|
||||
return ep
|
||||
|
||||
# 2. Filter to endpoints with valid credentials
|
||||
available = [ep for ep in endpoints if ep.resolved_api_key()]
|
||||
if not available:
|
||||
raise ValueError(f"No endpoints with valid credentials for provider '{provider_name}'")
|
||||
|
||||
if len(available) == 1:
|
||||
return available[0]
|
||||
|
||||
# 3. Weighted round-robin selection
|
||||
total_weight = sum(ep.weight for ep in available)
|
||||
if total_weight <= 0:
|
||||
return available[0]
|
||||
|
||||
key = provider_name
|
||||
pos = _round_robin_counters.get(key, 0)
|
||||
_round_robin_counters[key] = (pos + 1) % total_weight
|
||||
|
||||
# Walk through endpoints by cumulative weight
|
||||
cumulative = 0
|
||||
for ep in available:
|
||||
cumulative += ep.weight
|
||||
if pos < cumulative:
|
||||
return ep
|
||||
|
||||
return available[-1]
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Pydantic models for the new structured config
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class EndpointConfig(BaseModel):
|
||||
"""A single API endpoint within a provider — its own api_key, base_url, and optional weights."""
|
||||
|
||||
name: str = Field("", description="Endpoint name for reference (e.g. 'coding', 'general')")
|
||||
api_key: str = Field("", description="API key (literal or ${ENV_VAR} reference)")
|
||||
base_url: str = Field("", description="API base URL override (empty = provider default)")
|
||||
weight: int = Field(1, description="Load-balancing weight (higher = more traffic)")
|
||||
extra_body: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Extra JSON body fields for requests through this endpoint",
|
||||
)
|
||||
default_headers: dict[str, str] = Field(
|
||||
default_factory=dict,
|
||||
description="Custom HTTP headers for requests through this endpoint",
|
||||
)
|
||||
params: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Endpoint-level model/client parameter overrides",
|
||||
)
|
||||
|
||||
def resolved_api_key(self) -> str:
|
||||
"""Return the API key with environment variable references resolved."""
|
||||
return _resolve_env_ref(self.api_key) if self.api_key else ""
|
||||
|
||||
def resolved_base_url(self) -> str:
|
||||
"""Return the base URL with environment variable references resolved."""
|
||||
return _resolve_env_ref(self.base_url) if self.base_url else ""
|
||||
|
||||
|
||||
class AccessConfig(BaseModel):
|
||||
"""Plan/role access constraints for providers and models.
|
||||
|
||||
Empty lists mean unrestricted. Gateway treats the admin role as allowed.
|
||||
"""
|
||||
|
||||
allowed_plans: list[str] = Field(default_factory=list, description="Allowed subscription plans")
|
||||
allowed_roles: list[str] = Field(default_factory=list, description="Allowed user roles")
|
||||
|
||||
|
||||
class ModelEntry(BaseModel):
|
||||
"""A single model definition within a provider."""
|
||||
|
||||
id: str = Field(..., description="Full model ID sent to the API")
|
||||
alias: str = Field("", description="Short alias for easy reference")
|
||||
tier: str = Field("", description="Billing tier controlled by the server config")
|
||||
currency: str = Field("CNY", description="Settlement currency for model pricing")
|
||||
max_tokens: int = Field(4096, description="Maximum output tokens")
|
||||
supports_vision: bool = Field(False, description="Whether the model supports image inputs")
|
||||
supports_reasoning: bool = Field(False, description="Whether the model supports reasoning/thinking")
|
||||
endpoint: str = Field("", description="Preferred endpoint name (empty = auto-select via load balancing)")
|
||||
params: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Model-level parameter overrides (temperature, thinking, reasoning, etc.)",
|
||||
)
|
||||
pricing: dict[str, float | str | None] = Field(
|
||||
default_factory=dict,
|
||||
description="Optional billing price: input_per_million, output_per_million, cached_input_per_million.",
|
||||
)
|
||||
permission_mode: Literal["inherit", "custom"] = Field(
|
||||
"inherit",
|
||||
description="Model access mode: inherit provider permissions or use model-level access.",
|
||||
)
|
||||
access: AccessConfig = Field(
|
||||
default_factory=AccessConfig,
|
||||
description="Model-level access constraints used when permission_mode is custom.",
|
||||
)
|
||||
|
||||
@field_validator("id")
|
||||
@classmethod
|
||||
def id_not_empty(cls, v: str) -> str:
|
||||
if not v.strip():
|
||||
raise ValueError("model id must not be empty")
|
||||
return v.strip()
|
||||
|
||||
|
||||
class ProviderConfig(BaseModel):
|
||||
"""Configuration for a single LLM provider.
|
||||
|
||||
Supports two modes:
|
||||
1. Single endpoint (legacy): set api_key / base_url directly at provider level.
|
||||
2. Multi-endpoint (new): define an 'endpoints' list, each with its own
|
||||
api_key, base_url, weight. Models can pin to a specific endpoint or
|
||||
let the system auto-select via weighted round-robin for load balancing.
|
||||
"""
|
||||
|
||||
api_key: str = Field("", description="API key (literal or ${ENV_VAR} reference)")
|
||||
base_url: str = Field("", description="API base URL override (empty = provider default)")
|
||||
protocol: Literal["openai", "anthropic", "google-genai", "ollama", "openrouter"] = Field(
|
||||
"openai",
|
||||
description="LLM protocol to use for this provider",
|
||||
)
|
||||
extra_body: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Extra JSON body fields for all requests to this provider",
|
||||
)
|
||||
default_headers: dict[str, str] = Field(
|
||||
default_factory=dict,
|
||||
description="Custom HTTP headers for all requests to this provider",
|
||||
)
|
||||
params: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Provider-level model/client parameter defaults",
|
||||
)
|
||||
access: AccessConfig = Field(
|
||||
default_factory=AccessConfig,
|
||||
description="Provider-level access constraints.",
|
||||
)
|
||||
models: list[ModelEntry] = Field(
|
||||
default_factory=list,
|
||||
description="Models available under this provider",
|
||||
)
|
||||
endpoints: list[EndpointConfig] = Field(
|
||||
default_factory=list,
|
||||
description="Multiple API endpoints for load balancing (overrides top-level api_key/base_url when set)",
|
||||
)
|
||||
|
||||
def resolved_api_key(self) -> str:
|
||||
"""Return the API key with environment variable references resolved.
|
||||
|
||||
For multi-endpoint providers, returns the first available resolved key.
|
||||
"""
|
||||
# Multi-endpoint mode: return first non-empty key
|
||||
if self.endpoints:
|
||||
for ep in self.endpoints:
|
||||
key = ep.resolved_api_key()
|
||||
if key:
|
||||
return key
|
||||
return ""
|
||||
# Single-endpoint mode (legacy)
|
||||
return _resolve_env_ref(self.api_key) if self.api_key else ""
|
||||
|
||||
def get_resolved_endpoints(self) -> list[EndpointConfig]:
|
||||
"""Return only endpoints that have valid credentials."""
|
||||
return [ep for ep in self.endpoints if ep.resolved_api_key() or (not ep.api_key and self.resolved_api_key())]
|
||||
|
||||
def has_credentials(self) -> bool:
|
||||
"""Check if this provider has any usable credentials."""
|
||||
if self.endpoints:
|
||||
return any(ep.resolved_api_key() for ep in self.endpoints)
|
||||
return bool(self.resolved_api_key())
|
||||
|
||||
|
||||
class ModelDefaults(BaseModel):
|
||||
"""Global default parameters for all models."""
|
||||
|
||||
temperature: float | None = Field(None, description="Sampling temperature")
|
||||
max_tokens: int = Field(4096, description="Default max output tokens")
|
||||
reasoning_effort: str | None = Field("high", description="Reasoning effort level (low/medium/high/xhigh)")
|
||||
stream_usage: bool = Field(True, description="Whether to enable streaming token usage stats")
|
||||
|
||||
|
||||
class StructuredConfig(BaseModel):
|
||||
"""Top-level structured configuration for EvoScientist LLM settings."""
|
||||
|
||||
default_model: str = Field("", description="Default model ID or alias")
|
||||
providers: dict[str, ProviderConfig] = Field(
|
||||
default_factory=dict,
|
||||
description="Provider definitions keyed by name",
|
||||
)
|
||||
model_defaults: ModelDefaults = Field(
|
||||
default_factory=ModelDefaults,
|
||||
description="Global default model parameters",
|
||||
)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Config loading with layered discovery
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _get_global_settings_path() -> Path:
|
||||
"""Get global settings.yaml path via get_config_path()."""
|
||||
from EvoScientist.config.settings import get_config_path
|
||||
|
||||
return get_config_path()
|
||||
|
||||
|
||||
def _load_yaml_file(path: Path) -> dict:
|
||||
"""Load a YAML file, returning empty dict on failure."""
|
||||
if not path.is_file():
|
||||
return {}
|
||||
try:
|
||||
with open(path) as f:
|
||||
data = yaml.safe_load(f)
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _detect_new_config_format(data: dict) -> bool:
|
||||
"""Detect whether a YAML dict uses the new structured format.
|
||||
|
||||
The new format is identified by the presence of a 'providers' top-level key
|
||||
that is a dict (not a flat string value).
|
||||
"""
|
||||
providers = data.get("providers")
|
||||
return isinstance(providers, dict) and len(providers) > 0
|
||||
|
||||
|
||||
def _deep_merge(base: dict, override: dict) -> dict:
|
||||
"""Deep merge override into base. Override values take precedence."""
|
||||
result = base.copy()
|
||||
for key, value in override.items():
|
||||
if key in result and isinstance(result[key], dict) and isinstance(value, dict):
|
||||
result[key] = _deep_merge(result[key], value)
|
||||
elif key in result and isinstance(result[key], list) and isinstance(value, list):
|
||||
# For lists (like models), override completely replaces
|
||||
result[key] = value
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
def load_structured_config(
|
||||
cli_overrides: dict[str, Any] | None = None,
|
||||
) -> StructuredConfig:
|
||||
"""Load structured config from ``settings.yaml`` (single source of truth).
|
||||
|
||||
Returns code defaults if the file doesn't exist or is invalid.
|
||||
|
||||
Args:
|
||||
cli_overrides: Optional CLI argument overrides.
|
||||
|
||||
Returns:
|
||||
StructuredConfig instance.
|
||||
"""
|
||||
settings_path = _get_global_settings_path()
|
||||
settings_data = _load_yaml_file(settings_path)
|
||||
config = StructuredConfig()
|
||||
if _detect_new_config_format(settings_data):
|
||||
config = StructuredConfig(**settings_data)
|
||||
_remove_image_generation_models(config)
|
||||
|
||||
# Apply CLI overrides
|
||||
if cli_overrides:
|
||||
if "default_model" in cli_overrides and cli_overrides["default_model"]:
|
||||
config.default_model = cli_overrides["default_model"]
|
||||
|
||||
return config
|
||||
|
||||
|
||||
def _remove_image_generation_models(config: StructuredConfig) -> list[tuple[str, ProviderConfig, ModelEntry]]:
|
||||
"""Remove image-only models from chat LLM providers and return them."""
|
||||
from EvoScientist.config.image_models import is_image_generation_model
|
||||
|
||||
removed: list[tuple[str, ProviderConfig, ModelEntry]] = []
|
||||
for prov_name, provider in config.providers.items():
|
||||
chat_models: list[ModelEntry] = []
|
||||
for model in provider.models:
|
||||
if is_image_generation_model(model.id) or is_image_generation_model(model.alias):
|
||||
removed.append((prov_name, provider, model))
|
||||
else:
|
||||
chat_models.append(model)
|
||||
provider.models = chat_models
|
||||
if config.default_model and is_image_generation_model(config.default_model):
|
||||
config.default_model = ""
|
||||
return removed
|
||||
|
||||
|
||||
def move_image_models_to_dedicated_config(config: StructuredConfig) -> StructuredConfig:
|
||||
"""Move image-only model entries out of chat LLM config."""
|
||||
removed = _remove_image_generation_models(config)
|
||||
if not removed:
|
||||
return config
|
||||
|
||||
from EvoScientist.config.image_models import ImageModelEntry, merge_image_generation_models
|
||||
|
||||
entries: list[ImageModelEntry] = []
|
||||
for prov_name, provider, model in removed:
|
||||
api_key = provider.api_key
|
||||
base_url = provider.base_url
|
||||
if model.endpoint:
|
||||
for endpoint in provider.endpoints:
|
||||
if endpoint.name == model.endpoint:
|
||||
api_key = endpoint.api_key or api_key
|
||||
base_url = endpoint.base_url or base_url
|
||||
break
|
||||
entries.append(
|
||||
ImageModelEntry(
|
||||
id=model.id,
|
||||
name=model.alias or model.id,
|
||||
provider=prov_name,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
supports_generation=True,
|
||||
supports_edit=True,
|
||||
default_size=str(model.params.get("size") or "1024x1024"),
|
||||
default_quality=str(model.params.get("quality") or "auto"),
|
||||
params=model.params,
|
||||
)
|
||||
)
|
||||
merge_image_generation_models(entries, default_model=entries[0].id if entries else None)
|
||||
return config
|
||||
|
||||
|
||||
def save_structured_config(config: StructuredConfig | dict[str, Any]) -> StructuredConfig:
|
||||
"""Validate and atomically persist the global structured config."""
|
||||
validated = config if isinstance(config, StructuredConfig) else StructuredConfig.model_validate(config)
|
||||
validated = move_image_models_to_dedicated_config(validated)
|
||||
path = _get_global_settings_path()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
existing = _load_yaml_file(path)
|
||||
output = existing if isinstance(existing, dict) else {}
|
||||
for key in _FLAT_LLM_CONFIG_KEYS:
|
||||
output.pop(key, None)
|
||||
output.update(validated.model_dump(mode="python"))
|
||||
|
||||
yaml_text = yaml.safe_dump(
|
||||
output,
|
||||
allow_unicode=False,
|
||||
default_flow_style=False,
|
||||
sort_keys=False,
|
||||
)
|
||||
|
||||
fd, tmp_path = tempfile.mkstemp(
|
||||
dir=str(path.parent),
|
||||
prefix=".llm_config_",
|
||||
suffix=".yaml.tmp",
|
||||
text=True,
|
||||
)
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
f.write(yaml_text)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp_path, path)
|
||||
finally:
|
||||
if os.path.exists(tmp_path):
|
||||
os.unlink(tmp_path)
|
||||
|
||||
return validated
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Model registry: lookup and resolution
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class ResolvedModel(BaseModel):
|
||||
"""Fully resolved model information ready for get_chat_model()."""
|
||||
|
||||
provider_name: str = Field(..., description="Provider name (e.g. 'anthropic', 'deepseek')")
|
||||
model_id: str = Field(..., description="Full model ID for the API call")
|
||||
protocol: str = Field(..., description="Protocol to use (openai/anthropic/google-genai/ollama)")
|
||||
api_key: str = Field("", description="Resolved API key")
|
||||
base_url: str = Field("", description="Resolved base URL")
|
||||
endpoint_name: str = Field("", description="Name of the selected endpoint (empty for single-endpoint providers)")
|
||||
params: dict[str, Any] = Field(default_factory=dict, description="Merged parameters")
|
||||
supports_vision: bool = False
|
||||
supports_reasoning: bool = False
|
||||
max_tokens: int = 4096
|
||||
|
||||
|
||||
def _build_alias_index(config: StructuredConfig) -> dict[str, tuple[str, str]]:
|
||||
"""Build alias → (provider_name, model_id) index.
|
||||
|
||||
Also indexes model IDs directly so both alias and full ID can be looked up.
|
||||
"""
|
||||
index: dict[str, tuple[str, str]] = {}
|
||||
for prov_name, prov in config.providers.items():
|
||||
for model in prov.models:
|
||||
# Index by alias (if set)
|
||||
if model.alias:
|
||||
index[model.alias] = (prov_name, model.id)
|
||||
# Index by full model ID
|
||||
index[model.id] = (prov_name, model.id)
|
||||
return index
|
||||
|
||||
|
||||
def resolve_model(
|
||||
model_ref: str,
|
||||
config: StructuredConfig | None = None,
|
||||
runtime_params: dict[str, Any] | None = None,
|
||||
) -> ResolvedModel:
|
||||
"""Resolve a model reference to a fully configured ResolvedModel.
|
||||
|
||||
Args:
|
||||
model_ref: Model ID, alias, or provider-prefixed name (e.g. "anthropic/claude-sonnet-4-6").
|
||||
config: StructuredConfig to use (loaded automatically if None).
|
||||
runtime_params: Additional runtime parameter overrides.
|
||||
|
||||
Returns:
|
||||
ResolvedModel with all fields populated.
|
||||
|
||||
Raises:
|
||||
ValueError: If the model cannot be resolved.
|
||||
"""
|
||||
if config is None:
|
||||
config = load_structured_config()
|
||||
|
||||
runtime_params = runtime_params or {}
|
||||
|
||||
# Handle provider-prefixed references: "anthropic/claude-sonnet-4-6"
|
||||
explicit_provider = None
|
||||
if "/" in model_ref:
|
||||
parts = model_ref.split("/", 1)
|
||||
explicit_provider = parts[0]
|
||||
model_ref = parts[1]
|
||||
|
||||
# Try alias/ID lookup from structured configuration.
|
||||
alias_index = _build_alias_index(config)
|
||||
if model_ref in alias_index:
|
||||
prov_name, model_id = alias_index[model_ref]
|
||||
if explicit_provider and prov_name != explicit_provider:
|
||||
prov_name = explicit_provider
|
||||
model_id = model_ref
|
||||
elif explicit_provider:
|
||||
# Not in registry but user specified provider — use as-is
|
||||
prov_name = explicit_provider
|
||||
model_id = model_ref
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Model '{model_ref}' is not declared in structured LLM config. "
|
||||
"Add it under providers.*.models or call it as 'provider/model'."
|
||||
)
|
||||
|
||||
# Get provider config
|
||||
provider = config.providers.get(prov_name)
|
||||
if not provider:
|
||||
raise ValueError(
|
||||
f"Provider '{prov_name}' is not declared in structured LLM config."
|
||||
)
|
||||
|
||||
# Find the model entry (if registered)
|
||||
model_entry = None
|
||||
for m in provider.models:
|
||||
if m.id == model_id or m.alias == model_ref:
|
||||
model_entry = m
|
||||
break
|
||||
|
||||
# Three-level parameter merge: defaults ← model ← runtime
|
||||
params: dict[str, Any] = {}
|
||||
defaults = config.model_defaults
|
||||
|
||||
# Level 1: global defaults
|
||||
if defaults.temperature is not None:
|
||||
params["temperature"] = defaults.temperature
|
||||
params["max_tokens"] = defaults.max_tokens
|
||||
params["stream_usage"] = defaults.stream_usage
|
||||
|
||||
# Level 2: provider-level params
|
||||
params.update(provider.params)
|
||||
|
||||
# Level 3: model-level params
|
||||
if model_entry:
|
||||
params.update(model_entry.params)
|
||||
params["max_tokens"] = model_entry.max_tokens
|
||||
# Override defaults with model-level values
|
||||
if "temperature" not in model_entry.params and defaults.temperature is not None:
|
||||
params["temperature"] = defaults.temperature
|
||||
|
||||
# ── Resolve credentials from provider or endpoint ────────────
|
||||
resolved_api_key = provider.resolved_api_key()
|
||||
resolved_base_url = _resolve_env_ref(provider.base_url) if provider.base_url else ""
|
||||
resolved_extra_body = dict(provider.extra_body) if provider.extra_body else {}
|
||||
resolved_headers = dict(provider.default_headers) if provider.default_headers else {}
|
||||
selected_endpoint_name = ""
|
||||
|
||||
if provider.endpoints:
|
||||
# Multi-endpoint mode: select endpoint via load balancing
|
||||
preferred_ep = model_entry.endpoint if model_entry else ""
|
||||
available_endpoints = [ep for ep in provider.endpoints if ep.resolved_api_key()]
|
||||
if available_endpoints:
|
||||
selected = _select_endpoint(available_endpoints, prov_name, preferred_ep)
|
||||
resolved_api_key = selected.resolved_api_key()
|
||||
resolved_base_url = selected.resolved_base_url() or resolved_base_url
|
||||
selected_endpoint_name = selected.name
|
||||
# Merge endpoint-level extra_body and headers (provider-level first, endpoint overrides)
|
||||
if selected.extra_body:
|
||||
resolved_extra_body.update(selected.extra_body)
|
||||
if selected.default_headers:
|
||||
resolved_headers.update(selected.default_headers)
|
||||
if selected.params:
|
||||
params.update(selected.params)
|
||||
# Record endpoint call statistics
|
||||
_endpoint_stats.record(prov_name, selected.name or "default", model_id)
|
||||
logger.debug(
|
||||
"endpoint selected: provider=%s endpoint=%s model=%s",
|
||||
prov_name, selected.name or "default", model_id,
|
||||
)
|
||||
else:
|
||||
# Single-endpoint provider — still record for token attribution
|
||||
# so that usage_stats events carry correct endpoint info
|
||||
_endpoint_stats.record(prov_name, "default", model_id)
|
||||
|
||||
# Level 4: runtime params
|
||||
params.update(runtime_params)
|
||||
|
||||
# Store merged extra_body/headers in params for downstream consumption
|
||||
if resolved_extra_body:
|
||||
params["_extra_body"] = resolved_extra_body
|
||||
if resolved_headers:
|
||||
params["_default_headers"] = resolved_headers
|
||||
|
||||
return ResolvedModel(
|
||||
provider_name=prov_name,
|
||||
model_id=model_id,
|
||||
protocol=provider.protocol,
|
||||
api_key=resolved_api_key,
|
||||
base_url=resolved_base_url,
|
||||
endpoint_name=selected_endpoint_name,
|
||||
params=params,
|
||||
supports_vision=model_entry.supports_vision if model_entry else False,
|
||||
supports_reasoning=model_entry.supports_reasoning if model_entry else False,
|
||||
max_tokens=model_entry.max_tokens if model_entry else defaults.max_tokens,
|
||||
)
|
||||
|
||||
|
||||
def get_default_model(config: StructuredConfig | None = None) -> str:
|
||||
"""Get the default model reference from config."""
|
||||
if config is None:
|
||||
config = load_structured_config()
|
||||
return config.default_model
|
||||
|
||||
|
||||
def list_available_models(config: StructuredConfig | None = None) -> list[dict[str, Any]]:
|
||||
"""List all available models from configured providers.
|
||||
|
||||
Returns a list of dicts with keys: id, alias, provider, protocol,
|
||||
max_tokens, supports_vision, supports_reasoning.
|
||||
"""
|
||||
if config is None:
|
||||
config = load_structured_config()
|
||||
|
||||
models = []
|
||||
for prov_name, provider in config.providers.items():
|
||||
# Only include models from providers that have credentials
|
||||
# (either via top-level api_key or via at least one endpoint)
|
||||
has_creds = provider.has_credentials() or (
|
||||
prov_name == "ollama" and bool(provider.base_url)
|
||||
)
|
||||
if not has_creds:
|
||||
continue
|
||||
|
||||
for model in provider.models:
|
||||
models.append({
|
||||
"id": model.id,
|
||||
"alias": model.alias,
|
||||
"provider": prov_name,
|
||||
"protocol": provider.protocol,
|
||||
"max_tokens": model.max_tokens,
|
||||
"supports_vision": model.supports_vision,
|
||||
"supports_reasoning": model.supports_reasoning,
|
||||
"provider_access": provider.access.model_dump(mode="python"),
|
||||
"permission_mode": model.permission_mode,
|
||||
"model_access": model.access.model_dump(mode="python"),
|
||||
})
|
||||
return models
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Onboarding package.
|
||||
|
||||
The wizard's only package-level public entry point is :func:`run_onboard`.
|
||||
Everything else lives in submodules — import directly from them:
|
||||
|
||||
- :mod:`EvoScientist.config.onboard.wizard` — orchestrator, ``run_onboard``,
|
||||
``STEPS``, ``render_progress``
|
||||
- :mod:`EvoScientist.config.onboard.steps` — per-step functions
|
||||
- :mod:`EvoScientist.config.onboard.channels` — channel selection + setup
|
||||
- :mod:`EvoScientist.config.onboard.helpers` — API-key prompt,
|
||||
npx/node, LaTeX, iMessage helpers
|
||||
- :mod:`EvoScientist.config.onboard.style` — Rich styles + ``_checkbox_ask``
|
||||
- :mod:`EvoScientist.config.onboard.validators` — input validators
|
||||
- :mod:`EvoScientist.config.onboard.prompter` — ``NonInteractivePrompter``
|
||||
(CLI-answer container) + ``select_navigation_active`` for keyboard nav
|
||||
- :mod:`EvoScientist.config.onboard.constants` — canonical valid-value sets
|
||||
|
||||
This module used to re-export every symbol from every submodule for
|
||||
backward compat during the initial refactor; those re-exports have since
|
||||
been removed to keep the public surface narrow. New code should always
|
||||
import from the submodule that owns the symbol; test code should use
|
||||
``patch("EvoScientist.config.onboard.<submodule>.<name>")`` paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# Sole package-level public entry. ``EvoScientist.config`` re-exports this
|
||||
# (via lazy ``__getattr__``) so ``from EvoScientist.config import
|
||||
# run_onboard`` keeps working — that import path is used by the CLI and is
|
||||
# the only documented external API.
|
||||
from .wizard import run_onboard
|
||||
|
||||
__all__ = ["run_onboard"]
|
||||
@@ -0,0 +1,958 @@
|
||||
"""Channel selection + per-channel configuration.
|
||||
|
||||
`_step_channels` is the big one — over 700 lines that walk the user through
|
||||
selecting which messaging channels to enable and collecting credentials for
|
||||
each.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import questionary
|
||||
from questionary import Choice
|
||||
|
||||
from ..settings import EvoScientistConfig
|
||||
from .helpers import (
|
||||
_setup_imessage,
|
||||
)
|
||||
from .style import (
|
||||
QMARK,
|
||||
WIZARD_STYLE,
|
||||
console,
|
||||
)
|
||||
|
||||
|
||||
def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
||||
"""Step: Select channels to enable on startup.
|
||||
|
||||
Presents a multi-select list of supported channels.
|
||||
For each selected channel, prompts for required credentials
|
||||
and validates them via the channel's probe function.
|
||||
|
||||
Args:
|
||||
config: Current configuration.
|
||||
|
||||
Returns:
|
||||
Dict mapping config field names to their new values.
|
||||
Empty dict when the user skips or selects nothing.
|
||||
"""
|
||||
# Currently enabled channels
|
||||
_currently_enabled = {
|
||||
t.strip()
|
||||
for t in (getattr(config, "channel_enabled", "") or "").split(",")
|
||||
if t.strip()
|
||||
}
|
||||
# Legacy iMessage compat
|
||||
if (
|
||||
getattr(config, "imessage_enabled", False)
|
||||
and "imessage" not in _currently_enabled
|
||||
):
|
||||
_currently_enabled.add("imessage")
|
||||
|
||||
# Direct pip packages for each channel extra. Used to install the
|
||||
# exact dependency without requiring the evoscientist package itself
|
||||
# to be resolvable on PyPI (e.g. editable / dev installs).
|
||||
_CHANNEL_PIP_DEPS: dict[str, list[str]] = {
|
||||
"telegram": ["python-telegram-bot>=21.0"],
|
||||
"discord": ["discord.py>=2.3"],
|
||||
"slack": ["slack-sdk>=3.27", "aiohttp>=3.9"],
|
||||
"feishu": ["aiohttp>=3.9", "qrcode>=7.4"],
|
||||
"dingtalk": ["aiohttp>=3.9"],
|
||||
"wechat": [
|
||||
"pycryptodome>=3.20",
|
||||
"aiohttp>=3.9",
|
||||
"qrcode>=7.4",
|
||||
"certifi>=2024.0",
|
||||
],
|
||||
"qq": ["qq-botpy>=1.0", "cryptography>=41.0", "qrcode>=7.4"],
|
||||
}
|
||||
|
||||
# Channel definitions:
|
||||
# (value, display_name, required_fields, import_check, pip_extra)
|
||||
# required_fields entries are (field_name, prompt_label, is_secret).
|
||||
# ``is_secret=True`` triggers a password prompt (no echo, no default echo)
|
||||
# so bot tokens / OAuth secrets / IMAP+SMTP passwords don't leak into
|
||||
# terminal scrollback, screen recordings, or support sessions.
|
||||
_CHANNELS = [
|
||||
(
|
||||
"telegram",
|
||||
"Telegram",
|
||||
[("telegram_bot_token", "Bot token (from @BotFather)", True)],
|
||||
"telegram",
|
||||
"telegram",
|
||||
),
|
||||
(
|
||||
"discord",
|
||||
"Discord",
|
||||
[("discord_bot_token", "Bot token", True)],
|
||||
"discord",
|
||||
"discord",
|
||||
),
|
||||
(
|
||||
"slack",
|
||||
"Slack",
|
||||
[
|
||||
("slack_bot_token", "Bot token (xoxb-...)", True),
|
||||
("slack_app_token", "App token for Socket Mode (xapp-...)", True),
|
||||
],
|
||||
"slack_sdk",
|
||||
"slack",
|
||||
),
|
||||
(
|
||||
"feishu",
|
||||
"Feishu",
|
||||
[
|
||||
("feishu_app_id", "App ID", False),
|
||||
("feishu_app_secret", "App Secret", True),
|
||||
],
|
||||
"aiohttp",
|
||||
"feishu",
|
||||
),
|
||||
(
|
||||
"dingtalk",
|
||||
"DingTalk",
|
||||
[
|
||||
("dingtalk_client_id", "Client ID (AppKey)", False),
|
||||
("dingtalk_client_secret", "Client Secret (AppSecret)", True),
|
||||
],
|
||||
"aiohttp",
|
||||
"dingtalk",
|
||||
),
|
||||
(
|
||||
"wechat",
|
||||
"WeChat",
|
||||
[], # backend-specific fields prompted in the wechat branch below
|
||||
("aiohttp", "qrcode", "Crypto", "certifi"),
|
||||
"wechat",
|
||||
),
|
||||
(
|
||||
"email",
|
||||
"Email",
|
||||
[
|
||||
("email_imap_host", "IMAP host", False),
|
||||
("email_imap_username", "IMAP username", False),
|
||||
("email_imap_password", "IMAP password", True),
|
||||
("email_smtp_host", "SMTP host", False),
|
||||
("email_smtp_username", "SMTP username", False),
|
||||
("email_smtp_password", "SMTP password", True),
|
||||
("email_from_address", "From address", False),
|
||||
],
|
||||
None,
|
||||
None,
|
||||
),
|
||||
(
|
||||
"qq",
|
||||
"QQ",
|
||||
[
|
||||
("qq_app_id", "App ID", False),
|
||||
("qq_app_secret", "App Secret", True),
|
||||
],
|
||||
"botpy",
|
||||
"qq",
|
||||
),
|
||||
(
|
||||
"signal",
|
||||
"Signal",
|
||||
[("signal_phone_number", "Phone number (E.164)", False)],
|
||||
None,
|
||||
None,
|
||||
),
|
||||
("imessage", "iMessage", [], None, None), # handled via _setup_imessage()
|
||||
]
|
||||
|
||||
choices = [
|
||||
Choice(
|
||||
title=display,
|
||||
value=value,
|
||||
checked=value in _currently_enabled,
|
||||
)
|
||||
for value, display, *_ in _CHANNELS
|
||||
]
|
||||
|
||||
selected = questionary.checkbox(
|
||||
"Select channels to enable (Space to toggle, Enter to confirm):",
|
||||
choices=choices,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=QMARK,
|
||||
).ask()
|
||||
|
||||
if selected is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
updates: dict[str, object] = {}
|
||||
|
||||
if not selected:
|
||||
updates["channel_enabled"] = ""
|
||||
updates["imessage_enabled"] = False
|
||||
return updates
|
||||
|
||||
from ...mcp.registry import install_library, pip_install_hint
|
||||
|
||||
# Build a lookup for channel definitions
|
||||
_ch_lookup = {
|
||||
v: (v, d, fields, imp, extra) for v, d, fields, imp, extra in _CHANNELS
|
||||
}
|
||||
|
||||
enabled_channels: list[str] = []
|
||||
|
||||
for ch_name in selected:
|
||||
_, display, required_fields, import_check, pip_extra = _ch_lookup[ch_name]
|
||||
console.print(f"\n [bold cyan]── {display} ──[/bold cyan]")
|
||||
|
||||
# Check pip dependency before proceeding
|
||||
if import_check:
|
||||
_required_imports: tuple[str, ...] = (
|
||||
(import_check,)
|
||||
if isinstance(import_check, str)
|
||||
else tuple(import_check)
|
||||
)
|
||||
_pkg_ready = False
|
||||
try:
|
||||
for _module_name in _required_imports:
|
||||
__import__(_module_name)
|
||||
_pkg_ready = True
|
||||
except ImportError:
|
||||
console.print(" [yellow]✗ Required package not installed.[/yellow]")
|
||||
# Determine packages to install
|
||||
_pip_pkgs = _CHANNEL_PIP_DEPS.get(pip_extra, []) if pip_extra else []
|
||||
_pkg_display = (
|
||||
" ".join(f'"{p}"' for p in _pip_pkgs)
|
||||
if _pip_pkgs
|
||||
else f'"evoscientist[{pip_extra}]"'
|
||||
)
|
||||
install_now = questionary.confirm(
|
||||
f"Install {_pkg_display} now?",
|
||||
default=True,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if install_now is None:
|
||||
raise KeyboardInterrupt() from None
|
||||
if install_now:
|
||||
console.print(f" [dim]Installing {_pkg_display}...[/dim]")
|
||||
if _pip_pkgs:
|
||||
_ok = all(install_library(p) for p in _pip_pkgs)
|
||||
else:
|
||||
_ok = install_library(f"evoscientist[{pip_extra}]")
|
||||
if _ok:
|
||||
# Verify the imports actually work now
|
||||
try:
|
||||
for _module_name in _required_imports:
|
||||
__import__(_module_name)
|
||||
console.print(" [green]✓ Installed successfully.[/green]")
|
||||
_pkg_ready = True
|
||||
except ImportError:
|
||||
console.print(
|
||||
" [red]✗ Package installed but import failed.[/red]"
|
||||
)
|
||||
console.print(
|
||||
" [dim]Try restarting and running:[/dim] evosci channel setup"
|
||||
)
|
||||
else:
|
||||
console.print(" [red]✗ Installation failed.[/red]")
|
||||
console.print(
|
||||
f" [dim]Run manually:[/dim] {pip_install_hint()} {_pkg_display}"
|
||||
)
|
||||
if not _pkg_ready:
|
||||
# Previously-enabled channels are silently dropped from
|
||||
# ``channel_enabled`` if we just ``continue`` — warn.
|
||||
if ch_name in _currently_enabled:
|
||||
console.print(
|
||||
f" [bold yellow]⚠ {display} will be DISABLED[/bold yellow]"
|
||||
" [dim](dependency missing — re-run after install)[/dim]"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
f" [dim]Skipping {display} — dependency not installed.[/dim]"
|
||||
)
|
||||
continue
|
||||
|
||||
# Special handling for iMessage
|
||||
if ch_name == "imessage":
|
||||
ready = _setup_imessage()
|
||||
if not ready:
|
||||
console.print()
|
||||
enable_anyway = questionary.confirm(
|
||||
"Enable iMessage anyway? (will try to connect on startup)",
|
||||
default=False,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if enable_anyway is None:
|
||||
raise KeyboardInterrupt()
|
||||
if not enable_anyway:
|
||||
continue
|
||||
# Allowed senders
|
||||
senders = questionary.text(
|
||||
"Allowed senders (comma-separated, empty = all):",
|
||||
default=getattr(config, "imessage_allowed_senders", ""),
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if senders is None:
|
||||
raise KeyboardInterrupt()
|
||||
updates["imessage_enabled"] = True
|
||||
updates["imessage_allowed_senders"] = senders.strip()
|
||||
enabled_channels.append("imessage")
|
||||
continue
|
||||
|
||||
# QQ: offer scan-to-configure before falling back to manual entry.
|
||||
# The bot must already exist at q.qq.com — scanning binds the
|
||||
# developer's QQ account to it and returns app_id + client_secret.
|
||||
_qq_scanned = False
|
||||
_feishu_scanned = False
|
||||
if ch_name == "qq":
|
||||
scan_choices = [
|
||||
Choice(
|
||||
title="Scan QR code (recommended — auto-fill App ID & Secret)",
|
||||
value="scan",
|
||||
),
|
||||
Choice(title="Enter App ID and Secret manually", value="manual"),
|
||||
]
|
||||
scan_choice = questionary.select(
|
||||
"Configure QQ Bot:",
|
||||
choices=scan_choices,
|
||||
default="scan",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if scan_choice is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
if scan_choice == "scan":
|
||||
# Preflight: AES-GCM decryption needs `cryptography`.
|
||||
# `qrcode` is a soft dep — onboard.py degrades to URL-only display.
|
||||
try:
|
||||
import cryptography # noqa: F401
|
||||
except ImportError:
|
||||
console.print(
|
||||
' [yellow]✗ QR scan requires "cryptography".[/yellow]'
|
||||
)
|
||||
install_now = questionary.confirm(
|
||||
'Install "cryptography" now?',
|
||||
default=True,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if install_now is None:
|
||||
raise KeyboardInterrupt() from None
|
||||
if install_now and install_library("cryptography>=41.0"):
|
||||
console.print(" [green]✓ Installed cryptography.[/green]")
|
||||
else:
|
||||
console.print(
|
||||
" [yellow]⚠ Falling back to manual entry.[/yellow]"
|
||||
)
|
||||
scan_choice = "manual"
|
||||
|
||||
if scan_choice == "scan":
|
||||
from ...channels.qq.onboard import qr_register
|
||||
|
||||
console.print(
|
||||
" [dim]Make sure the bot is registered at"
|
||||
" https://q.qq.com first — scanning binds an"
|
||||
" existing app, it does not create one.[/dim]"
|
||||
)
|
||||
try:
|
||||
creds = qr_register()
|
||||
except Exception as exc:
|
||||
console.print(f" [red]✗ Scan failed: {exc}[/red]")
|
||||
creds = None
|
||||
|
||||
if creds:
|
||||
updates["qq_app_id"] = creds["app_id"]
|
||||
updates["qq_app_secret"] = creds["client_secret"]
|
||||
console.print(
|
||||
f" [green]✓ Bound QQ Bot (App ID: {creds['app_id']})[/green]"
|
||||
)
|
||||
_qq_scanned = True
|
||||
else:
|
||||
console.print(
|
||||
" [yellow]⚠ Scan did not complete — falling"
|
||||
" back to manual entry.[/yellow]"
|
||||
)
|
||||
|
||||
# Feishu: offer scan-to-create before falling back to manual entry.
|
||||
# Unlike QQ, this provisions a brand-new PersonalAgent app with the
|
||||
# required IM permissions attached, then returns app_id + app_secret.
|
||||
if ch_name == "feishu":
|
||||
scan_choices = [
|
||||
Choice(
|
||||
title="Scan QR code (recommended — auto-create app, fill App ID & Secret)",
|
||||
value="scan",
|
||||
),
|
||||
Choice(title="Enter App ID and Secret manually", value="manual"),
|
||||
]
|
||||
scan_choice = questionary.select(
|
||||
"Configure Feishu / Lark:",
|
||||
choices=scan_choices,
|
||||
default="scan",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if scan_choice is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
if scan_choice == "scan":
|
||||
# `qrcode` is the only soft dep needed — onboard prints the URL
|
||||
# if it's missing, but the UX is much worse, so offer to install.
|
||||
try:
|
||||
import qrcode # noqa: F401
|
||||
except ImportError:
|
||||
console.print(
|
||||
' [yellow]✗ QR scan looks best with "qrcode".[/yellow]'
|
||||
)
|
||||
install_now = questionary.confirm(
|
||||
'Install "qrcode" now?',
|
||||
default=True,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if install_now is None:
|
||||
raise KeyboardInterrupt() from None
|
||||
if install_now and install_library("qrcode>=7.4"):
|
||||
console.print(" [green]✓ Installed qrcode.[/green]")
|
||||
else:
|
||||
console.print(
|
||||
" [yellow]⚠ Falling back to manual entry.[/yellow]"
|
||||
)
|
||||
scan_choice = "manual"
|
||||
|
||||
if scan_choice == "scan":
|
||||
# Region selection — accounts.feishu.cn vs accounts.larksuite.com.
|
||||
# The poll endpoint auto-switches if the scanning user is on the
|
||||
# other tenant, so this is just a starting hint.
|
||||
region_choices = [
|
||||
Choice(title="Feishu (飞书, mainland China)", value="feishu"),
|
||||
Choice(title="Lark (overseas)", value="lark"),
|
||||
]
|
||||
region = questionary.select(
|
||||
"Region:",
|
||||
choices=region_choices,
|
||||
default="feishu",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if region is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
if scan_choice == "scan":
|
||||
from ...channels.feishu.onboard import qr_register
|
||||
|
||||
console.print(
|
||||
" [dim]A QR code will be printed below — open Feishu or"
|
||||
" Lark on your phone and scan it. The platform will"
|
||||
" auto-create a bot app with IM permissions and return"
|
||||
" the credentials here.[/dim]"
|
||||
)
|
||||
try:
|
||||
creds = qr_register(initial_domain=region)
|
||||
except Exception as exc:
|
||||
console.print(f" [red]✗ Scan failed: {exc}[/red]")
|
||||
creds = None
|
||||
|
||||
if creds:
|
||||
updates["feishu_app_id"] = creds["app_id"]
|
||||
updates["feishu_app_secret"] = creds["app_secret"]
|
||||
# Sync open-platform domain to the resolved region
|
||||
updates["feishu_domain"] = (
|
||||
"https://open.larksuite.com"
|
||||
if creds.get("domain") == "lark"
|
||||
else "https://open.feishu.cn"
|
||||
)
|
||||
bot_name = creds.get("bot_name")
|
||||
if bot_name:
|
||||
console.print(
|
||||
f' [green]✓ Bound Feishu bot "{bot_name}"'
|
||||
f" (App ID: {creds['app_id']})[/green]"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
f" [green]✓ Bound Feishu app"
|
||||
f" (App ID: {creds['app_id']})[/green]"
|
||||
)
|
||||
_feishu_scanned = True
|
||||
else:
|
||||
console.print(
|
||||
" [yellow]⚠ Scan did not complete — falling"
|
||||
" back to manual entry.[/yellow]"
|
||||
)
|
||||
|
||||
# WeChat: pick backend (wecom / wechatmp / personal), then prompt
|
||||
# backend-specific fields. Personal-WeChat has no static credentials —
|
||||
# we offer an interactive QR-scan that obtains and persists them.
|
||||
if ch_name == "wechat":
|
||||
backend_choices = [
|
||||
Choice(
|
||||
title="WeCom (企业微信应用) — most stable, official API",
|
||||
value="wecom",
|
||||
),
|
||||
Choice(
|
||||
title="Official Account (微信公众号) — public-facing bots",
|
||||
value="wechatmp",
|
||||
),
|
||||
Choice(
|
||||
title="Personal WeChat (个人微信, iLink) — QR-code scan login",
|
||||
value="personal",
|
||||
),
|
||||
]
|
||||
wechat_backend = questionary.select(
|
||||
"WeChat backend:",
|
||||
choices=backend_choices,
|
||||
default=getattr(config, "wechat_backend", "") or "wecom",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if wechat_backend is None:
|
||||
raise KeyboardInterrupt()
|
||||
updates["wechat_backend"] = wechat_backend
|
||||
|
||||
# Both WeCom and WeChat MP need the same non-empty-required
|
||||
# treatment as the generic required_fields loop below — newly
|
||||
# enabling either with blank credentials would leave the channel
|
||||
# half-configured and only fail at first message.
|
||||
wechat_newly_enabled = "wechat" not in _currently_enabled
|
||||
wechat_fields_for_backend: list[tuple[str, str, bool]] = []
|
||||
if wechat_backend == "wecom":
|
||||
wechat_fields_for_backend = [
|
||||
("wechat_wecom_corp_id", "WeCom Corp ID", False),
|
||||
("wechat_wecom_agent_id", "WeCom Agent ID", False),
|
||||
("wechat_wecom_secret", "WeCom Secret", True),
|
||||
]
|
||||
elif wechat_backend == "wechatmp":
|
||||
wechat_fields_for_backend = [
|
||||
("wechat_mp_app_id", "Official Account App ID", False),
|
||||
("wechat_mp_app_secret", "Official Account App Secret", True),
|
||||
]
|
||||
|
||||
if wechat_backend in ("wecom", "wechatmp"):
|
||||
for field_name, prompt_label, is_secret in wechat_fields_for_backend:
|
||||
current = getattr(config, field_name, "")
|
||||
while True:
|
||||
if is_secret:
|
||||
masked_hint = (
|
||||
f" (current: ***{current[-4:]})" if current else ""
|
||||
)
|
||||
value = questionary.password(
|
||||
f"{prompt_label}{masked_hint}:",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
else:
|
||||
value = questionary.text(
|
||||
f"{prompt_label}:",
|
||||
default=current,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if value is None:
|
||||
raise KeyboardInterrupt()
|
||||
value = value.strip()
|
||||
if not value and current:
|
||||
break # keep existing
|
||||
if not value and wechat_newly_enabled:
|
||||
console.print(
|
||||
f" [yellow]{prompt_label} is required to "
|
||||
"enable WeChat. Press Ctrl+C to cancel.[/yellow]"
|
||||
)
|
||||
continue
|
||||
updates[field_name] = value
|
||||
break
|
||||
elif wechat_backend == "personal":
|
||||
personal_choices = [
|
||||
Choice(
|
||||
title="Scan QR code now (recommended — login to a personal WeChat account)",
|
||||
value="scan",
|
||||
),
|
||||
Choice(
|
||||
title="I already have an account_id — enter it manually",
|
||||
value="manual",
|
||||
),
|
||||
]
|
||||
personal_choice = questionary.select(
|
||||
"Personal WeChat login:",
|
||||
choices=personal_choices,
|
||||
default="scan",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if personal_choice is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
if personal_choice == "scan":
|
||||
from ...channels.wechat.personal import _account_dir as _wp_dir
|
||||
|
||||
_accounts_path = _wp_dir()
|
||||
console.print(
|
||||
" [dim]A QR code will be printed below — open WeChat on"
|
||||
" your phone and scan it. The session token is saved"
|
||||
f" to {_accounts_path}.[/dim]"
|
||||
)
|
||||
try:
|
||||
import asyncio
|
||||
|
||||
from ...channels.wechat.personal import qr_login
|
||||
|
||||
creds = asyncio.run(qr_login())
|
||||
except Exception as exc:
|
||||
console.print(f" [red]✗ Scan failed: {exc}[/red]")
|
||||
creds = None
|
||||
|
||||
if creds:
|
||||
updates["wechat_personal_account_id"] = creds["account_id"]
|
||||
# Token is persisted on disk by qr_login(); the channel
|
||||
# reads it from the per-account store at runtime, so we
|
||||
# intentionally do NOT copy it into the main config here
|
||||
# (avoids stale duplicates and broader secret exposure).
|
||||
console.print(
|
||||
f" [green]✓ Logged in (account_id: "
|
||||
f"{creds['account_id'][:12]}…)[/green]"
|
||||
)
|
||||
else:
|
||||
console.print(
|
||||
" [yellow]⚠ QR login did not complete — falling"
|
||||
" back to manual entry.[/yellow]"
|
||||
)
|
||||
personal_choice = "manual"
|
||||
|
||||
if personal_choice == "manual":
|
||||
current_id = getattr(config, "wechat_personal_account_id", "")
|
||||
wechat_newly_enabled = "wechat" not in _currently_enabled
|
||||
while True:
|
||||
account_id = questionary.text(
|
||||
"iLink account_id (from a previous --qr-login run):",
|
||||
default=current_id,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if account_id is None:
|
||||
raise KeyboardInterrupt()
|
||||
account_id = account_id.strip()
|
||||
if not account_id and current_id:
|
||||
break # keep existing
|
||||
if not account_id and wechat_newly_enabled:
|
||||
console.print(
|
||||
" [yellow]account_id is required to enable "
|
||||
"WeChat Personal. Press Ctrl+C to cancel.[/yellow]"
|
||||
)
|
||||
continue
|
||||
updates["wechat_personal_account_id"] = account_id
|
||||
break
|
||||
|
||||
# Prompt for required fields. Secret fields use ``questionary.password``
|
||||
# so the entered value (and the existing one shown as a hint) are
|
||||
# never echoed to the terminal — see _CHANNELS docstring above.
|
||||
if not _qq_scanned and not _feishu_scanned:
|
||||
# "Newly enabled" = this channel wasn't in the user's prior
|
||||
# ``channel_enabled`` list. Required fields with no existing
|
||||
# value must be non-empty for newly enabled channels — saving
|
||||
# blanks leaves the channel half-configured and only surfaces
|
||||
# the problem on the first message.
|
||||
newly_enabled = ch_name not in _currently_enabled
|
||||
for field_name, prompt_label, is_secret in required_fields:
|
||||
current = getattr(config, field_name, "")
|
||||
while True:
|
||||
if is_secret:
|
||||
masked_hint = (
|
||||
f" (current: ***{current[-4:]})" if current else ""
|
||||
)
|
||||
value = questionary.password(
|
||||
f"{prompt_label}{masked_hint}:",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
else:
|
||||
value = questionary.text(
|
||||
f"{prompt_label}:",
|
||||
default=current,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if value is None:
|
||||
raise KeyboardInterrupt()
|
||||
value = value.strip()
|
||||
|
||||
# Empty input + existing value → keep existing (this is
|
||||
# the "re-run wizard, no change to this field" path).
|
||||
if not value and current:
|
||||
break
|
||||
# Empty input + newly enabling channel → not OK; the
|
||||
# channel would be enabled with broken creds. Re-prompt.
|
||||
if not value and newly_enabled:
|
||||
console.print(
|
||||
f" [yellow]{prompt_label} is required to enable "
|
||||
f"{display}. Press Ctrl+C to cancel instead.[/yellow]"
|
||||
)
|
||||
continue
|
||||
# Empty input + previously enabled but never set
|
||||
# (unlikely, but tolerate) → still allow blank-through
|
||||
# so the user isn't blocked re-running configure later.
|
||||
updates[field_name] = value
|
||||
break
|
||||
|
||||
# Feishu: subscription mode + optional fields
|
||||
if ch_name == "feishu":
|
||||
mode_choices = [
|
||||
Choice(
|
||||
title="Webhook (requires public IP / port forwarding)",
|
||||
value="webhook",
|
||||
),
|
||||
Choice(
|
||||
title="WebSocket long connection (no public IP needed)",
|
||||
value="websocket",
|
||||
),
|
||||
]
|
||||
sub_mode = questionary.select(
|
||||
"Subscription mode:",
|
||||
choices=mode_choices,
|
||||
default="webhook",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
if sub_mode is None:
|
||||
raise KeyboardInterrupt()
|
||||
updates["feishu_subscription_mode"] = sub_mode
|
||||
|
||||
if sub_mode == "websocket":
|
||||
# WebSocket mode needs lark-oapi SDK
|
||||
try:
|
||||
__import__("lark_oapi")
|
||||
except ImportError:
|
||||
console.print(
|
||||
' [yellow]✗ WebSocket mode requires "lark-oapi".[/yellow]'
|
||||
)
|
||||
install_sdk = questionary.confirm(
|
||||
'Install "lark-oapi>=1.4.0" now?',
|
||||
default=True,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if install_sdk is None:
|
||||
raise KeyboardInterrupt() from None
|
||||
if install_sdk:
|
||||
console.print(' [dim]Installing "lark-oapi"...[/dim]')
|
||||
if install_library("lark-oapi>=1.4.0"):
|
||||
console.print(" [green]✓ Installed successfully.[/green]")
|
||||
else:
|
||||
console.print(" [red]✗ Installation failed.[/red]")
|
||||
console.print(
|
||||
f" [dim]Run manually:[/dim] {pip_install_hint()} "
|
||||
'"lark-oapi>=1.4.0"'
|
||||
)
|
||||
else:
|
||||
# Webhook mode: prompt optional verification/encryption fields.
|
||||
# Both are credentials — use password() so they don't echo.
|
||||
console.print(
|
||||
" [dim]The following fields are optional"
|
||||
" (press Enter to skip):[/dim]"
|
||||
)
|
||||
for field_name, prompt_label in [
|
||||
("feishu_verification_token", "Verification Token (optional)"),
|
||||
("feishu_encrypt_key", "Encrypt Key (optional)"),
|
||||
]:
|
||||
current = getattr(config, field_name, "")
|
||||
masked_hint = f" (current: ***{current[-4:]})" if current else ""
|
||||
value = questionary.password(
|
||||
f"{prompt_label}{masked_hint}:",
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if value is None:
|
||||
raise KeyboardInterrupt()
|
||||
value = value.strip()
|
||||
if not value and current:
|
||||
# Keep existing value when user just presses Enter.
|
||||
continue
|
||||
updates[field_name] = value
|
||||
|
||||
# Allowed senders (common for all channels)
|
||||
senders_field = f"{ch_name}_allowed_senders"
|
||||
if hasattr(config, senders_field):
|
||||
senders = questionary.text(
|
||||
"Allowed senders (comma-separated, empty = all):",
|
||||
default=getattr(config, senders_field, ""),
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
).ask()
|
||||
if senders is None:
|
||||
raise KeyboardInterrupt()
|
||||
updates[senders_field] = senders.strip()
|
||||
|
||||
# Probe validation
|
||||
_probe_channel(ch_name, config, updates)
|
||||
|
||||
enabled_channels.append(ch_name)
|
||||
|
||||
updates["channel_enabled"] = ",".join(enabled_channels)
|
||||
# Keep legacy field in sync
|
||||
updates["imessage_enabled"] = "imessage" in enabled_channels
|
||||
|
||||
# --- Common prompt: send thinking (shown when any channel is enabled) ---
|
||||
if enabled_channels:
|
||||
console.print("\n [bold cyan]── Channel Settings ──[/bold cyan]")
|
||||
thinking_choices = [
|
||||
Choice(title="On (forward model reasoning)", value=True),
|
||||
Choice(title="Off (only send final responses)", value=False),
|
||||
]
|
||||
|
||||
send_thinking = questionary.select(
|
||||
"Send thinking panel in channel?",
|
||||
choices=thinking_choices,
|
||||
default=config.channel_send_thinking,
|
||||
style=WIZARD_STYLE,
|
||||
qmark=f" {QMARK}",
|
||||
use_indicator=True,
|
||||
).ask()
|
||||
|
||||
if send_thinking is None:
|
||||
raise KeyboardInterrupt()
|
||||
|
||||
updates["channel_send_thinking"] = send_thinking
|
||||
|
||||
return updates
|
||||
|
||||
|
||||
def _probe_channel(
|
||||
ch_name: str,
|
||||
config: EvoScientistConfig,
|
||||
updates: dict[str, object],
|
||||
) -> None:
|
||||
"""Run the probe for a channel type and print the result.
|
||||
|
||||
Non-fatal: prints a warning on failure but does not prevent enabling.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
def _val(key: str, fallback: str = "") -> str:
|
||||
"""Get a value from updates first, then config, then fallback."""
|
||||
if key in updates:
|
||||
return str(updates[key])
|
||||
return str(getattr(config, key, fallback))
|
||||
|
||||
console.print(" [dim]Validating credentials...[/dim]")
|
||||
|
||||
async def _run() -> tuple[bool, str]:
|
||||
if ch_name == "telegram":
|
||||
from ...channels.telegram.probe import validate_telegram_token
|
||||
|
||||
return await validate_telegram_token(
|
||||
_val("telegram_bot_token"),
|
||||
_val("telegram_proxy") or None,
|
||||
)
|
||||
elif ch_name == "discord":
|
||||
from ...channels.discord.probe import validate_discord_token
|
||||
|
||||
return await validate_discord_token(
|
||||
_val("discord_bot_token"),
|
||||
_val("discord_proxy") or None,
|
||||
)
|
||||
elif ch_name == "slack":
|
||||
from ...channels.slack.probe import validate_slack_tokens
|
||||
|
||||
return await validate_slack_tokens(
|
||||
_val("slack_bot_token"),
|
||||
_val("slack_app_token") or None,
|
||||
_val("slack_proxy") or None,
|
||||
)
|
||||
elif ch_name == "wechat":
|
||||
backend = _val("wechat_backend", "wecom")
|
||||
if backend == "wechatmp":
|
||||
from ...channels.wechat.probe import validate_wechat_mp
|
||||
|
||||
return await validate_wechat_mp(
|
||||
_val("wechat_mp_app_id"),
|
||||
_val("wechat_mp_app_secret"),
|
||||
_val("wechat_proxy") or None,
|
||||
)
|
||||
elif backend == "personal":
|
||||
from ...channels.wechat.probe import validate_wechat_personal
|
||||
|
||||
return await validate_wechat_personal(
|
||||
_val("wechat_personal_account_id"),
|
||||
_val("wechat_personal_token"),
|
||||
)
|
||||
else:
|
||||
from ...channels.wechat.probe import validate_wecom
|
||||
|
||||
return await validate_wecom(
|
||||
_val("wechat_wecom_corp_id"),
|
||||
_val("wechat_wecom_secret"),
|
||||
_val("wechat_proxy") or None,
|
||||
)
|
||||
elif ch_name == "feishu":
|
||||
from ...channels.feishu.probe import validate_feishu_credentials
|
||||
|
||||
return await validate_feishu_credentials(
|
||||
_val("feishu_app_id"),
|
||||
_val("feishu_app_secret"),
|
||||
_val("feishu_domain", "https://open.feishu.cn"),
|
||||
)
|
||||
elif ch_name == "dingtalk":
|
||||
from ...channels.dingtalk.probe import validate_dingtalk
|
||||
|
||||
return await validate_dingtalk(
|
||||
_val("dingtalk_client_id"),
|
||||
_val("dingtalk_client_secret"),
|
||||
_val("dingtalk_proxy") or None,
|
||||
)
|
||||
elif ch_name == "email":
|
||||
from ...channels.email.probe import validate_email_imap
|
||||
|
||||
return await validate_email_imap(
|
||||
_val("email_imap_host"),
|
||||
int(_val("email_imap_port", "993")),
|
||||
_val("email_imap_username"),
|
||||
_val("email_imap_password"),
|
||||
_val("email_imap_use_ssl", "True").lower() not in ("false", "0", "no"),
|
||||
)
|
||||
elif ch_name == "qq":
|
||||
from ...channels.qq.probe import validate_qq
|
||||
|
||||
return await validate_qq(
|
||||
_val("qq_app_id"),
|
||||
_val("qq_app_secret"),
|
||||
)
|
||||
elif ch_name == "signal":
|
||||
from ...channels.signal.probe import validate_signal
|
||||
|
||||
return await validate_signal(
|
||||
_val("signal_phone_number"),
|
||||
_val("signal_cli_path", "signal-cli"),
|
||||
int(_val("signal_rpc_port", "7583")),
|
||||
)
|
||||
else:
|
||||
return True, "No probe available"
|
||||
|
||||
try:
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
if loop.is_running():
|
||||
import nest_asyncio # type: ignore[import-untyped]
|
||||
|
||||
nest_asyncio.apply()
|
||||
except RuntimeError:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
|
||||
ok, detail = loop.run_until_complete(_run())
|
||||
if ok:
|
||||
console.print(f" [green]✓ {detail}[/green]")
|
||||
else:
|
||||
console.print(f" [yellow]⚠ {detail}[/yellow]")
|
||||
console.print(
|
||||
" [dim]Channel will still be enabled — check credentials later.[/dim]"
|
||||
)
|
||||
except Exception as e:
|
||||
console.print(f" [yellow]⚠ Could not validate: {e}[/yellow]")
|
||||
console.print(
|
||||
" [dim]Channel will still be enabled — check credentials later.[/dim]"
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Progress Rendering (for tests and potential future use)
|
||||
# =============================================================================
|
||||