mirror of
https://github.com/aaPanel/aaPanel.git
synced 2026-09-28 20:34:52 +02:00
Update to 7.27.0
1. Added Domains menu to manage Domains and SSL Certificates 2. Added Add Site automatic record creation when adding websites (Only supported by IPv4) 3. Added Vietnamese language 4. Added Indonesian language 5. Added HTTPS Protection function for Website (disables automatic HTTP to HTTPS redirection when enabled) 6. Added Mail Marketing -- Automation trigger tasks 7. Added Mail Marketing -- Groups Import, Export, Merge 8. Added Mail Marketing -- Subscribers to support paste import 9. Added Mail Marketing -- Template Import Export Duplicate 10. Added Mail server -- Other Settings add timed Auto Responder 11. Added Mail Marketing -- Suspend List to detect abnormal mailboxes 12. Added Mail Marketing -- Marketing Task to export error mail logs 13. Added Mail Domain -- SSL certificate expiration alert 14. Optimized Webmail management to display based on domain additions 15. Optimize Website Interface Button Integration 16. Optimize the response speed of WP Toolkit interface
This commit is contained in:
+353
-113
@@ -79,10 +79,16 @@ class acme_v2:
|
||||
_conf_file = 'config/letsencrypt.json'
|
||||
_conf_file_v2 = 'config/letsencrypt_v2.json'
|
||||
_request_type = 'curl'
|
||||
# ========== dns domian ============
|
||||
_log_path = f"{public.get_panel_path()}/logs/dns_domain_logs"
|
||||
_log_file = ""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, debug: bool = False):
|
||||
if not os.path.exists(self._conf_file_v2) and os.path.exists(self._conf_file):
|
||||
shutil.copyfile(self._conf_file, self._conf_file_v2)
|
||||
if not os.path.exists(self._log_path):
|
||||
os.makedirs(self._log_path)
|
||||
self._debug = debug
|
||||
if self._debug:
|
||||
self._url = 'https://acme-staging-v02.api.letsencrypt.org/directory'
|
||||
else:
|
||||
@@ -92,6 +98,20 @@ class acme_v2:
|
||||
self._nginx_cache_file_auth = {}
|
||||
self._can_use_lua = None
|
||||
self._well_known_check_cache = {}
|
||||
self._task_obj = None
|
||||
|
||||
def logger(self, log_str, mode="ab+"):
|
||||
if self._log_file:
|
||||
# 每个ssl独立log显示
|
||||
log = self._log_file
|
||||
else:
|
||||
# letsencrypt 公共log, 续签之类
|
||||
log = 'logs/letsencrypt.log'
|
||||
f = open(log, mode)
|
||||
log_str += "\n"
|
||||
f.write(log_str.encode('utf-8'))
|
||||
f.close()
|
||||
return True
|
||||
|
||||
def can_use_lua_module(self):
|
||||
if self._can_use_lua is None:
|
||||
@@ -246,6 +266,7 @@ class acme_v2:
|
||||
account = self._config['account'][k]
|
||||
account['email'] = self._config['email']
|
||||
self.set_crond()
|
||||
self.set_crond_v2()
|
||||
return account
|
||||
except Exception as ex:
|
||||
return public.return_msg_gettext(False, str(ex))
|
||||
@@ -396,6 +417,7 @@ class acme_v2:
|
||||
# 是否自动构造通配符
|
||||
if self._auto_wildcard:
|
||||
domains = self.auto_wildcard(domains)
|
||||
domains.sort()
|
||||
wildcard = []
|
||||
tmp_domains = []
|
||||
for domain in domains:
|
||||
@@ -576,6 +598,7 @@ class acme_v2:
|
||||
identifier_auth['auth_to'] = self._config['orders'][index]['auth_to']
|
||||
identifier_auth['type'] = self._config['orders'][index]['auth_type']
|
||||
# 设置验证信息
|
||||
# DNS Api add dns record
|
||||
self.set_auth_info(identifier_auth, index=index)
|
||||
|
||||
auths.append(identifier_auth)
|
||||
@@ -611,6 +634,7 @@ class acme_v2:
|
||||
index
|
||||
)
|
||||
else: # auth_to=dns-api
|
||||
# DNS Api add dns record
|
||||
self.create_dns_record(
|
||||
identifier_auth['auth_to'],
|
||||
identifier_auth['domain'],
|
||||
@@ -632,7 +656,7 @@ class acme_v2:
|
||||
return True
|
||||
acme_path = '{}/.well-known/acme-challenge'.format(self._config['orders'][index]['auth_to'])
|
||||
acme_path = acme_path.replace("//", '/')
|
||||
write_log('|-Verify the dir:{}'.format(acme_path))
|
||||
self.logger('|-Verify the dir:{}'.format(acme_path))
|
||||
if os.path.exists(acme_path):
|
||||
public.ExecShell("rm -f {}/*".format(acme_path))
|
||||
|
||||
@@ -832,36 +856,36 @@ if ( $well_known != "" ) {
|
||||
else:
|
||||
public.writeFile(file_check_config_path, old_data)
|
||||
|
||||
# 解析域名
|
||||
# def create_dns_record(self, auth_to, domain, dns_value):
|
||||
# # 如果为手动解析
|
||||
# if auth_to == 'dns':
|
||||
# return None
|
||||
#
|
||||
# from panelDnsapi import DnsMager
|
||||
#
|
||||
# self._dns_class = DnsMager().get_dns_obj_by_domain(domain)
|
||||
# self._dns_class.create_dns_record(public.de_punycode(domain), dns_value)
|
||||
# self._dns_domains.append({"domain": domain, "dns_value": dns_value})
|
||||
|
||||
# 解析域名
|
||||
# 解析挑战域名
|
||||
def create_dns_record(self, auth_to, domain, dns_value):
|
||||
# 如果为手动解析
|
||||
if auth_to == 'dns':
|
||||
return None
|
||||
|
||||
if auth_to.find('|') != -1:
|
||||
import panelDnsapi
|
||||
dns_name, key, secret = self.get_dnsapi(auth_to)
|
||||
self._dns_class = getattr(panelDnsapi, dns_name)(key, secret)
|
||||
self._dns_class.create_dns_record(public.de_punycode(domain), dns_value)
|
||||
else:
|
||||
pass
|
||||
# todo 待重构
|
||||
# from sslModel import dataModel
|
||||
# dataModel.main().add_dns_value_by_domain(domain, dns_value, is_let_txt=True)
|
||||
self._dns_domains.append({"domain": domain, "dns_value": dns_value})
|
||||
return
|
||||
from ssl_domainModelV2.model import DnsDomainProvider
|
||||
dns_name, account, token = auth_to.split('|')
|
||||
if not dns_name or not token: # account may be empty, cf limit
|
||||
raise Exception('dns_name or account or token is empty')
|
||||
try:
|
||||
provider = DnsDomainProvider.objects.filter(
|
||||
api_user=account, api_key=token, status=1,
|
||||
).first()
|
||||
if provider:
|
||||
# v2
|
||||
self._dns_class = provider.dns_obj
|
||||
self._dns_class.create_dns_record(public.de_punycode(domain), dns_value)
|
||||
self._dns_domains.append({"domain": domain, "dns_value": dns_value})
|
||||
else:
|
||||
# 旧调用方式
|
||||
import panelDnsapi
|
||||
dns_name, key, secret = self.get_dnsapi(auth_to)
|
||||
cf_limit_api = "/www/server/panel/data/cf_limit_api.pl"
|
||||
limit = True if os.path.exists(cf_limit_api) else False
|
||||
self._dns_class = getattr(panelDnsapi, dns_name)(key, secret, limit)
|
||||
self._dns_class.create_dns_record(public.de_punycode(domain), dns_value)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
print(traceback.format_exc())
|
||||
raise Exception("error: %s" % e)
|
||||
|
||||
# 解析DNSAPI信息
|
||||
def get_dnsapi(self, auth_to):
|
||||
@@ -886,7 +910,7 @@ if ( $well_known != "" ) {
|
||||
secret = tmp[2]
|
||||
return dns_name, key, secret
|
||||
|
||||
# 删除域名解析
|
||||
# 删除挑战域名解析
|
||||
def remove_dns_record(self):
|
||||
if not self._dns_domains:
|
||||
return None
|
||||
@@ -895,11 +919,6 @@ if ( $well_known != "" ) {
|
||||
if self._dns_class:
|
||||
self._dns_class.delete_dns_record(
|
||||
public.de_punycode(dns_info['domain']), dns_info['dns_value'])
|
||||
else:
|
||||
pass
|
||||
# todo 待重构
|
||||
# from sslModel import dataModel
|
||||
# dataModel.main().del_dns_value_by_domain(dns_info['domain'], is_let_txt=True)
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
@@ -941,7 +960,7 @@ if ( $well_known != "" ) {
|
||||
domain = get.domain
|
||||
self.get_apis()
|
||||
|
||||
write_log("|-domain is Verifying...:{}".format(domain))
|
||||
self.logger("|-domain is Verifying...:{}".format(domain))
|
||||
if index not in self._config['orders']:
|
||||
return public.return_msg_gettext(False, public.lang('The order does not exist!'))
|
||||
order = self._config['orders'][index]
|
||||
@@ -963,7 +982,7 @@ if ( $well_known != "" ) {
|
||||
auth['status'] = _return['status']
|
||||
return auth
|
||||
auth['status'] = res.json()['status']
|
||||
write_log("|-Verification succeeded!")
|
||||
self.logger("|-Verification succeeded!")
|
||||
return auth
|
||||
except StopIteration as e:
|
||||
auth['status'] = 'invalid'
|
||||
@@ -973,14 +992,14 @@ if ( $well_known != "" ) {
|
||||
msg[1] = json.loads(msg[1])
|
||||
else:
|
||||
msg = ex
|
||||
write_log(public.get_error_info())
|
||||
self.logger(public.get_error_info())
|
||||
auth['error'] = msg
|
||||
return public.return_msg_gettext(False, public.lang(msg))
|
||||
finally:
|
||||
self.save_config()
|
||||
if "pending" not in [i.get("status", "pending") for i in order['auths']]:
|
||||
return self.apply_dns_auth(get)
|
||||
write_log("|-This domain name was not found in the order: {}".format(domain))
|
||||
self.logger("|-This domain name was not found in the order: {}".format(domain))
|
||||
return public.return_msg_gettext(False, "This domain name was not found in the order: {}".format(domain))
|
||||
|
||||
# 检查验证状态
|
||||
@@ -989,18 +1008,18 @@ if ( $well_known != "" ) {
|
||||
number_of_checks = 0
|
||||
while True:
|
||||
if desired_status == ['valid', 'invalid']:
|
||||
write_log('|-{} Query verification results..'.format(str(number_of_checks + 1)))
|
||||
self.logger('|-{} Query verification results..'.format(str(number_of_checks + 1)))
|
||||
time.sleep(self._wait_time)
|
||||
check_authorization_status_response = self.acme_request(url, "")
|
||||
a_auth = check_authorization_status_response.json()
|
||||
if not isinstance(a_auth, dict):
|
||||
write_log(a_auth)
|
||||
self.logger(a_auth)
|
||||
continue
|
||||
authorization_status = a_auth["status"]
|
||||
number_of_checks += 1
|
||||
if authorization_status in desired_status:
|
||||
if authorization_status == "invalid":
|
||||
write_log("|-Verification failed")
|
||||
self.logger("|-Verification failed")
|
||||
try:
|
||||
if 'error' in a_auth['challenges'][0]:
|
||||
ret_title = a_auth['challenges'][0]['error']['detail']
|
||||
@@ -1030,12 +1049,12 @@ if ( $well_known != "" ) {
|
||||
str(self._wait_time)
|
||||
))
|
||||
if desired_status == ['valid', 'invalid']:
|
||||
write_log("|-Verification succeeded!")
|
||||
self.logger("|-Verification succeeded!")
|
||||
return check_authorization_status_response
|
||||
|
||||
# 格式化错误输出
|
||||
def get_error(self, error):
|
||||
write_log("error_result: " + str(error))
|
||||
self.logger("error_result: " + str(error))
|
||||
if error.find("Max checks allowed") >= 0:
|
||||
return public.lang(
|
||||
'CA cannot verify your domain name, please check if the domain name resolution is correct, or wait 5-10 minutes and try again.')
|
||||
@@ -1152,6 +1171,7 @@ if ( $well_known != "" ) {
|
||||
)
|
||||
)
|
||||
send_csr_response_json = send_csr_response.json()
|
||||
# ssl 证书地址
|
||||
certificate_url = send_csr_response_json.get("certificate", "")
|
||||
self._config['orders'][index]['certificate_url'] = certificate_url
|
||||
self.save_config()
|
||||
@@ -1205,14 +1225,19 @@ if ( $well_known != "" ) {
|
||||
)
|
||||
cert['save_path'] = self._config['orders'][index]['save_path']
|
||||
self.save_config()
|
||||
self.save_cert(cert, index)
|
||||
self.save_cert(cert, index) # 保存证书
|
||||
return cert
|
||||
|
||||
# 保存证书到文件
|
||||
def save_cert(self, cert, index):
|
||||
try:
|
||||
from ssl_manage import SSLManger
|
||||
SSLManger().save_by_data(cert['cert'] + cert['root'], cert['private_key'])
|
||||
# write db ssl_info
|
||||
SSLManger().save_by_data(
|
||||
certificate=cert['cert'] + cert['root'],
|
||||
private_key=cert['private_key'],
|
||||
log_file=self._log_file,
|
||||
)
|
||||
|
||||
domain_name = self._config['orders'][index]['domains'][0]
|
||||
path = self._config['orders'][index]['save_path']
|
||||
@@ -1250,9 +1275,13 @@ privkey.pem Paste into the key entry box
|
||||
fullchain.pem Paste into certificate input box
|
||||
'''
|
||||
public.writeFile(path + '/Description.txt', ps)
|
||||
# 替换新的证书文件和基本信息, 一旦替换去掉旧数据
|
||||
self.sub_all_cert(key_file, pem_file)
|
||||
except:
|
||||
write_log(public.get_error_info())
|
||||
except Exception as e:
|
||||
import traceback
|
||||
public.print_log(traceback.format_exc())
|
||||
public.print_log(f"---------save error {e}")
|
||||
self.logger(public.get_error_info())
|
||||
|
||||
def set_exclude_hash(self, order, exclude_hash):
|
||||
try:
|
||||
@@ -1379,7 +1408,7 @@ fullchain.pem Paste into certificate input box
|
||||
public.writeFile(
|
||||
to_key_file, public.readFile(key_file, 'rb'), 'wb')
|
||||
public.writeFile(to_info, json.dumps(cert_init))
|
||||
write_log(
|
||||
self.logger(
|
||||
'|-Detected that the certificate under {} '
|
||||
'overlaps with the certificate of this application and has an earlier expiration time, '
|
||||
'and has been replaced with a new certificate!'.format(to_path)
|
||||
@@ -1471,8 +1500,8 @@ fullchain.pem Paste into certificate input box
|
||||
|
||||
# 检查DNS记录
|
||||
def check_dns(self, domain, value, s_type='TXT'):
|
||||
write_log('|-Attempt to verify DNS records locally, '
|
||||
'domain name: {}, type: {} record value: {}'.format(domain, s_type, value))
|
||||
self.logger('|-Attempt to verify DNS records locally, '
|
||||
'domain name: {}, type: {} record value: {}'.format(domain, s_type, value))
|
||||
time.sleep(10)
|
||||
n = 0
|
||||
while n < 20:
|
||||
@@ -1483,9 +1512,9 @@ fullchain.pem Paste into certificate input box
|
||||
for j in ns.response.answer:
|
||||
for i in j.items:
|
||||
txt_value = i.to_text().replace('"', '').strip()
|
||||
write_log('|-Number of verifications: {}, value: {}'.format(n, txt_value))
|
||||
self.logger('|-Number of verifications: {}, value: {}'.format(n, txt_value))
|
||||
if txt_value == value:
|
||||
write_log("|-Local authentication succeeded!")
|
||||
self.logger("|-Local authentication succeeded!")
|
||||
return True
|
||||
except:
|
||||
try:
|
||||
@@ -1493,7 +1522,7 @@ fullchain.pem Paste into certificate input box
|
||||
except:
|
||||
return False
|
||||
time.sleep(3)
|
||||
write_log("|-Local authentication failed!")
|
||||
self.logger("|-Local authentication failed!")
|
||||
return True
|
||||
|
||||
# 创建CSR
|
||||
@@ -1582,7 +1611,7 @@ fullchain.pem Paste into certificate input box
|
||||
# 构造验证信息
|
||||
def get_identifier_auth(self, index, url, auth_info):
|
||||
s_type = self.get_auth_type(index)
|
||||
write_log('|-Verification type: {}'.format(s_type))
|
||||
self.logger('|-Verification type: {}'.format(s_type))
|
||||
domain = auth_info['identifier']['value']
|
||||
wildcard = False
|
||||
# 处理通配符
|
||||
@@ -1699,7 +1728,7 @@ fullchain.pem Paste into certificate input box
|
||||
payload64 = self.calculate_safe_base64(json.dumps(payload))
|
||||
protected = self.get_acme_header(url)
|
||||
protected64 = self.calculate_safe_base64(json.dumps(protected))
|
||||
signature = self.sign_message(
|
||||
signature = self.sign_message_new(
|
||||
message="{0}.{1}".format(protected64, payload64)
|
||||
) # bytes
|
||||
# signature = self.sign_message_new(
|
||||
@@ -1903,16 +1932,16 @@ fullchain.pem Paste into certificate input box
|
||||
# 申请证书
|
||||
def apply_cert(self, domains, auth_type='dns', auth_to='Dns_com|None|None', **args):
|
||||
index = ''
|
||||
write_log("", "wb+")
|
||||
self.logger("", "wb+")
|
||||
try:
|
||||
self.get_apis()
|
||||
index = None
|
||||
if 'index' in args:
|
||||
index = args['index']
|
||||
if not index: # 判断是否只想验证域名
|
||||
write_log(public.lang("|-Creating order.."))
|
||||
self.logger(public.lang("|-Creating order.."))
|
||||
index = self.create_order(domains, auth_type, auth_to)
|
||||
write_log(public.lang("|-Getting verification information.."))
|
||||
self.logger(public.lang("|-Getting verification information.."))
|
||||
self.get_auths(index) # add dns record or file record
|
||||
if auth_to == 'dns' and len(self._config['orders'][index]['auths']) > 0:
|
||||
auth_domains = [i["domain"].replace("*.", "") for i in self._config['orders'][index]['auths']]
|
||||
@@ -1923,16 +1952,16 @@ fullchain.pem Paste into certificate input box
|
||||
"Please verify the following domain names separately."
|
||||
)
|
||||
return self._config['orders'][index]
|
||||
write_log(public.lang("|-Verifying domain name.."))
|
||||
self.logger(public.lang("|-Verifying domain name.."))
|
||||
self.auth_domain(index)
|
||||
self.remove_dns_record()
|
||||
write_log(public.lang("|-Sending CSR.."))
|
||||
self.logger(public.lang("|-Sending CSR.."))
|
||||
self.send_csr(index)
|
||||
write_log(public.lang("|-Downloading certificate.."))
|
||||
self.logger(public.lang("|-Downloading certificate.."))
|
||||
cert = self.download_cert(index)
|
||||
cert['status'] = True
|
||||
cert['msg'] = public.lang("Application successful!")
|
||||
write_log(public.lang("|-Successful application, deploying to site.."))
|
||||
self.logger(public.lang("|-Successful application, deploying to site.."))
|
||||
return cert
|
||||
except Exception as ex:
|
||||
self.remove_dns_record()
|
||||
@@ -1942,7 +1971,7 @@ fullchain.pem Paste into certificate input box
|
||||
msg[1] = json.loads(msg[1])
|
||||
else:
|
||||
msg = ex
|
||||
write_log(public.get_error_info())
|
||||
self.logger(public.get_error_info())
|
||||
_res = {"status": False, "msg": msg, "index": index}
|
||||
return _res
|
||||
|
||||
@@ -2107,7 +2136,7 @@ fullchain.pem Paste into certificate input box
|
||||
args_obj = public.dict_obj()
|
||||
if not cron_id:
|
||||
cronPath = public.GetConfigValue('setup_path') + '/cron/' + echo
|
||||
shell = '{} -u /www/server/panel/class/acme_v2.py --renew_v2=1'.format(sys.executable)
|
||||
shell = '{} -u /www/server/panel/class/acme_v2.py --renew_v3=1'.format(sys.executable)
|
||||
public.writeFile(cronPath, shell)
|
||||
|
||||
# 使用随机时间
|
||||
@@ -2145,6 +2174,56 @@ fullchain.pem Paste into certificate input box
|
||||
except:
|
||||
pass
|
||||
|
||||
# 创建计划任务v2
|
||||
def set_crond_v2(self):
|
||||
try:
|
||||
echo = public.md5(public.md5('domain_ssl_renew_lets_ssl_bt'))
|
||||
find = public.M('crontab').where('echo=?', (echo,)).find()
|
||||
cron_id = find['id'] if find else None
|
||||
|
||||
import crontab
|
||||
import random
|
||||
args_obj = public.dict_obj()
|
||||
if not cron_id:
|
||||
cronPath = public.GetConfigValue('setup_path') + '/cron/' + echo
|
||||
shell = '{} -u /www/server/panel/class/acme_v2.py --renew_v3=1'.format(sys.executable)
|
||||
public.writeFile(cronPath, shell)
|
||||
|
||||
# 使用随机时间
|
||||
hour = random.randint(0, 23)
|
||||
minute = random.randint(1, 59)
|
||||
args_obj.id = public.M('crontab').add(
|
||||
'name,type,where1,where_hour,where_minute,echo,addtime,status,save,backupTo,sType,sName,sBody,urladdress',
|
||||
("Domain SSL Renew Let's Encrypt Certificate", 'day', '', hour, minute, echo,
|
||||
time.strftime('%Y-%m-%d %X', time.localtime()), 0, '', 'localhost', 'toShell', '', shell, ''))
|
||||
crontab.crontab().set_cron_status(args_obj)
|
||||
else:
|
||||
# 检查任务如果是0点10分执行,改为随机时间
|
||||
if find['where_hour'] == 0 and find['where_minute'] == 10:
|
||||
# 使用随机时间
|
||||
hour = random.randint(0, 23)
|
||||
minute = random.randint(1, 59)
|
||||
public.M('crontab').where('id=?', (cron_id,)).save('where_hour,where_minute,status',
|
||||
(hour, minute, 0))
|
||||
|
||||
# 停用任务
|
||||
args_obj.id = cron_id
|
||||
crontab.crontab().set_cron_status(args_obj)
|
||||
|
||||
# 启用任务
|
||||
public.M('crontab').where('id=?', (cron_id,)).setField('status', 1)
|
||||
crontab.crontab().set_cron_status(args_obj)
|
||||
|
||||
cron_path = public.get_cron_path()
|
||||
if os.path.exists(cron_path):
|
||||
cron_s = public.readFile(cron_path)
|
||||
if cron_s.find(echo) == -1:
|
||||
public.M('crontab').where('id=?', (cron_id,)).setField('status', 0)
|
||||
args_obj.id = cron_id
|
||||
crontab.crontab().set_cron_status(args_obj)
|
||||
except:
|
||||
pass
|
||||
|
||||
# 获取当前正在使用此证书的网站目录
|
||||
def get_ssl_used_site(self, index):
|
||||
hash_dic = self.get_exclude_hash(public.dict_obj())
|
||||
@@ -2234,8 +2313,8 @@ fullchain.pem Paste into certificate input box
|
||||
if index in self._config['orders'].keys(): continue # 已在订单列表
|
||||
|
||||
n += 1
|
||||
write_log("|-Renewing additional certificate {}, domain name:{}..".format(n, cert_init['subject']))
|
||||
write_log("|-Creating order..")
|
||||
self.logger("|-Renewing additional certificate {}, domain name:{}..".format(n, cert_init['subject']))
|
||||
self.logger("|-Creating order..")
|
||||
args.id = siteInfo['id']
|
||||
runPath = siteObj.GetRunPath(args)
|
||||
if runPath and not runPath in ['/']:
|
||||
@@ -2245,7 +2324,7 @@ fullchain.pem Paste into certificate input box
|
||||
|
||||
self.renew_cert_to(cert_init['dns'], 'http', path.replace('//', '/'))
|
||||
except:
|
||||
write_log("|-Renewal failed:")
|
||||
self.logger("|-Renewal failed:")
|
||||
|
||||
# 关闭强制https
|
||||
def close_httptohttps(self, siteName):
|
||||
@@ -2303,7 +2382,7 @@ fullchain.pem Paste into certificate input box
|
||||
siteName = public.M('sites').where('path=?', auth_to).getField('name')
|
||||
r_status = public.M('sites').where('path=?', auth_to).getField('status')
|
||||
if r_status != '1':
|
||||
write_log(
|
||||
self.logger(
|
||||
"|- This certificate uses the【file verification】method,"
|
||||
"but the website【{}】has not been started,"
|
||||
"so the renewal can only be skipped。".format(siteName)
|
||||
@@ -2328,10 +2407,10 @@ fullchain.pem Paste into certificate input box
|
||||
|
||||
isError = public.checkWebConfig()
|
||||
if isError is not True and public.get_webserver() == "nginx":
|
||||
write_log(
|
||||
self.logger(
|
||||
"|- The certificate uses the file verification method, but currently it cannot overload the nginx server configuration file and can only skip renewal.")
|
||||
write_log("|- The error message in the configuration file is as follows:")
|
||||
write_log(isError)
|
||||
self.logger("|- The error message in the configuration file is as follows:")
|
||||
self.logger(isError)
|
||||
|
||||
is_rep = self.close_httptohttps(siteName)
|
||||
try:
|
||||
@@ -2341,14 +2420,14 @@ fullchain.pem Paste into certificate input box
|
||||
auth_to.replace('//', '/'),
|
||||
index
|
||||
)
|
||||
write_log("|-Getting verification information..")
|
||||
self.logger("|-Getting verification information..")
|
||||
self.get_auths(index) # add dns record or file record
|
||||
write_log("|-Verifying domain name..")
|
||||
self.logger("|-Verifying domain name..")
|
||||
self.auth_domain(index)
|
||||
write_log("|-Sending CSR..")
|
||||
self.logger("|-Sending CSR..")
|
||||
self.remove_dns_record()
|
||||
self.send_csr(index)
|
||||
write_log("|-Downloading certificate..")
|
||||
self.logger("|-Downloading certificate..")
|
||||
cert = self.download_cert(index)
|
||||
self._config['orders'][index]['renew_time'] = int(time.time())
|
||||
|
||||
@@ -2360,7 +2439,7 @@ fullchain.pem Paste into certificate input box
|
||||
self.save_config()
|
||||
cert['status'] = True
|
||||
cert['msg'] = 'Renewed successfully!'
|
||||
write_log("|-Renewed successfully!!")
|
||||
self.logger("|-Renewed successfully!!")
|
||||
except Exception as e:
|
||||
|
||||
if str(e).find('please try again later') == -1: # 受其它证书影响和连接CA失败的的不记录重试次数
|
||||
@@ -2380,16 +2459,16 @@ fullchain.pem Paste into certificate input box
|
||||
else:
|
||||
msg = e
|
||||
err = {}
|
||||
write_log("|-" + msg)
|
||||
self.logger("|-" + msg)
|
||||
return {"status": False, "msg": msg, "err": err}
|
||||
finally:
|
||||
if is_rep: self.rep_httptohttps(siteName)
|
||||
write_log("-" * 70)
|
||||
self.logger("-" * 70)
|
||||
return cert
|
||||
|
||||
# 续签证书
|
||||
def renew_cert(self, index, cycle=None):
|
||||
write_log("", "wb+")
|
||||
self.logger("", "wb+")
|
||||
set_status = False
|
||||
index_info = None
|
||||
try:
|
||||
@@ -2399,7 +2478,7 @@ fullchain.pem Paste into certificate input box
|
||||
if type(index) != str:
|
||||
index = index.index
|
||||
if not index in self._config['orders']:
|
||||
write_log("|-指定订单号不存在,无法续签!")
|
||||
self.logger("|-指定订单号不存在,无法续签!")
|
||||
self.set_auto_renew_status(index, -1, "指定订单号不存在,无法续签!")
|
||||
raise Exception(
|
||||
public.lang("The specified order number does not exist and cannot be renewed!")
|
||||
@@ -2410,7 +2489,7 @@ fullchain.pem Paste into certificate input box
|
||||
self.set_auto_renew_status(
|
||||
index, 0, "|-the expiration date is greater than{}days,skip the renewal!".format(cycle)
|
||||
)
|
||||
write_log("|-过期时间大于{}天,跳过续签!".format(cycle))
|
||||
self.logger("|-过期时间大于{}天,跳过续签!".format(cycle))
|
||||
return
|
||||
order_index.append(index)
|
||||
else:
|
||||
@@ -2437,7 +2516,7 @@ fullchain.pem Paste into certificate input box
|
||||
if not public.M('domain').where("name=?", (domain,)).count() and not public.M('binding').where(
|
||||
"domain=?", domain).count():
|
||||
_auth_to = None
|
||||
write_log("|-Skip deleted domain names: {}".format(self._config['orders'][i]['domains']))
|
||||
self.logger("|-Skip deleted domain names: {}".format(self._config['orders'][i]['domains']))
|
||||
if not _auth_to: continue
|
||||
|
||||
self._config['orders'][i]['auth_to'] = _auth_to
|
||||
@@ -2446,7 +2525,7 @@ fullchain.pem Paste into certificate input box
|
||||
if 'next_retry_time' in self._config['orders'][i]:
|
||||
timeout = self._config['orders'][i]['next_retry_time'] - int(time.time())
|
||||
if timeout > 0:
|
||||
write_log(
|
||||
self.logger(
|
||||
'|-Skipping domain name: {} this time, due to last renewal failure, we need to wait {} hours before trying again'.format(
|
||||
self._config['orders'][i]['domains'], int(timeout / 60 / 60)))
|
||||
continue
|
||||
@@ -2454,12 +2533,12 @@ fullchain.pem Paste into certificate input box
|
||||
# 加入到续签订单
|
||||
order_index.append(i)
|
||||
if not order_index:
|
||||
write_log("|-No SSL certificate found within 30 days!")
|
||||
self.logger("|-No SSL certificate found within 30 days!")
|
||||
self.get_apis()
|
||||
self.renew_cert_other()
|
||||
write_log("|-All tasks have been processed!")
|
||||
self.logger("|-All tasks have been processed!")
|
||||
return
|
||||
write_log("|-A total of {} certificates need to be renewed".format(len(order_index)))
|
||||
self.logger("|-A total of {} certificates need to be renewed".format(len(order_index)))
|
||||
n = 0
|
||||
self.get_apis()
|
||||
cert = None
|
||||
@@ -2470,7 +2549,7 @@ fullchain.pem Paste into certificate input box
|
||||
self._config['orders'][index]['auth_to'],
|
||||
self._config['orders'][index]['auth_type'])
|
||||
if len(domains) == 0:
|
||||
write_log(
|
||||
self.logger(
|
||||
"|-The domain names under the {} certificate are all unused (these domains are: [%s]) and have been skipped.".format(
|
||||
n, ",".join(self._config['orders'][index]['domains'])))
|
||||
err_msg = "All domain names are unused and have been skipped."
|
||||
@@ -2480,11 +2559,11 @@ fullchain.pem Paste into certificate input box
|
||||
else:
|
||||
index_info = self._config['orders'][index]
|
||||
self._config['orders'][index]['domains'] = domains
|
||||
write_log(
|
||||
self.logger(
|
||||
'|-Renewing certificate number of {},domain: {}..'.format(n, str(
|
||||
self._config['orders'][index]['domains']))
|
||||
)
|
||||
write_log("|-Creating order..")
|
||||
self.logger("|-Creating order..")
|
||||
cert = self.renew_cert_to(self._config['orders'][index]['domains'],
|
||||
self._config['orders'][index]['auth_type'],
|
||||
self._config['orders'][index]['auth_to'], index)
|
||||
@@ -2507,7 +2586,7 @@ fullchain.pem Paste into certificate input box
|
||||
msg[1] = json.loads(msg[1])
|
||||
else:
|
||||
msg = ex
|
||||
write_log(public.get_error_info())
|
||||
self.logger(public.get_error_info())
|
||||
if set_status:
|
||||
self.set_auto_renew_status(index, -1, msg)
|
||||
if index_info:
|
||||
@@ -2532,10 +2611,14 @@ fullchain.pem Paste into certificate input box
|
||||
from sslModel import base
|
||||
for ssl_hash in hash_list:
|
||||
s += 1
|
||||
write_log(f"|-Renewing the {s} certificate,There are {len(hash_list)} certificates in total...")
|
||||
write_log_old(f"|-Renewing the {s} certificate,There are {len(hash_list)} certificates in total...")
|
||||
cert_data = public.M('ssl_info').where('hash=?', ssl_hash).find()
|
||||
if not cert_data:
|
||||
write_log(
|
||||
from ssl_manage import ssl_db
|
||||
cert_data = ssl_db.connection().where('hash=?', ssl_hash).find()
|
||||
|
||||
if not cert_data:
|
||||
write_log_old(
|
||||
"|-【{}】The specified certificate information was not found"
|
||||
" and the renewal cannot be carried out!".format(ssl_hash)
|
||||
)
|
||||
@@ -2544,8 +2627,8 @@ fullchain.pem Paste into certificate input box
|
||||
cert_info = json.loads(cert_data['info'])
|
||||
auth_info = json.loads(cert_data['auth_info'])
|
||||
except:
|
||||
write_log(public.get_error_info())
|
||||
write_log(
|
||||
write_log_old(public.get_error_info())
|
||||
write_log_old(
|
||||
"|-【{}】The format of the certificate information is incorrect and"
|
||||
" renewal is not possible. Please try to renew it manually!".format(ssl_hash)
|
||||
)
|
||||
@@ -2553,12 +2636,12 @@ fullchain.pem Paste into certificate input box
|
||||
|
||||
if cert_info.get('issuer') not in ("R3", "R8", "R11", "R10", "R5") and cert_info.get(
|
||||
'issuer_O') != "Let's Encrypt":
|
||||
write_log("|-【{}】It's not a Let's Encrypt certificate and cannot be renewed!".format(ssl_hash))
|
||||
write_log_old("|-【{}】It's not a Let's Encrypt certificate and cannot be renewed!".format(ssl_hash))
|
||||
continue
|
||||
# 计算 30 天后的日期
|
||||
future_date = (datetime.datetime.now().date() + datetime.timedelta(days=cycle)).strftime('%Y-%m-%d')
|
||||
if future_date < cert_info['notAfter']:
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】The expiration date is greater than {} days,"
|
||||
" so there is no need for renewal!".format(ssl_hash, cycle)
|
||||
)
|
||||
@@ -2566,7 +2649,7 @@ fullchain.pem Paste into certificate input box
|
||||
# 判断是否有泛域名
|
||||
wildcard = False
|
||||
if "*" in ",".join(cert_info['dns']):
|
||||
write_log(
|
||||
write_log_old(
|
||||
f"|-【{ssl_hash}】 with domain【{cert_data.get('subject', '')}】There is a wildcard domain name, "
|
||||
f"and only the DNS verification method can be used for renewal!"
|
||||
)
|
||||
@@ -2584,14 +2667,14 @@ fullchain.pem Paste into certificate input box
|
||||
auth_domains.append(i)
|
||||
else:
|
||||
root_domain, _, _ = base.sslBase().extract_zone(i)
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-The root domain name【{}】is not bound to the dns-api, "
|
||||
"Skip the domain name: {}!".format(root_domain, i)
|
||||
)
|
||||
continue
|
||||
|
||||
if not auth_domains:
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】None of the domain names are bound to the dns-api,"
|
||||
" so the DNS verification renewal cannot be used.".format(ssl_hash)
|
||||
)
|
||||
@@ -2603,7 +2686,7 @@ fullchain.pem Paste into certificate input box
|
||||
|
||||
# dns api auth
|
||||
if set(auth_domains) == set(cert_info['dns']):
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】All domain names have been bound to the dns-api. "
|
||||
"An attempt is being made to use the DNS verification method for renewal.".format(ssl_hash)
|
||||
)
|
||||
@@ -2616,24 +2699,24 @@ fullchain.pem Paste into certificate input box
|
||||
domains = ""
|
||||
site_info = {}
|
||||
sites = {}
|
||||
write_log("|-【{}】Checking whether file verification is available.".format(ssl_hash))
|
||||
write_log_old("|-【{}】Checking whether file verification is available.".format(ssl_hash))
|
||||
for domain in cert_info['dns']:
|
||||
domain_info = public.M('domain').where("name=?", domain).find()
|
||||
if not domain_info:
|
||||
write_log("|-The domain name【{}】does not exist. Skip it!".format(domain))
|
||||
write_log_old("|-The domain name【{}】does not exist. Skip it!".format(domain))
|
||||
continue
|
||||
if not sites.get(domain_info["pid"]):
|
||||
sites[domain_info["pid"]] = [domain]
|
||||
else:
|
||||
sites[domain_info["pid"]].append(domain)
|
||||
if not sites:
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】No available file verification sites were found. "
|
||||
"Please try to renew it manually.".format(ssl_hash)
|
||||
)
|
||||
# 暂时不做多网站文件验证
|
||||
if len(sites.keys()) > 1:
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】It has been detected that the verification domain names "
|
||||
"are scattered across multiple sites. "
|
||||
"Multiple-site file verification is currently not supported!".format(ssl_hash)
|
||||
@@ -2641,27 +2724,27 @@ fullchain.pem Paste into certificate input box
|
||||
for site_id, domains in sites.items():
|
||||
site_info = public.M('sites').where("id=?", site_id).find()
|
||||
if not site_info:
|
||||
write_log("|-site【{}】do not exist. Skip it!".format(site_id))
|
||||
write_log_old("|-site【{}】do not exist. Skip it!".format(site_id))
|
||||
break
|
||||
if set(domains) == set(cert_info['dns']) or len(domains) > len(auth_domains):
|
||||
http_auth = True
|
||||
break
|
||||
|
||||
if http_auth and domains and site_info: # http auth
|
||||
write_log("|-【{}】Trying to use file verification for renewal!".format(ssl_hash))
|
||||
write_log_old("|-【{}】Trying to use file verification for renewal!".format(ssl_hash))
|
||||
self.get_apis()
|
||||
cert = self.renew_cert_to(domains=domains, auth_type="http", auth_to=site_info['path'])
|
||||
if cert.get('status') is False: continue
|
||||
continue
|
||||
elif dns_auth: # dns auth
|
||||
write_log("|-【{}】Trying to use DNS verification for renewal!".format(ssl_hash))
|
||||
write_log_old("|-【{}】Trying to use DNS verification for renewal!".format(ssl_hash))
|
||||
self.get_apis()
|
||||
cert = self.renew_cert_to(domains=auth_domains, auth_type="dns", auth_to="dns-api")
|
||||
if cert.get('status') is False:
|
||||
continue
|
||||
continue
|
||||
else:
|
||||
write_log(
|
||||
write_log_old(
|
||||
"|-【{}】No available verification methods were found."
|
||||
" Please try to renew it manually.".format(ssl_hash)
|
||||
)
|
||||
@@ -2900,6 +2983,139 @@ fullchain.pem Paste into certificate input box
|
||||
get.ssl_hash = ssl_hash
|
||||
return panelSSL.panelSSL().SetCertToSite(get)
|
||||
|
||||
def _generate_own_log(self, domains: list, auth_to: str):
|
||||
from hashlib import md5
|
||||
try:
|
||||
md5_obj = md5()
|
||||
dns, user, key = auth_to.split('|')
|
||||
body = f"{dns}{user}{key}{domains}"
|
||||
md5_obj.update(body.encode("utf-8"))
|
||||
self._log_file = f"{self._log_path}/{md5_obj.hexdigest()}.log"
|
||||
except Exception as e:
|
||||
public.print_log("error % s" % e)
|
||||
|
||||
def _set_task(self, val: int):
|
||||
if self._task_obj:
|
||||
self._task_obj.task_transfer(set_status=val)
|
||||
|
||||
# ========================= new ==============================
|
||||
def renew_cert_v3(self, index, cycle=30):
|
||||
if not 'class_v2/' in sys.path:
|
||||
sys.path.insert(0, 'class_v2/')
|
||||
from ssl_domainModelV2.model import DnsDomainSSL, DnsDomainProvider
|
||||
cycle = 30 if not cycle else int(cycle)
|
||||
if index:
|
||||
ssl_obj = DnsDomainSSL.objects.filter(hash=index)
|
||||
else:
|
||||
# 开了自动续签的ssl对象
|
||||
ssl_obj = DnsDomainSSL.objects.filter(auto_renew=1)
|
||||
s = 0
|
||||
count = ssl_obj.count()
|
||||
write_log("", "wb+")
|
||||
for ssl in ssl_obj:
|
||||
s += 1
|
||||
write_log(f"|-Renewing the {s} certificate,There are {count} certificates in total...")
|
||||
if ssl.info.get("issuer") not in (
|
||||
"R3", "R8", "R11", "R10", "R5"
|
||||
) and ssl.info.get("issuer_O") != "Let's Encrypt":
|
||||
write_log(
|
||||
f"|- Domain Subject:【{ssl.subject}】It's not a Let's Encrypt certificate and cannot be renewed!")
|
||||
continue
|
||||
# 计算 30 天后的日期
|
||||
after_ts = round((time.time() + 86400 * cycle) * 1000)
|
||||
if ssl.not_after_ts > after_ts:
|
||||
write_log(
|
||||
f"|- Domain Subject:【{ssl.subject}】The expiration date is greater than {cycle} days,"
|
||||
f" so there is no need for renewal!"
|
||||
)
|
||||
continue
|
||||
|
||||
if ssl.provider_id == 0:
|
||||
write_log(f"|- Domain Subject:【{ssl.subject}】is not bound to the dns-api, skip it...")
|
||||
continue
|
||||
|
||||
provider = DnsDomainProvider.objects.find_one(id=ssl.provider_id)
|
||||
if not provider:
|
||||
write_log(f"|- Domain Subject:【{ssl.subject}】The provider api info is not bound, skip it...")
|
||||
continue
|
||||
|
||||
write_log(f"|- Domain Subject:【{ssl.subject}】Trying to use DNS verification for renewal!")
|
||||
try:
|
||||
res = self.apply_cert_dns_domain(
|
||||
domains=ssl.dns,
|
||||
auth_to=f"{provider.name}|{provider.api_user}|{provider.api_key}",
|
||||
task_obj=None,
|
||||
auto_wildcard=False,
|
||||
)
|
||||
if res.get("status"):
|
||||
write_log(f"|- Domain Subject:【{ssl.subject}】Renewal SSL certificate Successfully!")
|
||||
else:
|
||||
write_log(
|
||||
f"|- Domain Subject:【{ssl.subject}】Renewal SSL certificate Error: {res.get('msg')[0]}"
|
||||
)
|
||||
except Exception as e:
|
||||
write_log(f"|- Domain Subject:【{ssl.subject}】Renewal SSL certificate Failed:{str(e)}")
|
||||
continue
|
||||
return
|
||||
|
||||
def apply_cert_dns_domain(
|
||||
self,
|
||||
domains: list,
|
||||
auth_to: str = '||',
|
||||
auth_type='dns',
|
||||
task_obj=None,
|
||||
auto_wildcard=False,
|
||||
**args
|
||||
):
|
||||
"""
|
||||
dns domain v2 申请证书
|
||||
"""
|
||||
# generate own log
|
||||
self._generate_own_log(domains, auth_to)
|
||||
self.logger("", "wb+")
|
||||
self._task_obj = task_obj
|
||||
index = ''
|
||||
self._auto_wildcard = auto_wildcard
|
||||
try:
|
||||
self.get_apis()
|
||||
self._set_task(5)
|
||||
index = None
|
||||
if 'index' in args:
|
||||
index = args['index']
|
||||
if not index: # 判断是否只想验证域名
|
||||
self.logger(public.lang("|-Creating order.."))
|
||||
index = self.create_order(domains, auth_type, auth_to)
|
||||
self._set_task(10)
|
||||
self.logger(public.lang("|-Getting verification information.."))
|
||||
# DNS Api add dns record
|
||||
self.get_auths(index)
|
||||
self._set_task(30)
|
||||
self.logger(public.lang("|-Verifying domain name.."))
|
||||
self.auth_domain(index)
|
||||
self._set_task(50)
|
||||
self.remove_dns_record()
|
||||
self.logger(public.lang("|-Sending CSR.."))
|
||||
self.send_csr(index)
|
||||
self._set_task(70)
|
||||
self.logger(public.lang("|-Downloading certificate.."))
|
||||
cert = self.download_cert(index)
|
||||
self._set_task(90)
|
||||
cert['status'] = True
|
||||
cert['msg'] = public.lang("Application successful!")
|
||||
self.logger(public.lang("|-Successful application, deploying to site.."))
|
||||
return cert
|
||||
except Exception as ex:
|
||||
self.remove_dns_record()
|
||||
ex = str(ex)
|
||||
if ex.find(">>>>") != -1:
|
||||
msg = ex.split(">>>>")
|
||||
msg[1] = json.loads(msg[1])
|
||||
else:
|
||||
msg = ex
|
||||
self.logger(public.get_error_info())
|
||||
_res = {"status": False, "msg": msg, "index": index}
|
||||
return _res
|
||||
|
||||
|
||||
def _test_domains(domains, auth_to, auth_type):
|
||||
# 检查站点域名变更情况, 若有删除域名,则在续签时,删除已经不使用的域名,再执行续签任务
|
||||
@@ -2937,16 +3153,26 @@ def echo_err(msg):
|
||||
exit()
|
||||
|
||||
|
||||
def write_log_old(log_str, mode="ab+"):
|
||||
if __name__ == "__main__":
|
||||
print(log_str)
|
||||
return
|
||||
_log_file = 'logs/letsencrypt_old.log'
|
||||
with open(_log_file, mode, encoding="utf-8") as f:
|
||||
log_str += "\n"
|
||||
f.write(log_str)
|
||||
return True
|
||||
|
||||
|
||||
# 写日志
|
||||
def write_log(log_str, mode="ab+"):
|
||||
if __name__ == "__main__":
|
||||
print(log_str)
|
||||
return
|
||||
_log_file = 'logs/letsencrypt.log'
|
||||
f = open(_log_file, mode)
|
||||
log_str += "\n"
|
||||
f.write(log_str.encode('utf-8'))
|
||||
f.close()
|
||||
with open(_log_file, mode, encoding="utf-8") as f:
|
||||
log_str += "\n"
|
||||
f.write(log_str)
|
||||
return True
|
||||
|
||||
|
||||
@@ -2966,6 +3192,8 @@ if __name__ == "__main__":
|
||||
p.add_argument('--index', default=None, help=public.lang("Specify the order index"), dest="index")
|
||||
p.add_argument('--renew', default=None, help=public.lang("renew certificate"), dest="renew")
|
||||
p.add_argument('--renew_v2', default=None, help=public.lang("renew certificate v2"), dest="renew_v2")
|
||||
|
||||
p.add_argument('--renew_v3', default=None, help=public.lang("renew certificate v3"), dest="renew_v3")
|
||||
p.add_argument('--revoke', default=None, help=public.lang("Revoke certificate"), dest="revoke")
|
||||
p.add_argument('--cycle', default=None, help=public.lang("Renew when the expiration time is lte certain time"),
|
||||
dest="cycle")
|
||||
@@ -2980,7 +3208,18 @@ if __name__ == "__main__":
|
||||
write_log(result)
|
||||
exit()
|
||||
|
||||
if args.renew_v2 or args.renew: # force run v2
|
||||
if args.renew_v3:
|
||||
sys.path.append(public.get_panel_path())
|
||||
p = acme_v2()
|
||||
if args.cycle:
|
||||
try:
|
||||
int(args.cycle)
|
||||
except:
|
||||
args.cycle = None
|
||||
p.renew_cert_v3(args.index, args.cycle)
|
||||
exit()
|
||||
# 计划弃置
|
||||
elif args.renew_v2 or args.renew: # force run v2
|
||||
sys.path.append(public.get_panel_path())
|
||||
p = acme_v2()
|
||||
if args.cycle:
|
||||
@@ -2990,6 +3229,7 @@ if __name__ == "__main__":
|
||||
args.cycle = None
|
||||
p.renew_cert_v2(args.index, args.cycle)
|
||||
exit()
|
||||
|
||||
else:
|
||||
try:
|
||||
if not args.index:
|
||||
|
||||
+1
-1
@@ -27,7 +27,7 @@ class panelSetup:
|
||||
if ua.find('spider') != -1 or g.ua.find('bot') != -1:
|
||||
return abort(403)
|
||||
|
||||
g.version = '7.25.0'
|
||||
g.version = '7.27.0'
|
||||
g.title = public.GetConfigValue('title')
|
||||
g.uri = request.path
|
||||
g.debug = os.path.exists('data/debug.pl')
|
||||
|
||||
@@ -93,6 +93,7 @@ def acme_crond_reinit():
|
||||
|
||||
import acme_v2
|
||||
acme_v2.acme_v2().set_crond()
|
||||
acme_v2.acme_v2().set_crond_v2()
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
+361
-269
@@ -86,6 +86,8 @@ extract_zone = ExtractZoneTool()
|
||||
class BaseDns(object):
|
||||
def __init__(self):
|
||||
self.dns_provider_name = self.__class__.__name__
|
||||
self.api_user = ""
|
||||
self.api_key = ""
|
||||
|
||||
def log_response(self, response):
|
||||
try:
|
||||
@@ -100,16 +102,25 @@ class BaseDns(object):
|
||||
def delete_dns_record(self, domain_name, domain_dns_value):
|
||||
raise NotImplementedError("delete_dns_record method must be implemented.")
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data):
|
||||
raise NotImplementedError("new method must be implemented.")
|
||||
|
||||
def remove_record(self, domain, host, s_type):
|
||||
raise NotImplementedError("remove_record method must be implemented.")
|
||||
|
||||
def add_record_for_creat_site(self, domain, server_ip):
|
||||
raise NotImplementedError("add_record_for_creat_site method must be implemented.")
|
||||
|
||||
# =============== 域名管理同步信息 =====================
|
||||
def get_domains(self):
|
||||
raise NotImplementedError("get_domains method must be implemented.")
|
||||
|
||||
def get_dns_record(self, domain_name):
|
||||
raise NotImplementedError("get_dns_record method must be implemented.")
|
||||
|
||||
def create_org_record(self, domain_name, record, record_value, record_type, ttl, **kwargs):
|
||||
raise NotImplementedError("create_org_record method must be implemented.")
|
||||
|
||||
def remove_record(self, domain_name, record, record_type):
|
||||
raise NotImplementedError("remove_record method must be implemented.")
|
||||
|
||||
def update_record(self, domain_name, record, new_record, ttl=1, **kwargs):
|
||||
raise NotImplementedError("update_record method must be implemented.")
|
||||
|
||||
def raise_resp_error(self, response: requests.Response):
|
||||
raise ValueError(
|
||||
"Error {dns_name}: status_code={status_code} response={response}".format(
|
||||
@@ -176,13 +187,6 @@ class TencentCloudDns(BaseDns):
|
||||
except TencentCloudSDKException as err:
|
||||
return public.returnMsg(False, public.lang('add fail, msg: {}'.format(err)))
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict):
|
||||
secret_id = conf_data.get("secret_id", "")
|
||||
secret_key = conf_data.get("secret_key", "")
|
||||
|
||||
return cls(secret_id, secret_key)
|
||||
|
||||
|
||||
# 未验证
|
||||
# noinspection PyUnresolvedReferences
|
||||
@@ -247,14 +251,6 @@ class HuaweiCloudDns(BaseDns):
|
||||
|
||||
return data
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict) -> BaseDns:
|
||||
ak = conf_data.get("ak", None) or conf_data.get("AccessKey", "")
|
||||
sk = conf_data.get("sk", None) or conf_data.get("SecretKey", "")
|
||||
project_id = conf_data.get("project_id", None) or conf_data.get("project_id", "")
|
||||
|
||||
return cls(ak, sk, project_id)
|
||||
|
||||
|
||||
# 未验证
|
||||
class DNSPodDns(BaseDns):
|
||||
@@ -342,24 +338,16 @@ class DNSPodDns(BaseDns):
|
||||
domain_name, zone, _ = extract_zone(domain)
|
||||
self.add_record(domain_name, zone, server_ip, "A")
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict) -> BaseDns:
|
||||
key = conf_data.get("key", None) or conf_data.get("ID", "")
|
||||
secret = conf_data.get("secret", None) or conf_data.get("Token", "")
|
||||
base_url = "https://dnsapi.cn/"
|
||||
|
||||
return cls(key, secret, base_url)
|
||||
|
||||
|
||||
class NameCheapDns(BaseDns):
|
||||
dns_provider_name = "namecheap"
|
||||
_type = 0 # 0:lest 1:锐成
|
||||
|
||||
def __init__(self, api_user, api_key):
|
||||
def __init__(self, api_user, api_key, **kwargs):
|
||||
super().__init__()
|
||||
self.timeout = 30
|
||||
self.api_user = api_user
|
||||
self.api_key = api_key
|
||||
self.timeout = 30
|
||||
self.base_url = "https://api.namecheap.com/xml.response"
|
||||
|
||||
def _get_hosts(self, domain_name) -> list:
|
||||
@@ -375,22 +363,20 @@ class NameCheapDns(BaseDns):
|
||||
resp = requests.get(url=self.base_url, params=params, timeout=self.timeout)
|
||||
if resp.status_code != 200:
|
||||
self.raise_resp_error(resp)
|
||||
import xml.etree.ElementTree as EtTree
|
||||
from xml.etree.ElementTree import ParseError as ETParseError
|
||||
hosts = []
|
||||
index = 0
|
||||
tree_root = resp.text.replace('xmlns="http://api.namecheap.com/xml.response"', '')
|
||||
try:
|
||||
hosts_info = EtTree.fromstring(tree_root).findall(".//host")
|
||||
except ETParseError:
|
||||
hosts_info = []
|
||||
|
||||
hosts_info = self._generate_xml_tree(resp.text, ".//host")
|
||||
for host in hosts_info:
|
||||
index += 1
|
||||
try:
|
||||
ttl = int(host.get("TTL", 1))
|
||||
except Exception:
|
||||
ttl = 1
|
||||
hosts.append({
|
||||
f"HostName{index}": host.get("Name"),
|
||||
f"RecordType{index}": host.get("Type"),
|
||||
f"Address{index}": host.get("Address"),
|
||||
f"TTL{index}": ttl,
|
||||
})
|
||||
return hosts
|
||||
|
||||
@@ -420,6 +406,7 @@ class NameCheapDns(BaseDns):
|
||||
self.raise_resp_error(setHosts_resp)
|
||||
|
||||
def create_dns_record(self, domain_name, domain_dns_value):
|
||||
# acme 调用
|
||||
domain_name = domain_name.lstrip("*.")
|
||||
root, _, acme_txt = extract_zone(domain_name)
|
||||
if self._type != 0:
|
||||
@@ -429,32 +416,8 @@ class NameCheapDns(BaseDns):
|
||||
s_type = "TXT"
|
||||
return self.add_record(root, s_type, acme_txt, domain_dns_value)
|
||||
|
||||
def remove_record(self, domain_name, dns_name, s_type="TXT"):
|
||||
hosts_info = self._get_hosts(domain_name)
|
||||
new_hosts = []
|
||||
for host in hosts_info:
|
||||
if dns_name in host.values() and s_type in host.values():
|
||||
continue
|
||||
else:
|
||||
new_hosts.append(host)
|
||||
if new_hosts:
|
||||
new_params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
"Command": "namecheap.domains.dns.setHosts",
|
||||
"SLD": domain_name.split(".")[0],
|
||||
"TLD": domain_name.split(".")[1],
|
||||
"DomainName": domain_name,
|
||||
}
|
||||
for host in new_hosts:
|
||||
new_params.update(host)
|
||||
setHosts_resp = requests.get(url=self.base_url, params=new_params, timeout=self.timeout)
|
||||
if setHosts_resp.status_code != 200:
|
||||
self.raise_resp_error(setHosts_resp)
|
||||
|
||||
def delete_dns_record(self, domain_name, dns_value=None):
|
||||
# 移除挑战值
|
||||
domain_name = domain_name.lstrip("*.")
|
||||
dns_name = "_acme-challenge" + "." + domain_name
|
||||
self.remove_record(domain_name, dns_name, 'TXT')
|
||||
@@ -464,55 +427,209 @@ class NameCheapDns(BaseDns):
|
||||
root, zone, _ = extract_zone(domain)
|
||||
self.add_record(root, "A", zone, server_ip)
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict):
|
||||
api_user = conf_data.get("Account")
|
||||
api_key = conf_data.get("ApiKey")
|
||||
if api_user is None or api_key is None:
|
||||
raise Exception(public.lang("Account, ApiKey not found"))
|
||||
return cls(api_user, api_key)
|
||||
# =============== 域名管理 ====================
|
||||
@staticmethod
|
||||
def _generate_xml_tree(resp_body: str, findall: str):
|
||||
import xml.etree.ElementTree as EtTree
|
||||
from xml.etree.ElementTree import ParseError as ETParseError
|
||||
tree_root = resp_body.replace('xmlns="http://api.namecheap.com/xml.response"', '')
|
||||
try:
|
||||
targets = EtTree.fromstring(tree_root).findall(findall)
|
||||
return targets
|
||||
except ETParseError:
|
||||
return []
|
||||
|
||||
def get_domains(self) -> list:
|
||||
# 获取账号下所有域名, 判断域名nameserver归属, 并且所有返回均为xml
|
||||
params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
"Command": "namecheap.domains.getList", # returns a list of domains for the particular user
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
}
|
||||
resp = requests.get(url=self.base_url, params=params, timeout=self.timeout)
|
||||
if resp.status_code != 200:
|
||||
return []
|
||||
domains = self._generate_xml_tree(resp.text, ".//Domain")
|
||||
domains = [
|
||||
domain.get("Name") for domain in domains if domain.get("IsExpired") == "false"
|
||||
]
|
||||
res = []
|
||||
for d in domains:
|
||||
try:
|
||||
params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
# gets a list of DNS servers associated with the requested domain.
|
||||
"Command": "namecheap.domains.dns.getList",
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
"SLD": d.split(".")[0],
|
||||
"TLD": d.split(".")[1],
|
||||
}
|
||||
resp = requests.get(url=self.base_url, params=params, timeout=self.timeout)
|
||||
if resp.status_code != 200:
|
||||
continue
|
||||
tree = self._generate_xml_tree(resp.text, ".//DomainDNSGetListResult")
|
||||
for t in tree:
|
||||
if t.get("Domain") == d and t.get("IsUsingOurDNS") == "true":
|
||||
res.append(d)
|
||||
break
|
||||
time.sleep(1)
|
||||
except Exception as e:
|
||||
public.print_log(f"get_domains error {e}")
|
||||
continue
|
||||
return res
|
||||
|
||||
def get_dns_record(self, domain_name):
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
try:
|
||||
records = self._get_hosts(domain_name)
|
||||
except Exception as e:
|
||||
public.print_log(f"get_dns_record error {e}")
|
||||
records = []
|
||||
res = []
|
||||
for r in records:
|
||||
temp = {}
|
||||
for k, v in r.items():
|
||||
if k.startswith("HostName"):
|
||||
temp["record"] = v
|
||||
elif k.startswith("RecordType"):
|
||||
temp["record_type"] = v
|
||||
elif k.startswith("Address"):
|
||||
temp["record_value"] = v
|
||||
elif k.startswith("TTL"):
|
||||
temp["ttl"] = r.get("ttl", 1)
|
||||
else:
|
||||
temp[k] = v
|
||||
res.append(temp)
|
||||
return res
|
||||
|
||||
def __set_hosts_with_params(self, domain_name: str, new_params: dict):
|
||||
try:
|
||||
setHosts_resp = requests.get(url=self.base_url, params=new_params, timeout=self.timeout)
|
||||
except Exception as e:
|
||||
return {"status": False, "msg": str(e)}
|
||||
if any([
|
||||
setHosts_resp.status_code != 200,
|
||||
f'Domain="{domain_name}" IsSuccess="true"' not in setHosts_resp.text
|
||||
]):
|
||||
return {"status": False, "msg": setHosts_resp.text}
|
||||
return {"status": True, "msg": setHosts_resp.text}
|
||||
|
||||
# 创建record
|
||||
def create_org_record(self, domain_name, record, record_value, record_type, ttl=1, **kwargs):
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
hosts = self._get_hosts(domain_name)
|
||||
add_index = len(hosts) + 1
|
||||
params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
"Command": "namecheap.domains.dns.setHosts",
|
||||
"SLD": domain_name.split(".")[0],
|
||||
"TLD": domain_name.split(".")[1],
|
||||
"DomainName": domain_name,
|
||||
}
|
||||
for index, host in enumerate(hosts):
|
||||
params[f"HostName{index + 1}"] = host[f"HostName{index + 1}"]
|
||||
params[f"Address{index + 1}"] = host[f"Address{index + 1}"]
|
||||
params[f"RecordType{index + 1}"] = host[f"RecordType{index + 1}"]
|
||||
params[f"TTL{index + 1}"] = host[f"TTL{index + 1}"]
|
||||
|
||||
params[f"HostName{add_index}"] = record
|
||||
params[f"Address{add_index}"] = record_value
|
||||
params[f"RecordType{add_index}"] = record_type
|
||||
params[f"TTL{add_index}"] = ttl
|
||||
return self.__set_hosts_with_params(domain_name, params)
|
||||
|
||||
# 删除record
|
||||
def remove_record(self, domain_name, record, record_type="TXT") -> dict:
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
hosts_info = self._get_hosts(domain_name)
|
||||
new_hosts = []
|
||||
for host in hosts_info:
|
||||
if record in host.values() and record_type in host.values():
|
||||
continue
|
||||
else:
|
||||
new_hosts.append(host)
|
||||
if not new_hosts:
|
||||
# is empty
|
||||
return {"status": True, "msg": "Dns Record is empty."}
|
||||
new_params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
"Command": "namecheap.domains.dns.setHosts",
|
||||
"SLD": domain_name.split(".")[0],
|
||||
"TLD": domain_name.split(".")[1],
|
||||
"DomainName": domain_name,
|
||||
}
|
||||
for host in new_hosts:
|
||||
new_params.update(host)
|
||||
return self.__set_hosts_with_params(domain_name, new_params)
|
||||
|
||||
# 更新record
|
||||
def update_record(self, domain_name, record: dict, new_record: dict, **kwargs):
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
hosts_info = self._get_hosts(domain_name)
|
||||
new_hosts = []
|
||||
for index, host in enumerate(hosts_info):
|
||||
if all([
|
||||
record.get("record") in host.values(),
|
||||
record.get("record_type") in host.values(),
|
||||
record.get("record_value") in host.values(),
|
||||
]):
|
||||
host[f"HostName{index + 1}"] = new_record.get("record")
|
||||
host[f"RecordType{index + 1}"] = new_record.get("record_type")
|
||||
host[f"Address{index + 1}"] = new_record.get("record_value")
|
||||
host[f"TTL{index + 1}"] = kwargs.get("ttl", 1)
|
||||
new_hosts.append(host)
|
||||
else:
|
||||
new_hosts.append(host)
|
||||
if not new_hosts: # is empty
|
||||
return {"status": True, "msg": "Dns Record is empty."}
|
||||
new_params = {
|
||||
"ApiUser": self.api_user,
|
||||
"ApiKey": self.api_key,
|
||||
"UserName": self.api_user,
|
||||
"ClientIp": public.GetLocalIp(),
|
||||
"Command": "namecheap.domains.dns.setHosts",
|
||||
"SLD": domain_name.split(".")[0],
|
||||
"TLD": domain_name.split(".")[1],
|
||||
"DomainName": domain_name,
|
||||
}
|
||||
for host in new_hosts:
|
||||
new_params.update(host)
|
||||
return self.__set_hosts_with_params(domain_name, new_params)
|
||||
|
||||
|
||||
class CloudFlareDns(BaseDns):
|
||||
dns_provider_name = "cloudflare"
|
||||
_type = 0 # 0:lest 1:锐成
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
CLOUDFLARE_EMAIL,
|
||||
CLOUDFLARE_API_KEY,
|
||||
CLOUDFLARE_API_BASE_URL="https://api.cloudflare.com/client/v4/",
|
||||
):
|
||||
self.CLOUDFLARE_DNS_ZONE_ID = None
|
||||
self.CLOUDFLARE_EMAIL = CLOUDFLARE_EMAIL
|
||||
self.CLOUDFLARE_API_KEY = CLOUDFLARE_API_KEY
|
||||
self.CLOUDFLARE_API_BASE_URL = CLOUDFLARE_API_BASE_URL
|
||||
self.HTTP_TIMEOUT = 65 # seconds
|
||||
|
||||
try:
|
||||
import urllib.parse as urlparse
|
||||
except:
|
||||
import urlparse
|
||||
|
||||
if CLOUDFLARE_API_BASE_URL[-1] != "/":
|
||||
self.CLOUDFLARE_API_BASE_URL = CLOUDFLARE_API_BASE_URL + "/"
|
||||
else:
|
||||
self.CLOUDFLARE_API_BASE_URL = CLOUDFLARE_API_BASE_URL
|
||||
super(CloudFlareDns, self).__init__()
|
||||
def __init__(self, api_user, api_key, limit: bool = True, **kwargs):
|
||||
super().__init__()
|
||||
self.cf_zone_id = None
|
||||
self.api_user = api_user
|
||||
self.api_key = api_key
|
||||
self.limit = limit
|
||||
self.cf_base_url = "https://api.cloudflare.com/client/v4/"
|
||||
self.time_out = 65 # seconds
|
||||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
# api limit True
|
||||
if os.path.exists('/www/server/panel/data/cf_limit_api.pl'):
|
||||
return {"Authorization": "Bearer " + self.CLOUDFLARE_API_KEY}
|
||||
# if self.CLOUDFLARE_EMAIL is None and isinstance(self.CLOUDFLARE_API_KEY, str):
|
||||
# return
|
||||
else: # api limit False
|
||||
return {"X-Auth-Email": self.CLOUDFLARE_EMAIL, "X-Auth-Key": self.CLOUDFLARE_API_KEY}
|
||||
if self.limit is True:
|
||||
return {"Authorization": "Bearer " + self.api_key}
|
||||
else: # api limit False, is global permissions
|
||||
return {"X-Auth-Email": self.api_user, "X-Auth-Key": self.api_key}
|
||||
|
||||
def find_dns_zone(self, domain_name):
|
||||
url = urljoin(self.CLOUDFLARE_API_BASE_URL, "zones?status=active&per_page=1000")
|
||||
url = self.cf_base_url + "zones?status=active&per_page=1000"
|
||||
headers = self._get_auth_headers()
|
||||
find_dns_zone_response = requests.get(url, headers=headers, timeout=self.HTTP_TIMEOUT)
|
||||
find_dns_zone_response = requests.get(url, headers=headers, timeout=self.time_out)
|
||||
if find_dns_zone_response.status_code != 200:
|
||||
raise ValueError(
|
||||
"Error creating cloudflare dns record: status_code={status_code} response={response}".format(
|
||||
@@ -524,8 +641,8 @@ class CloudFlareDns(BaseDns):
|
||||
result = find_dns_zone_response.json()["result"]
|
||||
for i in result:
|
||||
if i["name"] in domain_name:
|
||||
setattr(self, "CLOUDFLARE_DNS_ZONE_ID", i["id"])
|
||||
if isinstance(self.CLOUDFLARE_DNS_ZONE_ID, type(None)):
|
||||
setattr(self, "cf_zone_id", i["id"])
|
||||
if isinstance(self.cf_zone_id, type(None)):
|
||||
raise ValueError(
|
||||
"Error unable to get DNS zone for domain_name={domain_name}: status_code={status_code} response={response}".format(
|
||||
domain_name=domain_name,
|
||||
@@ -535,9 +652,10 @@ class CloudFlareDns(BaseDns):
|
||||
)
|
||||
|
||||
def add_record(self, domain_name, value, s_type):
|
||||
self.find_dns_zone(domain_name)
|
||||
url = urljoin(
|
||||
self.CLOUDFLARE_API_BASE_URL,
|
||||
"zones/{0}/dns_records".format(self.CLOUDFLARE_DNS_ZONE_ID),
|
||||
self.cf_base_url,
|
||||
"zones/{0}/dns_records".format(self.cf_zone_id),
|
||||
)
|
||||
headers = self._get_auth_headers()
|
||||
body = {
|
||||
@@ -547,17 +665,18 @@ class CloudFlareDns(BaseDns):
|
||||
}
|
||||
|
||||
create_resp = requests.post(
|
||||
url, headers=headers, json=body, timeout=self.HTTP_TIMEOUT
|
||||
url, headers=headers, json=body, timeout=self.time_out
|
||||
)
|
||||
if create_resp.status_code != 200:
|
||||
self.raise_resp_error(create_resp)
|
||||
|
||||
def create_dns_record(self, domain_name, domain_dns_value):
|
||||
# acme 调用
|
||||
domain_name = domain_name.lstrip("*.")
|
||||
self.find_dns_zone(domain_name)
|
||||
url = urljoin(
|
||||
self.CLOUDFLARE_API_BASE_URL,
|
||||
"zones/{0}/dns_records".format(self.CLOUDFLARE_DNS_ZONE_ID),
|
||||
self.cf_base_url,
|
||||
"zones/{0}/dns_records".format(self.cf_zone_id),
|
||||
)
|
||||
headers = self._get_auth_headers()
|
||||
body = {
|
||||
@@ -571,7 +690,7 @@ class CloudFlareDns(BaseDns):
|
||||
body['name'] = acme_txt.replace('_acme-challenge.', '')
|
||||
|
||||
create_cloudflare_dns_record_response = requests.post(
|
||||
url, headers=headers, json=body, timeout=self.HTTP_TIMEOUT
|
||||
url, headers=headers, json=body, timeout=self.time_out
|
||||
)
|
||||
if create_cloudflare_dns_record_response.status_code != 200:
|
||||
# raise error so that we do not continue to make calls to ACME
|
||||
@@ -583,30 +702,8 @@ class CloudFlareDns(BaseDns):
|
||||
)
|
||||
)
|
||||
|
||||
def remove_record(self, domain_name, dns_name, s_type):
|
||||
headers = self._get_auth_headers()
|
||||
list_dns_payload = {"type": s_type, "name": dns_name}
|
||||
list_dns_url = urljoin(
|
||||
self.CLOUDFLARE_API_BASE_URL,
|
||||
"zones/{0}/dns_records".format(self.CLOUDFLARE_DNS_ZONE_ID),
|
||||
)
|
||||
|
||||
list_dns_response = requests.get(
|
||||
list_dns_url, params=list_dns_payload, headers=headers, timeout=self.HTTP_TIMEOUT
|
||||
)
|
||||
|
||||
for i in range(0, len(list_dns_response.json()["result"])):
|
||||
dns_record_id = list_dns_response.json()["result"][i]["id"]
|
||||
url = urljoin(
|
||||
self.CLOUDFLARE_API_BASE_URL,
|
||||
"zones/{0}/dns_records/{1}".format(self.CLOUDFLARE_DNS_ZONE_ID, dns_record_id),
|
||||
)
|
||||
headers = self._get_auth_headers()
|
||||
requests.delete(
|
||||
url, headers=headers, timeout=self.HTTP_TIMEOUT
|
||||
)
|
||||
|
||||
def delete_dns_record(self, domain_name, domain_dns_value):
|
||||
# 移除挑战值
|
||||
domain_name = domain_name.lstrip("*.")
|
||||
dns_name = "_acme-challenge" + "." + domain_name
|
||||
self.remove_record(domain_name, dns_name, 'TXT')
|
||||
@@ -616,133 +713,142 @@ class CloudFlareDns(BaseDns):
|
||||
self.find_dns_zone(domain_name)
|
||||
self.add_record(zone, server_ip, "A")
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict) -> BaseDns:
|
||||
key = conf_data.get("key", None) or conf_data.get("E-Mail", None) or conf_data.get("E-MAIL", None)
|
||||
secret = conf_data.get("secret", None) or conf_data.get("API Key", None) or conf_data.get("API KEY", None)
|
||||
base_url = "https://api.cloudflare.com/client/v4/"
|
||||
# =============== 域名管理 ====================
|
||||
|
||||
if key is None and secret is None:
|
||||
secret = conf_data.get("API Token", None) # 处理api - token的情况
|
||||
if key is None and secret is None:
|
||||
raise Exception(public.lang("api key, api secret not found"))
|
||||
return cls(key, secret, base_url)
|
||||
def get_domains(self) -> list:
|
||||
url = self.cf_base_url + "zones?status=active&per_page=1000"
|
||||
headers = self._get_auth_headers()
|
||||
res = requests.get(url, headers=headers, timeout=self.time_out)
|
||||
if res.status_code != 200:
|
||||
return []
|
||||
try:
|
||||
result = res.json().get("result", [])
|
||||
except Exception as e:
|
||||
public.print_log(f"cloudflare get_domains error {e}")
|
||||
result = []
|
||||
return [i.get("name", "") for i in result]
|
||||
|
||||
def get_dns_record(self, domain_name: str) -> list:
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
self.find_dns_zone(domain_name)
|
||||
url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records"
|
||||
result = []
|
||||
page = 1
|
||||
per_page = 500
|
||||
fail_count = 0
|
||||
while True:
|
||||
params = {"page": page, "per_page": per_page}
|
||||
try:
|
||||
response = requests.get(
|
||||
url, headers=self._get_auth_headers(), params=params
|
||||
)
|
||||
data = response.json()
|
||||
if data.get("success"):
|
||||
records = data.get("result", [])
|
||||
result.extend([
|
||||
{
|
||||
"record": i.get("name", ""),
|
||||
"record_value": i.get("content", ""),
|
||||
"record_type": i.get("type", ""),
|
||||
"proxy": i.get("proxied", False),
|
||||
"ttl": i.get("ttl", 1),
|
||||
} for i in records
|
||||
])
|
||||
if len(records) < per_page:
|
||||
break
|
||||
page += 1
|
||||
else:
|
||||
fail_count += 1
|
||||
if fail_count >= 3:
|
||||
break
|
||||
except requests.RequestException as e:
|
||||
print("get_dns_record error", e)
|
||||
break
|
||||
return result
|
||||
|
||||
# 官方不支持
|
||||
class GoDaddyDns(BaseDns):
|
||||
_type = 0 # 0:lest 1:锐成
|
||||
http_timeout = 65
|
||||
_debug = False
|
||||
|
||||
def __init__(self, sso_key: str, sso_secret: str, base_url='https://api.godaddy.com'):
|
||||
self.sso_key = sso_key
|
||||
self.sso_secret = sso_secret
|
||||
# self.base_url = "https://api.ote-godaddy.com/"
|
||||
self.base_url = base_url
|
||||
super(GoDaddyDns, self).__init__()
|
||||
self._headers = None
|
||||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
if self._headers is not None:
|
||||
return self._headers
|
||||
self._headers = {
|
||||
"Authorization": "sso-key {}:{}".format(self.sso_key, self.sso_secret)
|
||||
# 创建record
|
||||
def create_org_record(self, domain_name, record, record_value, record_type, ttl=1, proxied=0):
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
proxied = True if proxied == 1 else False
|
||||
self.find_dns_zone(domain_name)
|
||||
url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records"
|
||||
headers = self._get_auth_headers()
|
||||
body = {
|
||||
"content": record_value,
|
||||
"name": record,
|
||||
"proxied": proxied,
|
||||
"ttl": ttl,
|
||||
"type": record_type
|
||||
}
|
||||
return self._headers
|
||||
try:
|
||||
create_res = requests.post(url, headers=headers, json=body, timeout=self.time_out)
|
||||
create_res = create_res.json()
|
||||
if create_res.get("success"):
|
||||
return {"status": True, "msg": create_res}
|
||||
return {"status": False, "msg": str(create_res.get("errors"))}
|
||||
except requests.exceptions.HTTPError as http_err:
|
||||
return {"status": False, "msg": http_err}
|
||||
except Exception as e:
|
||||
return {"status": False, "msg": str(e)}
|
||||
|
||||
def create_dns_record(self, domain_name, domain_dns_value):
|
||||
domain_name = domain_name.lstrip("*.")
|
||||
root, zone, acme_txt = extract_zone(domain_name)
|
||||
|
||||
url = urljoin(
|
||||
self.base_url,
|
||||
"/v1/domains/{}/records".format(root),
|
||||
)
|
||||
# 删除record
|
||||
def remove_record(self, domain_name, record, record_type="TXT") -> dict:
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
self.find_dns_zone(domain_name)
|
||||
headers = self._get_auth_headers()
|
||||
body = [
|
||||
{
|
||||
"data": domain_dns_value,
|
||||
"name": acme_txt,
|
||||
"type": "TXT",
|
||||
}
|
||||
]
|
||||
if self._type == 1:
|
||||
body[0]['type'] = 'CNAME'
|
||||
root, _, acme_txt = extract_zone(domain_name)
|
||||
body[0]['name'] = acme_txt.replace('_acme-challenge.', '')
|
||||
|
||||
create_godaday_dns_record_response = requests.patch(
|
||||
url, headers=headers, json=body, timeout=self.http_timeout
|
||||
list_dns_payload = {"type": record_type, "name": record}
|
||||
list_dns_url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records"
|
||||
list_dns_response = requests.get(
|
||||
list_dns_url, params=list_dns_payload, headers=headers, timeout=self.time_out
|
||||
)
|
||||
if create_godaday_dns_record_response.status_code != 200:
|
||||
# raise error so that we do not continue to make calls to ACME
|
||||
# server
|
||||
raise ValueError(
|
||||
"Error creating GoDaddyDns dns record: status_code={status_code} response={response}".format(
|
||||
status_code=create_godaday_dns_record_response.status_code,
|
||||
response=self.log_response(create_godaday_dns_record_response),
|
||||
)
|
||||
)
|
||||
try:
|
||||
for i in range(0, len(list_dns_response.json()["result"])):
|
||||
dns_record_id = list_dns_response.json()["result"][i]["id"]
|
||||
url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records/{dns_record_id}"
|
||||
remove_res = requests.delete(url, headers=headers, timeout=self.time_out)
|
||||
remove_res = remove_res.json()
|
||||
if remove_res.get("success"):
|
||||
return {"status": True, "msg": remove_res}
|
||||
return {"status": False, "msg": str(remove_res.get("errors"))}
|
||||
# is empty
|
||||
return {"status": True, "msg": "Dns Record is empty."}
|
||||
except requests.exceptions.HTTPError as http_err:
|
||||
return {"status": False, "msg": http_err}
|
||||
except Exception as e:
|
||||
return {"status": False, "msg": str(e)}
|
||||
|
||||
def add_record(self, root, host, value, s_type):
|
||||
url = urljoin(
|
||||
self.base_url,
|
||||
"/v1/domains/{}/records".format(root),
|
||||
)
|
||||
headers = self._get_auth_headers()
|
||||
|
||||
body = [{
|
||||
"type": s_type,
|
||||
"name": host,
|
||||
"data": "{0}".format(value),
|
||||
}]
|
||||
|
||||
create_cloudflare_dns_record_response = requests.patch(
|
||||
url, headers=headers, json=body, timeout=self.http_timeout
|
||||
)
|
||||
if create_cloudflare_dns_record_response.status_code != 200:
|
||||
raise ValueError(
|
||||
"Error creating cloudflare dns record: status_code={status_code} response={response}".format(
|
||||
status_code=create_cloudflare_dns_record_response.status_code,
|
||||
response=self.log_response(create_cloudflare_dns_record_response),
|
||||
)
|
||||
)
|
||||
|
||||
def remove_record(self, domain_name, dns_name, s_type):
|
||||
headers = self._get_auth_headers()
|
||||
|
||||
list_dns_url = urljoin(
|
||||
self.base_url,
|
||||
"/v1/domains/{}/records/{}/{}".format(domain_name, s_type, dns_name),
|
||||
)
|
||||
|
||||
dns_response = requests.delete(
|
||||
list_dns_url, headers=headers, timeout=self.http_timeout
|
||||
)
|
||||
if dns_response.status_code != 200:
|
||||
raise ValueError(
|
||||
"Error creating cloudflare dns record: status_code={status_code} response={response}".format(
|
||||
status_code=dns_response.status_code,
|
||||
response=self.log_response(dns_response),
|
||||
)
|
||||
)
|
||||
|
||||
def add_record_for_creat_site(self, domain, server_ip):
|
||||
root, zone, _ = extract_zone(domain)
|
||||
self.add_record(root, zone, server_ip, "A")
|
||||
|
||||
def delete_dns_record(self, domain_name, domain_dns_value):
|
||||
root, zone, acme_txt = extract_zone(domain_name)
|
||||
self.remove_record(root, acme_txt, 'TXT')
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict) -> BaseDns:
|
||||
key = conf_data.get("key", None) or conf_data.get("Key", "")
|
||||
secret = conf_data.get("secret", None) or conf_data.get("Secret", "")
|
||||
base_url = "https://api.godaddy.com"
|
||||
|
||||
return cls(key, secret, base_url)
|
||||
# 更新record
|
||||
def update_record(self, domain_name, record: dict, new_record: dict, **kwargs):
|
||||
domain_name, _, _ = extract_zone(domain_name)
|
||||
self.find_dns_zone(domain_name)
|
||||
record_type = record.get("record_type")
|
||||
record_name = record.get("record")
|
||||
get_url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records?type={record_type}&name={record_name}"
|
||||
try:
|
||||
# get record id
|
||||
get_response = requests.get(get_url, headers=self._get_auth_headers())
|
||||
get_result = get_response.json()
|
||||
if get_result.get("success") and get_result.get("result"):
|
||||
record_id = get_result["result"][0]["id"]
|
||||
update_url = self.cf_base_url + f"zones/{self.cf_zone_id}/dns_records/{record_id}"
|
||||
body = {
|
||||
"type": new_record.get("record_type"),
|
||||
"name": new_record.get("record"),
|
||||
"content": new_record.get("record_value"),
|
||||
"ttl": new_record.get("ttl", 1),
|
||||
"proxied": True if new_record.get("proxy") == 1 else False,
|
||||
}
|
||||
update_response = requests.put(update_url, headers=self._get_auth_headers(), json=body)
|
||||
update_result = update_response.json()
|
||||
if update_result.get("success"):
|
||||
return {"status": True, "msg": update_result}
|
||||
return {"status": False, "msg": str(update_result.get("errors"))}
|
||||
else:
|
||||
return {"status": False, "msg": "Dns Record Not Found!"}
|
||||
except requests.exceptions.HTTPError as http_err:
|
||||
return {"status": False, "msg": http_err}
|
||||
except Exception as err:
|
||||
return {"status": False, "msg": err}
|
||||
|
||||
|
||||
# 未验证
|
||||
@@ -875,13 +981,6 @@ class AliyunDns(object):
|
||||
root, zone, _ = extract_zone(domain)
|
||||
self.add_record(root, "A", zone, server_ip)
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data: dict) -> "AliyunDns":
|
||||
key = conf_data.get("key", None) or conf_data.get("AccessKey", "")
|
||||
secret = conf_data.get("secret", None) or conf_data.get("SecretKey", "")
|
||||
|
||||
return cls(key, secret)
|
||||
|
||||
|
||||
# 未验证
|
||||
class CloudxnsDns(object):
|
||||
@@ -1007,13 +1106,6 @@ class DNSLADns(BaseDns):
|
||||
self.domain_list = None
|
||||
super(DNSLADns, self).__init__()
|
||||
|
||||
@classmethod
|
||||
def new(cls, conf_data) -> BaseDns:
|
||||
key = conf_data.get("key", None) or conf_data.get("APIID", "")
|
||||
secret = conf_data.get("secret", None) or conf_data.get("API密钥", "")
|
||||
|
||||
return cls(key, secret)
|
||||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
if self._token is None:
|
||||
self._token = base64.b64encode("{}:{}".format(self.api_id, self.api_secret).encode("utf-8")).decode("utf-8")
|
||||
@@ -1190,7 +1282,7 @@ class DnsMager(object):
|
||||
"AliyunDns": AliyunDns,
|
||||
"DNSPodDns": DNSPodDns,
|
||||
"CloudFlareDns": CloudFlareDns,
|
||||
"GoDaddyDns": GoDaddyDns,
|
||||
# "GoDaddyDns": GoDaddyDns,
|
||||
"DNSLADns": DNSLADns,
|
||||
"HuaweiCloudDns": HuaweiCloudDns,
|
||||
"TencentCloudDns": TencentCloudDns,
|
||||
@@ -1273,7 +1365,7 @@ class DnsMager(object):
|
||||
for rule_name, rule in rule_map.items():
|
||||
tmp = {}
|
||||
for r_key, r_value in rule.items():
|
||||
account_res = re.search(r_value + "\s*=\s*'(.+)'", account)
|
||||
account_res = re.search(r_value + r"\s*=\s*'(.+)'", account)
|
||||
if account_res:
|
||||
tmp[r_key] = account_res.groups()[0]
|
||||
|
||||
|
||||
+13
-5
@@ -1111,7 +1111,7 @@ class panelSSL:
|
||||
if "ssl_hash" in Info:
|
||||
get.ssl_hash = Info['ssl_hash']
|
||||
result = self.SetCertToSite(get)
|
||||
if not result:
|
||||
if not result or result.get("status") is False:
|
||||
set_result['status'] = False
|
||||
failnum += 1
|
||||
faildList.append(set_result)
|
||||
@@ -1174,6 +1174,9 @@ class panelSSL:
|
||||
public.serviceReload()
|
||||
return public.return_msg_gettext(True, public.lang("Setup successfully!"))
|
||||
except Exception as ex:
|
||||
import traceback
|
||||
public.print_log(traceback.format_exc())
|
||||
public.print_log(f"error : {ex}")
|
||||
if 'isBatch' in get: return False
|
||||
return public.returnMsg(False, 'SET_ERROR,' + public.get_error_info())
|
||||
|
||||
@@ -1224,10 +1227,15 @@ class panelSSL:
|
||||
# 读取证书
|
||||
def GetCert(self, get):
|
||||
vpath = os.path.join('/www/server/panel/vhost/ssl', get.certName.replace("*.", ''))
|
||||
if not os.path.exists(vpath): return public.return_msg_gettext(False, public.lang("Certificate does NOT exist!"))
|
||||
data = {}
|
||||
data['privkey'] = public.readFile(vpath + '/privkey.pem')
|
||||
data['fullchain'] = public.readFile(vpath + '/fullchain.pem')
|
||||
if not os.path.exists(vpath):
|
||||
return public.return_msg_gettext(False, public.lang("Certificate does NOT exist!"))
|
||||
# vpath = os.path.join('/www/server/panel/vhost/ssl_saved', get.ssl_hash)
|
||||
# if not os.path.exists(vpath):
|
||||
# return public.return_msg_gettext(False, public.lang("Certificate does NOT exist!"))
|
||||
data = {
|
||||
'privkey': public.readFile(vpath + '/privkey.pem'),
|
||||
'fullchain': public.readFile(vpath + '/fullchain.pem'),
|
||||
}
|
||||
return data
|
||||
|
||||
# 获取证书名称
|
||||
|
||||
+4
-4
@@ -1713,12 +1713,12 @@ listener Default%s{
|
||||
if result != apis:
|
||||
public.writeFile('./config/dns_api.json', json.dumps(result))
|
||||
|
||||
for index, item in enumerate(result):
|
||||
for index, item in enumerate(apis):
|
||||
if item.get("title", "") == "Manual resolution":
|
||||
target_dict = result.pop(index)
|
||||
result.insert(0, target_dict)
|
||||
target_dict = apis.pop(index)
|
||||
apis.insert(0, target_dict)
|
||||
break
|
||||
return result
|
||||
return apis
|
||||
|
||||
# 设置DNS-API
|
||||
def SetDnsApi(self, get):
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
# coding: utf-8
|
||||
|
||||
from .fields import *
|
||||
from .model import aaModel
|
||||
from .manager import Q
|
||||
__version__ = "1.1.0"
|
||||
|
||||
__all__ = [
|
||||
"aaModel", "Q"
|
||||
] + fields.__all__
|
||||
@@ -0,0 +1,362 @@
|
||||
# coding: utf-8
|
||||
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field as dataclass_field
|
||||
from datetime import datetime
|
||||
from typing import Any, TypeVar, List, Optional
|
||||
|
||||
from public.exceptions import HintException
|
||||
|
||||
__all__ = [
|
||||
"StrField",
|
||||
"IntField",
|
||||
"FloatField",
|
||||
"BlobField",
|
||||
"BoolField",
|
||||
"ListField",
|
||||
"DictField",
|
||||
"DateTimeStrField",
|
||||
]
|
||||
|
||||
M = TypeVar("M", bound="aaModel")
|
||||
|
||||
|
||||
def json_func(v_type: type, value: Any, forward: bool = True):
|
||||
try:
|
||||
if forward is True:
|
||||
if isinstance(value, v_type):
|
||||
return json.dumps(value)
|
||||
else:
|
||||
if isinstance(value, str):
|
||||
return json.loads(value)
|
||||
return value
|
||||
except TypeError as t:
|
||||
print("type error %s" % t)
|
||||
return value
|
||||
except Exception as e:
|
||||
print("error %s" % e)
|
||||
raise e
|
||||
|
||||
|
||||
@dataclass
|
||||
class aaField(object):
|
||||
"""
|
||||
字段基类
|
||||
default 默认值
|
||||
ps 字段说明
|
||||
null 是否null
|
||||
primary_key 是否主键
|
||||
foreign_key 外键
|
||||
field_name 字段key
|
||||
field_type sql类型
|
||||
py_type py类型
|
||||
compare 比较
|
||||
transform 转换工具
|
||||
"""
|
||||
default: Any = None
|
||||
ps: str = None
|
||||
null: bool = False
|
||||
primary_key: bool = False
|
||||
field_name: str = None
|
||||
field_type: str = None
|
||||
py_type: type = None
|
||||
compare: tuple = None
|
||||
serialized: Callable = None
|
||||
# todo
|
||||
# require: bool = False
|
||||
# foreign_key: str = None
|
||||
|
||||
def __set_name__(self, owner: M, name: str):
|
||||
self.__model = owner
|
||||
self.field_name = str(name)
|
||||
|
||||
def __get__(self, instance: object, owner):
|
||||
if instance is None:
|
||||
return self
|
||||
try:
|
||||
return instance.__dict__[self.field_name]
|
||||
except KeyError:
|
||||
raise AttributeError(f'{self.field_name} is not set')
|
||||
|
||||
def __set__(self, instance, value):
|
||||
instance.__dict__[self.field_name] = value
|
||||
|
||||
def __delete__(self, instance):
|
||||
try:
|
||||
del instance.__dict__[self.field_name]
|
||||
except KeyError:
|
||||
raise AttributeError(f"{instance} dont have attr '{self.field_name}'")
|
||||
|
||||
def _raise_error(self, raise_exp: bool = True) -> bool:
|
||||
if raise_exp is True:
|
||||
err = f"'{self.field_name}' TypeError! It should be '{self.py_type.__name__}'"
|
||||
raise HintException(err)
|
||||
else:
|
||||
return False
|
||||
|
||||
def _check_type(self, target: Any, raise_exp=True) -> bool:
|
||||
target = target if not isinstance(target, Callable) else target()
|
||||
flag = True if any([isinstance(target, x) for x in self.py_types]) else False
|
||||
return True if flag is True else self._raise_error(raise_exp=raise_exp)
|
||||
|
||||
@property
|
||||
def default_val_sql(self) -> str:
|
||||
default_v = self.get_default_val(check=True)
|
||||
default_v = default_v if self.serialized is None else self.serialized(default_v)
|
||||
return f"DEFAULT '{default_v}'" if isinstance(default_v, str) else f"DEFAULT {default_v}"
|
||||
|
||||
@property
|
||||
def py_types(self) -> List[type]:
|
||||
if self.null is True:
|
||||
original = [type(self.default), self.py_type, type(None)]
|
||||
else:
|
||||
original = [type(self.default), self.py_type]
|
||||
return list(set(original))
|
||||
|
||||
def get_default_val(self, check: bool = False) -> Optional[Any]:
|
||||
if check:
|
||||
self.check_org_type(raise_exp=check)
|
||||
if self.default is not None:
|
||||
return self.default if not isinstance(self.default, Callable) else self.default()
|
||||
else:
|
||||
if self.null is True:
|
||||
return None
|
||||
else:
|
||||
raise TypeError(
|
||||
f"\n1: field '{self.field_name}' is not null, must have a default value"
|
||||
f"\n2: you can add the '{self.field_name}' field's params null=True"
|
||||
)
|
||||
|
||||
def model_check_type(self, target: Any, raise_exp=True) -> bool:
|
||||
"""
|
||||
检查模型当前类型结构
|
||||
"""
|
||||
return self._check_type(target, raise_exp=raise_exp)
|
||||
|
||||
def check_org_type(self, raise_exp: bool = True) -> bool:
|
||||
"""
|
||||
检查初始化的类型结构
|
||||
"""
|
||||
return self._check_type(self.default, raise_exp=raise_exp)
|
||||
|
||||
|
||||
@dataclass()
|
||||
class StrField(aaField):
|
||||
"""
|
||||
String field
|
||||
|
||||
field__like="a",
|
||||
field__ne="a",
|
||||
field__in=["a", "b", "c"]
|
||||
field__not_in=["a", "b", "c"]
|
||||
field__startswith="a"
|
||||
field__endswith="a"
|
||||
"""
|
||||
default: str | None = ""
|
||||
field_type: str = "TEXT"
|
||||
py_type: type = str
|
||||
max_length: int = 255 # not limit now
|
||||
min_length: int = 0 # not limit now
|
||||
compare: tuple[str] = (
|
||||
"like",
|
||||
"ne",
|
||||
"in",
|
||||
"not_in",
|
||||
"startswith",
|
||||
"endswith",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class IntField(aaField):
|
||||
"""
|
||||
Int field
|
||||
|
||||
field__gt=1,
|
||||
field__gte=1,
|
||||
field__lt=1,
|
||||
field__lte=1,
|
||||
field__ne=1,
|
||||
field__in=[1, 2, 3]
|
||||
field__not_in=[1, 2, 3]
|
||||
"""
|
||||
default: int | None | Any = 0
|
||||
field_type: str = "INTEGER"
|
||||
py_type: type = int
|
||||
max: int = 0 # not limit now
|
||||
min: int = 0 # not limit now
|
||||
compare: tuple[str] = (
|
||||
"gt",
|
||||
"lt",
|
||||
"gte",
|
||||
"lte",
|
||||
"ne",
|
||||
"in",
|
||||
"not_in",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FloatField(aaField):
|
||||
"""
|
||||
Float field
|
||||
"""
|
||||
default: float | None = 0.0
|
||||
field_type: str = "REAL"
|
||||
py_type: type = float
|
||||
max: float = 0.0 # not limit now
|
||||
min: float = 0.0 # not limit now
|
||||
compare: tuple[str] = (
|
||||
"gt",
|
||||
"lt",
|
||||
"gte",
|
||||
"lte",
|
||||
"ne",
|
||||
"in",
|
||||
"not_in",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlobField(aaField):
|
||||
"""
|
||||
Blob field
|
||||
"""
|
||||
default: bytes | None = b''
|
||||
field_type: str = "BLOB"
|
||||
py_type: type = bytes
|
||||
|
||||
|
||||
@dataclass
|
||||
class BoolField(aaField):
|
||||
"""
|
||||
Bool field
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _serialized(value: bool | int, forward: bool = True) -> bool | int:
|
||||
if forward is True:
|
||||
if value is True:
|
||||
return 1
|
||||
elif value is False:
|
||||
return 0
|
||||
else:
|
||||
if value == 1:
|
||||
return True
|
||||
elif value == 0:
|
||||
return False
|
||||
return value
|
||||
|
||||
default: bool = True
|
||||
field_type: str = "INTEGER"
|
||||
serialized: Callable = _serialized
|
||||
py_type: type = bool
|
||||
|
||||
|
||||
@dataclass
|
||||
class ListField(aaField):
|
||||
"""
|
||||
List field
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _serialized(value: list | str, forward: bool = True) -> list | Any:
|
||||
return json_func(list, value, forward)
|
||||
|
||||
default: list = dataclass_field(default_factory=list)
|
||||
field_type: str = "TEXT"
|
||||
serialized: Callable = _serialized
|
||||
py_type: type = list
|
||||
compare: tuple[str] = (
|
||||
"has_element",
|
||||
)
|
||||
update: tuple[str] = (
|
||||
"append",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DictField(aaField):
|
||||
"""
|
||||
Dict field
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _serialized(value: dict | str, forward: bool = True) -> dict | Any:
|
||||
return json_func(dict, value, forward)
|
||||
|
||||
default: dict = dataclass_field(default_factory=dict)
|
||||
field_type: str = "TEXT"
|
||||
serialized: Callable = _serialized
|
||||
py_type: type = dict
|
||||
compare: tuple[str] = (
|
||||
"has_key",
|
||||
"has_value",
|
||||
# "has_key_value",
|
||||
)
|
||||
update: tuple[str] = (
|
||||
"update",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DateTimeStrField(aaField):
|
||||
"""
|
||||
时间戳
|
||||
auto_now_add=True 创建时间自动添加
|
||||
auto_now=True 更新时间自动更新
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _current_timestamp(cls):
|
||||
return time.strftime(cls.format, time.localtime())
|
||||
|
||||
@staticmethod
|
||||
def _dynamic(obj, val):
|
||||
if hasattr(obj, "auto_now_add") and obj.auto_now_add is True: # 创建时间
|
||||
return val
|
||||
elif hasattr(obj, "auto_now") and obj.auto_now is True: # 更新时间
|
||||
return obj.get_default_val()
|
||||
else:
|
||||
return val
|
||||
|
||||
@staticmethod
|
||||
def _serialized(value: int | str, forward: bool = True) -> str | int:
|
||||
try:
|
||||
if forward is True:
|
||||
if isinstance(value, str): # save will be str
|
||||
return int(
|
||||
datetime.strptime(value, DateTimeStrField.format).timestamp() * DateTimeStrField.accuracy)
|
||||
else:
|
||||
if isinstance(value, int):
|
||||
return datetime.fromtimestamp(value / DateTimeStrField.accuracy).strftime(DateTimeStrField.format)
|
||||
return value
|
||||
except Exception as e:
|
||||
print("type error %s" % e)
|
||||
return value
|
||||
|
||||
serialized: Callable = _serialized
|
||||
default: str | Callable = ""
|
||||
field_type: str = "INTEGER"
|
||||
py_type: type = str
|
||||
dynamic: bool = True
|
||||
auto_now_add: bool = False
|
||||
auto_now: bool = False
|
||||
accuracy: int = 1000
|
||||
format: str = "%Y-%m-%d %H:%M:%S"
|
||||
compare: tuple[str] = (
|
||||
"gt",
|
||||
"lt",
|
||||
"gte",
|
||||
"lte",
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
if self.auto_now_add is True and self.auto_now is True:
|
||||
raise TypeError("auto_now_add and auto_now can not be used at the same time")
|
||||
if self.auto_now is True:
|
||||
self.default = self._current_timestamp
|
||||
elif self.auto_now_add is True:
|
||||
self.default = self._current_timestamp
|
||||
@@ -0,0 +1,634 @@
|
||||
# coding: utf-8
|
||||
|
||||
import uuid
|
||||
from functools import reduce
|
||||
from typing import Optional, TypeVar, Generic, Any, List, Dict, Generator, Tuple
|
||||
|
||||
from public.exceptions import HintException, PanelError
|
||||
from public.sqlite_easy import Db
|
||||
|
||||
__all__ = ["aaManager", "Q"]
|
||||
|
||||
M = TypeVar("M", bound="aaModel")
|
||||
|
||||
|
||||
class Operator:
|
||||
def __init__(self, model_class: M, query: Db.query):
|
||||
self._model_class: M = model_class
|
||||
self._query = query
|
||||
self._tb = self._model_class.__table_name__
|
||||
self._fields = self._model_class.__fields__
|
||||
self._serializes = self._model_class.__serializes__
|
||||
|
||||
def _q_error(self, key: str, act: str, val: Any, sp_act: tuple):
|
||||
raise HintException(
|
||||
"field: '%s' is not support '%s', you can use: %s" % (key, act, sp_act)
|
||||
)
|
||||
|
||||
def __check_js_type(self, key: str, val: Any) -> None:
|
||||
for t in [str, int]:
|
||||
if isinstance(val, t):
|
||||
break
|
||||
else:
|
||||
raise HintException(
|
||||
"%s's fields %s is not support compare value type: %s" % (self._model_class.__name__, key, val)
|
||||
)
|
||||
|
||||
# todo
|
||||
def __json_operator(self, key: str, compare: str, val: Any, sp_compare: tuple):
|
||||
as_name = f"{self._tb}_{key}_je"
|
||||
operate_map = {
|
||||
# ListField
|
||||
"has_element": f"{as_name}.value = {val}",
|
||||
# DictField
|
||||
"has_key": f"{as_name}.key = '{val}'",
|
||||
"has_value": f"{as_name}.value = {val}",
|
||||
}
|
||||
if operate_map.get(compare):
|
||||
self.__check_js_type(key, val)
|
||||
# self._query.join(
|
||||
# f"json_each({self._tb}.{key}) AS {as_name}",
|
||||
# operate_map.get(compare),
|
||||
# )
|
||||
self._query.join(
|
||||
f"json_each({self._tb}.{key}) AS {as_name}", operate_map.get(compare)
|
||||
)
|
||||
return "", []
|
||||
else:
|
||||
self._q_error(key, compare, val, sp_compare)
|
||||
|
||||
def __compare_operator(self, key: str, compare: str, val: Any, sp_compare: tuple):
|
||||
def trans(v):
|
||||
try:
|
||||
return f"({v[0]})" if isinstance(v, list) and len(v) == 1 else tuple(v)
|
||||
except TypeError:
|
||||
return v
|
||||
|
||||
operate = {
|
||||
"like": (f"{key} LIKE ?", f"%{val}%"),
|
||||
"gt": (f"{key} > ?", val),
|
||||
"lt": (f"{key} < ?", val),
|
||||
"gte": (f"{key} >= ?", val),
|
||||
"lte": (f"{key} <= ?", val),
|
||||
"ne": (f"{key} != ?", val),
|
||||
"in": (f'{key} IN {trans(val)}', ()),
|
||||
"not_in": (f"{key} NOT IN {trans(val)}", ()),
|
||||
"startswith": (f"{key} LIKE ?", f"{val}%"),
|
||||
"endswith": (f"{key} LIKE ?", f"%{val}")
|
||||
}
|
||||
if operate.get(compare):
|
||||
return operate[compare]
|
||||
else:
|
||||
self._q_error(key, compare, val, sp_compare)
|
||||
|
||||
def __compare_reducer(self, key: str, compare: str, val: Any):
|
||||
sp_compare = getattr(self._fields.get(key), "compare")
|
||||
if sp_compare and compare in sp_compare:
|
||||
if not (self._serializes and self._serializes.get(key)):
|
||||
return self.__compare_operator(key, compare.lower(), val, sp_compare)
|
||||
elif hasattr(self._fields.get(key), "dynamic") and self._serializes.get(key):
|
||||
real_val = self._serializes[key].serialized(value=val, forward=True)
|
||||
return self.__compare_operator(key, compare.lower(), real_val, sp_compare)
|
||||
else:
|
||||
self._q_error(key, compare, val, ())
|
||||
# return self.__json_operator(key, compare.lower(), val, sp_compare)
|
||||
else:
|
||||
self._q_error(key, compare, val, sp_compare)
|
||||
|
||||
def __equal_reducer(self, key: str, val: Any):
|
||||
if self._serializes and key in self._serializes:
|
||||
val = self._serializes[key].serialized(value=val, forward=True)
|
||||
if val is not None:
|
||||
return f"{self._tb}.{key} = ?", [val]
|
||||
else:
|
||||
return f"{self._tb}.{key} IS NULL", []
|
||||
|
||||
def __split_condition(self, condition: Dict[str, Any]):
|
||||
for k, v in condition.items():
|
||||
key, compare = (k, None) if "__" not in k else k.split("__")
|
||||
if not key or not self._fields.get(key):
|
||||
print("Filter: %s's fields is not found: '%s' it will be pass" % (self._model_class.__name__, k))
|
||||
# raise AttributeError("%s's fields is not found: '%s'" % (self._model_class.__name__, k))
|
||||
else:
|
||||
yield key, compare, v
|
||||
|
||||
def reducer_process(self, condition: Dict[str, Any]) -> Tuple[str, tuple]:
|
||||
for key, compare, v in self.__split_condition(condition):
|
||||
if not compare:
|
||||
sql, params = self.__equal_reducer(key=key, val=v)
|
||||
else:
|
||||
if v is None:
|
||||
raise HintException("do not try to use 'None' value to compare.")
|
||||
sql, params = self.__compare_reducer(key=key, compare=compare, val=v)
|
||||
if sql:
|
||||
yield sql, params
|
||||
|
||||
|
||||
class Q:
|
||||
"""
|
||||
嵌套查询
|
||||
AND优先级大于OR, 括号改变优先级
|
||||
example: model.object.filter( Q(a=1) & (Q(b=2) | Q(c=3)) )
|
||||
"""
|
||||
AND = "AND"
|
||||
OR = "OR"
|
||||
|
||||
def __init__(self, *args, _connector=None, **kwargs):
|
||||
self.children = []
|
||||
if args:
|
||||
self.children.extend(args) # Q
|
||||
if kwargs:
|
||||
self.children.append(kwargs) # conditions
|
||||
self._connector = _connector or self.AND
|
||||
|
||||
def __and__(self, other):
|
||||
return Q(self, other, _connector=Q.AND)
|
||||
|
||||
def __or__(self, other):
|
||||
return Q(self, other, _connector=Q.OR)
|
||||
|
||||
def resolve(self, operator, query):
|
||||
for child in self.children:
|
||||
with query.where_nest(logic=self._connector) as n:
|
||||
if not isinstance(child, Q):
|
||||
for s, p in operator.reducer_process(child):
|
||||
if s:
|
||||
n.where(s, p)
|
||||
else:
|
||||
child.resolve(operator, n)
|
||||
|
||||
|
||||
class QuerySet(Generic[M]):
|
||||
"""
|
||||
查询集
|
||||
"""
|
||||
|
||||
def __init__(self, model_class: M, query: Db.query):
|
||||
self._model_class: M = model_class
|
||||
self._tb = self._model_class.__table_name__
|
||||
self._query = query
|
||||
self._cache = None
|
||||
|
||||
def __len__(self):
|
||||
return len(self.__execute())
|
||||
|
||||
def __bool__(self):
|
||||
return bool(self.__len__())
|
||||
|
||||
def __iter__(self) -> Generator[M, None, None]:
|
||||
yield from self.__execute()
|
||||
|
||||
def __getitem__(self, index: Optional[int | slice]) -> Optional[M | List[M]]:
|
||||
"""
|
||||
查询结果切片
|
||||
"""
|
||||
if isinstance(index, int):
|
||||
if self._cache is not None:
|
||||
try:
|
||||
return self._cache[index]
|
||||
except IndexError:
|
||||
raise HintException("list index out of range")
|
||||
else:
|
||||
new_q = self._query.fork()
|
||||
new_q.limit(index + 1, index)
|
||||
temp = new_q.find()
|
||||
return self._model_class(
|
||||
**self._model_class._serialized_data(temp)
|
||||
) if temp else None
|
||||
elif isinstance(index, slice):
|
||||
if self._cache is not None:
|
||||
try:
|
||||
return self._cache[index.start:index.stop]
|
||||
except IndexError:
|
||||
raise HintException("list index out of range")
|
||||
else:
|
||||
self.__execute()
|
||||
return self._cache[index.start:index.stop]
|
||||
else:
|
||||
return None
|
||||
|
||||
def __execute(self) -> Optional[List[M]]:
|
||||
if self._cache is None:
|
||||
try:
|
||||
if len(self._query._SqliteEasy__OPT_FIELD._Field__FIELDS) == 0:
|
||||
self._query.field(f"'{self._tb}'.*")
|
||||
self._cache = [
|
||||
self._model_class(
|
||||
**self._model_class._serialized_data(i)
|
||||
) for i in self._query.select()
|
||||
]
|
||||
return self._cache
|
||||
except Exception as e:
|
||||
print("db query error => %s" % str(e))
|
||||
raise PanelError(e)
|
||||
else:
|
||||
return self._cache
|
||||
|
||||
def filter(self, *args: Dict[str, Any] | Q, **kwargs: Dict[str, Any]) -> "QuerySet[M]":
|
||||
"""
|
||||
过滤
|
||||
:return: QuerySet
|
||||
"""
|
||||
operator = Operator(model_class=self._model_class, query=self._query)
|
||||
# args
|
||||
for i in args:
|
||||
if isinstance(i, Q):
|
||||
i.resolve(operator, self._query)
|
||||
elif isinstance(i, dict):
|
||||
for s, p in operator.reducer_process(i):
|
||||
if s:
|
||||
self._query.where(s, p)
|
||||
else:
|
||||
raise HintException("args: '%s' must be 'dict'" % args)
|
||||
# kwargs
|
||||
for s, p in operator.reducer_process(kwargs):
|
||||
if s:
|
||||
self._query.where(s, p)
|
||||
return self
|
||||
|
||||
def limit(self, num: int) -> "QuerySet[M]":
|
||||
"""
|
||||
限制
|
||||
:return: QuerySet
|
||||
"""
|
||||
self._query.limit(num)
|
||||
return self
|
||||
|
||||
def offset(self, num: int) -> "QuerySet[M]":
|
||||
"""
|
||||
偏移量
|
||||
:return: QuerySet
|
||||
"""
|
||||
self._query.skip(num)
|
||||
return self
|
||||
|
||||
def distinct(self) -> "QuerySet[M]":
|
||||
"""
|
||||
以指定字段去重
|
||||
"""
|
||||
return self
|
||||
|
||||
def order_by(self, *args) -> "QuerySet[M]":
|
||||
"""
|
||||
排序
|
||||
:param args: "filed" ASC "-filed" DESC
|
||||
:return: QuerySet
|
||||
"""
|
||||
reduce(
|
||||
lambda q, c: q.order(f"{self._tb}.{c[1:]}", "DESC") if c[:1] == "-"
|
||||
else q.order(f"{self._tb}.{c}"), args, self._query
|
||||
)
|
||||
return self
|
||||
|
||||
def fields(self, *args) -> "QuerySet[M]":
|
||||
self._query.field(*(f"{self._tb}.{c}" for c in args))
|
||||
# todo 序列化后出现初始值
|
||||
return self
|
||||
|
||||
def first(self) -> Optional[M]:
|
||||
"""
|
||||
获取第一条数据
|
||||
:return: dict
|
||||
"""
|
||||
if self._cache is None:
|
||||
new_q = self._query.fork()
|
||||
if len(new_q._SqliteEasy__OPT_FIELD._Field__FIELDS) == 0:
|
||||
new_q.field(f"'{self._tb}'.*")
|
||||
data = new_q.find()
|
||||
if not data:
|
||||
return None
|
||||
return self._model_class(
|
||||
**self._model_class._serialized_data(data)
|
||||
)
|
||||
else:
|
||||
return self._cache[0] if len(self._cache) != 0 else None
|
||||
|
||||
def get_field(self, key_name: str) -> Optional[Any]:
|
||||
"""
|
||||
获取第一条数据的指定字段的值
|
||||
:param key_name: 字段名
|
||||
:return: Any
|
||||
"""
|
||||
# noinspection PyUnresolvedReferences
|
||||
f = self.first().as_dict()
|
||||
return f.get(key_name) if f else None
|
||||
|
||||
def update(self, *args, **kwargs) -> int:
|
||||
"""
|
||||
更新数据
|
||||
:return: int
|
||||
"""
|
||||
self._cache = None
|
||||
if args and kwargs:
|
||||
raise HintException("args and kwargs can not be used at the same time")
|
||||
if args:
|
||||
if len(args) != 1:
|
||||
raise HintException("%s too many args" % (args,))
|
||||
elif not isinstance(args[0], dict):
|
||||
raise HintException("%s must be a dict" % (args[0],))
|
||||
target = args[0]
|
||||
elif kwargs:
|
||||
target = kwargs
|
||||
else:
|
||||
target = None
|
||||
if not target:
|
||||
return 0
|
||||
count = 0
|
||||
for i in self._model_class._serialized_data(self._query.select()):
|
||||
self._model_class(**{**i, **target}).save()
|
||||
count += 1
|
||||
return count
|
||||
|
||||
def delete(self) -> int:
|
||||
"""
|
||||
删除数据
|
||||
:return: int
|
||||
"""
|
||||
self._cache = None
|
||||
count = self._query.delete()
|
||||
return count
|
||||
|
||||
def exists(self) -> bool:
|
||||
"""
|
||||
存在数据
|
||||
:return: bool
|
||||
"""
|
||||
return bool(self.__execute())
|
||||
|
||||
def count(self) -> int:
|
||||
"""
|
||||
获取数量
|
||||
:return: int
|
||||
"""
|
||||
return self._query.fork().count()
|
||||
|
||||
def as_list(self) -> list:
|
||||
"""
|
||||
转列表
|
||||
:return: list
|
||||
"""
|
||||
if self._cache is None:
|
||||
self.__execute()
|
||||
return [x.as_dict() for x in self._cache]
|
||||
|
||||
|
||||
class aaObjects(Generic[M]):
|
||||
"""
|
||||
管理器
|
||||
"""
|
||||
__m_map__ = {}
|
||||
|
||||
def __new__(cls, args):
|
||||
if hasattr(args, "__table_name__") and not cls.__m_map__.get(args.__table_name__):
|
||||
cls.__m_map__[args.__table_name__] = aaMigrate(args).run_migrate()
|
||||
return super(aaObjects, cls).__new__(cls)
|
||||
else:
|
||||
return super(aaObjects, cls).__new__(cls)
|
||||
|
||||
def __init__(self, model: M):
|
||||
self._model = model
|
||||
self.__q = None
|
||||
|
||||
@property
|
||||
def _query(self) -> Db.query:
|
||||
if not self.__q:
|
||||
q = Db(self._model.__db_name__).query()
|
||||
self.__q = q.table(self._model.__table_name__)
|
||||
return self.__q
|
||||
else:
|
||||
return self.__q
|
||||
|
||||
def _insert(self, val_data) -> int:
|
||||
return self._query.insert(val_data)
|
||||
|
||||
def insert(self, data: Dict[str, Any], raise_exp: bool = True) -> dict:
|
||||
"""
|
||||
插入单条数据
|
||||
:data dict
|
||||
:raise_exp bool 抛字段类型检查异常
|
||||
:return 插入的数据
|
||||
"""
|
||||
model_obj = self._model(**data)
|
||||
insert_res = self._insert(
|
||||
model_obj._validate(raise_exp=raise_exp)
|
||||
)
|
||||
if insert_res:
|
||||
return {
|
||||
self._model.__primary_key__: insert_res, **model_obj.as_dict()
|
||||
}
|
||||
else:
|
||||
if raise_exp:
|
||||
raise HintException(insert_res)
|
||||
else:
|
||||
return {}
|
||||
|
||||
def insert_many(self, data: List[Dict[str, Any]], raise_exp: bool = True) -> int:
|
||||
"""
|
||||
批量插入数据
|
||||
:data list
|
||||
:raise_exp bool 不抛异常则跳过异常继续插入
|
||||
:return: int 影响行数
|
||||
"""
|
||||
for i in data:
|
||||
try:
|
||||
self._model(**i)._validate(raise_exp=raise_exp)
|
||||
except Exception as e:
|
||||
if raise_exp:
|
||||
raise HintException(f"{e}, data: {i}")
|
||||
else:
|
||||
data.remove(i)
|
||||
continue
|
||||
return self._query.insert_all(data)
|
||||
|
||||
def find_one(self, **kwargs) -> Optional[M]:
|
||||
"""
|
||||
过滤查询一行数据
|
||||
:kwargs dict
|
||||
:return: QuerySet | None
|
||||
"""
|
||||
# noinspection PyUnresolvedReferences
|
||||
return QuerySet(self._model, self._query).filter(**kwargs).first()
|
||||
|
||||
def filter(self, *args, **kwargs) -> "QuerySet[M]":
|
||||
"""
|
||||
过滤
|
||||
:kwargs dict
|
||||
:return: QuerySet
|
||||
"""
|
||||
return QuerySet(self._model, self._query).filter(*args, **kwargs)
|
||||
|
||||
def all(self) -> "QuerySet[M]":
|
||||
"""
|
||||
所有数据
|
||||
return: QuerySet
|
||||
"""
|
||||
return QuerySet(self._model, self._query)
|
||||
|
||||
|
||||
class aaMigrate:
|
||||
"""
|
||||
同步表字段
|
||||
"""
|
||||
NULL_MAP = {False: "NOT NULL", True: "NULL"}
|
||||
|
||||
def __init__(self, model: M):
|
||||
self.__model = model
|
||||
self.__table = self.__model.__table_name__
|
||||
self.__fields = self.__model.__fields__
|
||||
self.__client = Db(model.__db_name__)
|
||||
self.__query = None
|
||||
|
||||
def run_migrate(self) -> bool:
|
||||
"""
|
||||
迁移
|
||||
"""
|
||||
if not self.__model:
|
||||
raise PanelError("Model is None")
|
||||
if not hasattr(self.__model, '__db_name__'):
|
||||
raise PanelError(f"{self.__model.__class__.__name__} need 'db_name'")
|
||||
if not hasattr(self.__model, '__table_name__'):
|
||||
raise PanelError(f"{self.__model.__class__.__name__} need 'table_name'")
|
||||
if not hasattr(self.__model, '__fields__'):
|
||||
raise PanelError(f"{self.__model.__class__.__name__} need 'fields'")
|
||||
|
||||
try:
|
||||
self.__table_exists()
|
||||
self.__index_exists()
|
||||
except Exception as e:
|
||||
raise PanelError(e)
|
||||
finally:
|
||||
self.__query.close()
|
||||
self.__client.close()
|
||||
return True
|
||||
|
||||
def __new_tb_transform_sql(self, tb_name: str) -> str:
|
||||
"""
|
||||
转sql
|
||||
"""
|
||||
field_sql = ""
|
||||
pk_flag = 0
|
||||
for key, val in self.__fields.items():
|
||||
if key == "index":
|
||||
raise PanelError("'%s' is a reserved word in SQL. do not use it" % key)
|
||||
if val.primary_key is False:
|
||||
field_sql += f"`{key}` {val.field_type} {self.NULL_MAP.get(val.null)} {val.default_val_sql}, "
|
||||
else: # is primary_key
|
||||
pk_flag += 1
|
||||
if val.field_type != "INTEGER":
|
||||
raise PanelError("'primary_key' only support IntegerField now")
|
||||
field_sql += f"'{key}' {val.field_type} PRIMARY KEY AUTOINCREMENT, "
|
||||
if not field_sql:
|
||||
return ""
|
||||
if pk_flag != 1:
|
||||
raise PanelError("primary_key not found, and must be only one")
|
||||
field_sql = field_sql.rstrip(", ")
|
||||
sql = f"""CREATE TABLE IF NOT EXISTS `{tb_name}` ({field_sql});"""
|
||||
return sql
|
||||
|
||||
def __fields_exist(self, add_fields_map: dict = None, del_fields: set = None) -> None:
|
||||
"""
|
||||
字段处理
|
||||
"""
|
||||
if not del_fields:
|
||||
for k, v in add_fields_map.items():
|
||||
add_sql = (f"ALTER TABLE `{self.__table}` "
|
||||
f"ADD COLUMN `{k}` {v.field_type} {v.default_val_sql} {self.NULL_MAP.get(v.null)};")
|
||||
self.__query.execute(add_sql)
|
||||
else:
|
||||
temp_tb = f"table_{uuid.uuid4().hex}"
|
||||
new = self.__new_tb_transform_sql(temp_tb)
|
||||
if new:
|
||||
try:
|
||||
self.__query.autocommit(autocommit=False)
|
||||
self.__query.execute("BEGIN;")
|
||||
self.__query.execute(new)
|
||||
new_keys = tuple(self.__model.__fields__.keys())
|
||||
copy_sql = (f"INSERT INTO `{temp_tb}` {new_keys} "
|
||||
f"SELECT {', '.join(new_keys)} FROM `{self.__table}`;")
|
||||
self.__query.execute(copy_sql)
|
||||
self.__query.execute(f"DROP TABLE `{self.__table}`;")
|
||||
self.__query.execute(f"ALTER TABLE `{temp_tb}` RENAME TO `{self.__table}`;")
|
||||
self.__query.commit()
|
||||
except Exception as e:
|
||||
import traceback
|
||||
print(traceback.format_exc())
|
||||
self.__query.rollback()
|
||||
raise e
|
||||
|
||||
def __table_exists(self) -> None:
|
||||
"""
|
||||
表迁移
|
||||
"""
|
||||
self.__query = self.__client.query().table("sqlite_master")
|
||||
if self.__query.where("type=? AND name=?", ("table", self.__table)).count() != 1:
|
||||
sql = self.__new_tb_transform_sql(self.__table)
|
||||
if sql:
|
||||
self.__query.execute(sql)
|
||||
else: # has table
|
||||
self.__query.table(self.__table)
|
||||
set_cur = set(self.__fields.keys())
|
||||
set_db = set(self.__query.get_columns())
|
||||
add_fields = set_cur - set_db
|
||||
del_fields = set_db - set_cur
|
||||
add_fields_map = {k: v for k, v in self.__fields.items() if k in add_fields}
|
||||
self.__fields_exist(add_fields_map, del_fields)
|
||||
|
||||
def __trans_index_key(self, index_info: tuple | str) -> str:
|
||||
def __if_raise_error(item: str):
|
||||
if not self.__model.__fields__.get(item):
|
||||
raise PanelError(f"create index error, '{item}' is not in model's fields")
|
||||
|
||||
col_sql = ""
|
||||
if isinstance(index_info, tuple):
|
||||
for item in index_info:
|
||||
__if_raise_error(item)
|
||||
col_sql += f"`{item}`,"
|
||||
elif isinstance(index_info, str):
|
||||
__if_raise_error(index_info)
|
||||
col_sql = f"`{index_info}`"
|
||||
else:
|
||||
raise PanelError("model's index error, should be like ['key1', ('key2', 'key3')]")
|
||||
return "(" + col_sql.rstrip(",") + ")"
|
||||
|
||||
def __index_exists(self) -> bool:
|
||||
"""
|
||||
索引
|
||||
"""
|
||||
try:
|
||||
self.__query.table(self.__table)
|
||||
cur = self.__query.query(f"PRAGMA index_list(`{self.__table}`);")
|
||||
if cur:
|
||||
current_index = [x.get("name") for x in cur]
|
||||
else:
|
||||
return False
|
||||
sql = ""
|
||||
for index_info in self.__model.__index_keys__:
|
||||
col_sql = self.__trans_index_key(index_info)
|
||||
index_name = f"idx_{self.__table}_{'_'.join([col.strip('` ') for col in col_sql.strip('()').split(',')])}"
|
||||
if index_name not in current_index:
|
||||
idx_sql = f"CREATE INDEX IF NOT EXISTS `{index_name}` ON `{self.__table}` {col_sql}; "
|
||||
sql += idx_sql
|
||||
else:
|
||||
current_index.remove(index_name)
|
||||
drop = '; '.join([f"DROP INDEX IF EXISTS `{index}`" for index in current_index]) + ';'
|
||||
if drop != ";":
|
||||
sql += drop
|
||||
if sql:
|
||||
self.__query.execute_script(sql)
|
||||
return True
|
||||
except Exception:
|
||||
import traceback
|
||||
print(traceback.format_exc())
|
||||
|
||||
|
||||
class aaManager:
|
||||
def __get__(self, instance, cls: M):
|
||||
if instance is None:
|
||||
try:
|
||||
return aaObjects(cls)
|
||||
except Exception:
|
||||
import traceback
|
||||
raise PanelError(traceback.format_exc())
|
||||
raise RuntimeError(
|
||||
f"object manager can't accessible from '{cls.__name__}' instances"
|
||||
)
|
||||
@@ -0,0 +1,260 @@
|
||||
# coding: utf-8
|
||||
|
||||
import copy
|
||||
from typing import Self, Generator, Optional, Dict, Any
|
||||
|
||||
from .fields import aaField
|
||||
from .manager import aaManager
|
||||
|
||||
__all__ = ["aaModel"]
|
||||
|
||||
from public.exceptions import HintException
|
||||
|
||||
|
||||
def generate_table_name(class_name: str) -> str:
|
||||
"""
|
||||
驼峰名转表名
|
||||
"""
|
||||
return ''.join(['_' + c.lower() if c.isupper() else c for c in class_name]).lstrip('_')
|
||||
|
||||
|
||||
class aaMetaClass(type):
|
||||
__abstract__: bool
|
||||
__db_name__: str
|
||||
__table_name__: str
|
||||
__fields__: dict
|
||||
__primary_key__: str
|
||||
__serializes__: dict
|
||||
__index_keys__: list
|
||||
|
||||
# __foreign_keys__: object
|
||||
|
||||
def __new__(cls, name, bases, attrs):
|
||||
if attrs.get("__abstract__") is True:
|
||||
return super().__new__(cls, name, bases, attrs)
|
||||
attrs.update({"__abstract__": False})
|
||||
new_class = super().__new__(cls, name, bases, attrs)
|
||||
cls.__fields_process(obj=new_class, name=name, attrs=attrs)
|
||||
cls.__database_process(obj=new_class, name=name, attrs=attrs)
|
||||
return new_class
|
||||
|
||||
def __setattr__(cls, key, value):
|
||||
if key == '__abstract__':
|
||||
raise AttributeError("can't set attribute '__abstract__'")
|
||||
return super().__setattr__(key, value)
|
||||
|
||||
@classmethod
|
||||
def __fields_process(cls, obj: "aaMetaClass", name: str, attrs: dict):
|
||||
pk = ""
|
||||
fields = {}
|
||||
for k, v in attrs.items():
|
||||
if isinstance(v, aaField):
|
||||
if k in fields:
|
||||
raise HintException(f"model {name} field '{k}' is already defined")
|
||||
fields[k] = v
|
||||
if v.primary_key:
|
||||
if pk:
|
||||
raise HintException(f"model {name} can only have one primary key")
|
||||
else:
|
||||
pk = k
|
||||
if not pk:
|
||||
raise HintException(f"sth wrong with {name}'s primary key, please check the model")
|
||||
setattr(obj, "__primary_key__", pk)
|
||||
setattr(obj, "__fields__", fields)
|
||||
setattr(obj, "__serializes__", cls.__get_serialized(fields))
|
||||
|
||||
@classmethod
|
||||
def __database_process(cls, obj: "aaMetaClass", name: str, attrs: dict):
|
||||
db_name, tb_name, idx = "default", generate_table_name(name), []
|
||||
meta = attrs.get("_Meta")
|
||||
if meta:
|
||||
if hasattr(meta, "db_name"):
|
||||
db_name = meta.db_name
|
||||
if hasattr(meta, "table_name"):
|
||||
tb_name = meta.table_name
|
||||
if hasattr(meta, "index"):
|
||||
idx = meta.index
|
||||
setattr(obj, "__db_name__", db_name)
|
||||
setattr(obj, "__table_name__", tb_name)
|
||||
setattr(obj, "__index_keys__", idx)
|
||||
|
||||
@staticmethod
|
||||
def __get_serialized(fields: Dict[str, aaField]) -> Dict[str, aaField]:
|
||||
return dict(filter(
|
||||
lambda x: x[1].serialized is not None, {**fields}.items()
|
||||
))
|
||||
|
||||
|
||||
class aaCusModel(metaclass=aaMetaClass):
|
||||
__abstract__ = True
|
||||
objects = aaManager()
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
if self.__abstract__ is True:
|
||||
raise RuntimeError(f'{self.__class__.__name__} class can not be init')
|
||||
for f, v in self._generate_init(kwargs):
|
||||
setattr(self, f.field_name, v)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<'{self.__class__.__name__}' Model Object, {self.__dict__}>"
|
||||
|
||||
def _generate_init(self, val_data: dict) -> Generator:
|
||||
for name, field in {**self.__class__.__fields__}.items():
|
||||
val = val_data.pop(name) if name in val_data else field.get_default_val()
|
||||
if field.primary_key is True and val == 0:
|
||||
continue # skip default id val
|
||||
if field.primary_key is True and val != 0:
|
||||
try:
|
||||
val = int(val)
|
||||
except Exception:
|
||||
pass
|
||||
yield field, val
|
||||
if val_data: # other field
|
||||
# raise AttributeError(f"model '{self.__class__.__name__}' has no field {val_data}")
|
||||
pass
|
||||
|
||||
|
||||
class aaModel(aaCusModel):
|
||||
"""
|
||||
基础模型
|
||||
|
||||
:example:
|
||||
class MyTestModel(aaModel):
|
||||
id = IntField(primary_key=True)
|
||||
name = StrField(ps="名字")
|
||||
status = BoolField(default=True, ps="状态")
|
||||
float_number = FloatField(default=0.05, ps="浮点")
|
||||
|
||||
class _Meta:
|
||||
db_name = "default" 默认为 default.db 文件
|
||||
table_name = "my_table" 默认为类名驼峰转表名 my_test_model
|
||||
index = [("status", 1)] 索引
|
||||
|
||||
"""
|
||||
__abstract__: bool = True
|
||||
__destroyed: bool = False
|
||||
id: int = None
|
||||
|
||||
@staticmethod
|
||||
def check_destroyed(func):
|
||||
def wrapper(self, *args, **kwargs):
|
||||
if getattr(self, "__destroyed", False):
|
||||
raise RuntimeError(f"Cannot call {func.__name__}() on destroyed object")
|
||||
return func(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
@classmethod
|
||||
@check_destroyed
|
||||
def _output(cls, data: dict) -> dict:
|
||||
return {
|
||||
k: v if not cls.__serializes__.get(k) else cls.__serializes__.get(k).serialized(v, False) for k, v in
|
||||
data.items()
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@check_destroyed
|
||||
def _serialized_data(cls, data: Optional[dict | list]) -> Optional[dict | list]:
|
||||
if not data or not hasattr(cls, "__serializes__"):
|
||||
return data
|
||||
if isinstance(data, list):
|
||||
return [cls._output(d) for d in data]
|
||||
elif isinstance(data, dict):
|
||||
return cls._output(data)
|
||||
else:
|
||||
return data
|
||||
|
||||
@check_destroyed
|
||||
def _validate(self, raise_exp: bool = True) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
模型验证
|
||||
"""
|
||||
body = {}
|
||||
for f, cur_val in self._generate_init(copy.deepcopy(self.__dict__)):
|
||||
try:
|
||||
f.model_check_type(target=cur_val, raise_exp=True)
|
||||
# 1, dynamic generated
|
||||
if hasattr(f, "dynamic") and f.dynamic is True:
|
||||
cur_val = f._dynamic(f, cur_val)
|
||||
setattr(self, f.field_name, cur_val)
|
||||
# 2, serialized
|
||||
body[f.field_name] = f.serialized(cur_val, True) if f.serialized else cur_val
|
||||
except TypeError as t:
|
||||
if raise_exp:
|
||||
raise t
|
||||
else:
|
||||
return {}
|
||||
except Exception as e:
|
||||
raise e
|
||||
return body
|
||||
|
||||
def _before_save(self) -> bool:
|
||||
# override
|
||||
return True
|
||||
|
||||
def _after_save(self) -> None:
|
||||
# override
|
||||
pass
|
||||
|
||||
@check_destroyed
|
||||
def save(self, raise_exp: bool = True) -> Optional[Self]:
|
||||
"""
|
||||
模型数据, 不存在则 保存 , 存在则 更新
|
||||
:raise_exp 抛异常
|
||||
:return: model object 字段类型异常等问题返回 None
|
||||
"""
|
||||
if self.__class__.__abstract__ is True:
|
||||
raise RuntimeError(f'{self.__class__.__name__} class can not be save')
|
||||
try:
|
||||
cls = self.__class__
|
||||
validate = self._validate(raise_exp=raise_exp)
|
||||
primary_key = cls.__primary_key__
|
||||
if validate and primary_key:
|
||||
if primary_key in validate:
|
||||
target_id = validate.pop(primary_key)
|
||||
exist = cls.objects._query.where(f"{primary_key}=?", (target_id,)).exists()
|
||||
if exist: # update
|
||||
res = cls.objects._query.where(f"{primary_key}=?", (target_id,)).update(validate)
|
||||
if res == 1:
|
||||
self._after_save()
|
||||
return self
|
||||
else:
|
||||
if raise_exp:
|
||||
raise HintException("update error")
|
||||
else:
|
||||
return None
|
||||
else:
|
||||
validate[primary_key] = target_id
|
||||
# save
|
||||
if not self._before_save():
|
||||
return None
|
||||
new_id = cls.objects._query.insert(validate)
|
||||
self._after_save()
|
||||
return cls(**{primary_key: new_id, **self.__dict__})
|
||||
return None
|
||||
except TypeError as t:
|
||||
if raise_exp:
|
||||
raise t
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
raise HintException(e)
|
||||
|
||||
@check_destroyed
|
||||
def delete(self) -> int:
|
||||
try:
|
||||
self.__class__.objects._query.where(
|
||||
f"{self.__class__.__primary_key__}=?", (self.id,)
|
||||
).delete()
|
||||
setattr(self, "__destroyed", True)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
return 0
|
||||
return 1
|
||||
|
||||
@check_destroyed
|
||||
def as_dict(self) -> dict:
|
||||
"""
|
||||
转字典
|
||||
"""
|
||||
return self.__dict__ if self.__dict__ is not None else {}
|
||||
+76
-53
@@ -16,6 +16,9 @@ from .tools import is_number
|
||||
from .structures import aap_t_simple_result, aap_t_mysql_dump_info, aap_t_http_multipart
|
||||
import gzip
|
||||
import fcntl
|
||||
import shutil
|
||||
import tempfile
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
|
||||
@@ -1131,72 +1134,63 @@ def read_file_each_reverse(filename: str, using_gzip: bool = False):
|
||||
if not os.path.exists(filename):
|
||||
raise ValueError(lang('file not found: {}', filename))
|
||||
|
||||
import shutil
|
||||
if filename.endswith('.gz'):
|
||||
using_gzip = True
|
||||
|
||||
using_tmp_file = False
|
||||
def _open_file():
|
||||
# if using gzip read file
|
||||
# decompress to tmp file
|
||||
if using_gzip:
|
||||
tmp_fp = tempfile.NamedTemporaryFile('wb+')
|
||||
|
||||
# if using gzip read file
|
||||
# decompress to tmp file
|
||||
if using_gzip:
|
||||
tmp_path = make_panel_tmp_path()
|
||||
tmp_file = '{}/{}'.format(tmp_path, GetRandomString(16))
|
||||
using_tmp_file = True
|
||||
|
||||
with open(tmp_file, 'wb') as fp:
|
||||
with gzip.open(filename, 'rb') as gz_fp:
|
||||
shutil.copyfileobj(gz_fp, fp)
|
||||
shutil.copyfileobj(gz_fp, tmp_fp)
|
||||
|
||||
filename = tmp_file
|
||||
return tmp_fp
|
||||
|
||||
try:
|
||||
with open(filename, 'rb') as fp:
|
||||
chunk_size = 4096
|
||||
end_pos = fp.seek(0, 2)
|
||||
loops = int(end_pos / chunk_size)
|
||||
last = b''
|
||||
i = 0
|
||||
return open(filename, 'rb')
|
||||
|
||||
while i < loops:
|
||||
fp.seek((chunk_size + chunk_size * i) * -1, 2)
|
||||
with _open_file() as fp:
|
||||
chunk_size = 4096
|
||||
end_pos = fp.seek(0, 2)
|
||||
loops = int(end_pos / chunk_size)
|
||||
last = b''
|
||||
i = 0
|
||||
|
||||
bs = fp.read(chunk_size)
|
||||
while i < loops:
|
||||
fp.seek(end_pos + (chunk_size + chunk_size * i) * -1)
|
||||
|
||||
lines = (bs + last).decode('utf-8', 'ignore').split('\n')
|
||||
last = lines[0].encode('utf-8', 'ignore')
|
||||
k = len(lines)
|
||||
bs = fp.read(chunk_size)
|
||||
|
||||
while k > 1:
|
||||
yield lines.pop()
|
||||
k -= 1
|
||||
lines = (bs + last).decode('utf-8', 'ignore').split('\n')
|
||||
last = lines[0].encode('utf-8', 'ignore')
|
||||
k = len(lines)
|
||||
|
||||
i += 1
|
||||
while k > 1:
|
||||
yield lines.pop()
|
||||
k -= 1
|
||||
|
||||
# once flow
|
||||
# handle remain rows
|
||||
for j in range(1):
|
||||
if i < loops:
|
||||
break
|
||||
i += 1
|
||||
|
||||
remainder = end_pos % chunk_size
|
||||
if i < loops:
|
||||
return
|
||||
|
||||
if remainder == 0:
|
||||
break
|
||||
remainder = end_pos % chunk_size
|
||||
|
||||
# move cursor to top
|
||||
fp.seek(0, 0)
|
||||
if remainder == 0:
|
||||
return
|
||||
|
||||
bs = fp.read(remainder)
|
||||
# move cursor to top
|
||||
fp.seek(0, 0)
|
||||
|
||||
lines = (bs + last).decode('utf-8', 'ignore').split('\n')
|
||||
last = b''
|
||||
k = len(lines)
|
||||
bs = fp.read(remainder)
|
||||
|
||||
while k > 0:
|
||||
yield lines.pop()
|
||||
k -= 1
|
||||
finally:
|
||||
if using_tmp_file:
|
||||
shutil.rmtree(os.path.dirname(filename))
|
||||
lines = (bs + last).decode('utf-8', 'ignore').split('\n')
|
||||
k = len(lines)
|
||||
|
||||
while k > 0:
|
||||
yield lines.pop()
|
||||
k -= 1
|
||||
|
||||
|
||||
# 验证证书
|
||||
@@ -4512,8 +4506,8 @@ class dict_obj:
|
||||
return hasattr(self, key)
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
if key in key_filter_list:
|
||||
raise ValueError("wrong field name")
|
||||
# if key in key_filter_list:
|
||||
# raise ValueError("wrong field name")
|
||||
|
||||
if not re_key_match.match(key) or re_key_match2.match(key):
|
||||
raise ValueError("wrong field name")
|
||||
@@ -4574,8 +4568,8 @@ class dict_obj:
|
||||
|
||||
def set(self, key, value):
|
||||
if not isinstance(value, str) or not isinstance(key, str): return False
|
||||
if key in key_filter_list:
|
||||
raise ValueError("wrong field name")
|
||||
# if key in key_filter_list:
|
||||
# raise ValueError("wrong field name")
|
||||
if not re_key_match.match(key) or re_key_match2.match(key):
|
||||
raise ValueError("wrong field name")
|
||||
return setattr(self, key, value)
|
||||
@@ -9137,3 +9131,32 @@ def snow_flake(machine_id: int = 0) -> int:
|
||||
fp.write(str(cur_snow_flake_time))
|
||||
|
||||
return (int(cur_snow_flake_time - 1739376000000) << 22) | ((int(machine_id) & ((1 << 10) - 1)) << 12) | (int(snow_flake_sequence) & ((1 << 12) - 1))
|
||||
|
||||
|
||||
# 根据时间区间生成日期字符串序列
|
||||
def gen_date_sequence_by_time_section(start_time: int = -1, end_time: int = -1, date_format: str = '%Y%m%d'):
|
||||
if end_time < 1:
|
||||
end_time = int(time.time())
|
||||
|
||||
if end_time < start_time:
|
||||
raise ValueError(lang('end_time must greater than start_time'))
|
||||
|
||||
for i in range(start_time, end_time + (end_time % 86400), 86400):
|
||||
yield datetime.fromtimestamp(i).strftime(date_format)
|
||||
|
||||
|
||||
def check_field_exists(db_obj,table_name, field_name):
|
||||
"""
|
||||
@name 检查表字段是否存在
|
||||
@param db_obj 数据库对象
|
||||
@param table_name 表名
|
||||
@param field_name 要检查的字段
|
||||
"""
|
||||
try:
|
||||
res = db_obj.query("PRAGMA table_info({})".format(table_name))
|
||||
for val in res:
|
||||
if field_name == val[1]:
|
||||
return True
|
||||
except:
|
||||
pass
|
||||
return False
|
||||
@@ -30,13 +30,18 @@ def hook_import():
|
||||
panel_path = public.get_panel_path()
|
||||
pyfile = '{}.py'.format(str(name).strip().replace('.', '/'))
|
||||
|
||||
realpath = os.path.join(panel_path, 'class', pyfile)
|
||||
cond = os.path.exists(realpath)
|
||||
realpath = ''
|
||||
cond = False
|
||||
|
||||
if not cond:
|
||||
realpath = os.path.join(panel_path, 'class_v2', pyfile)
|
||||
for p in set(sys.path):
|
||||
realpath = os.path.join(panel_path, p, pyfile)
|
||||
cond = os.path.exists(realpath)
|
||||
|
||||
if not cond:
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
if not cond:
|
||||
realpath = os.path.join(panel_path, pyfile)
|
||||
cond = os.path.exists(realpath)
|
||||
|
||||
@@ -2112,10 +2112,10 @@ class SqliteEasy:
|
||||
|
||||
# 嵌套where
|
||||
@contextmanager
|
||||
def where_nest(self):
|
||||
def where_nest(self, logic: str = 'and'):
|
||||
query = SqliteEasy(self.__DB)
|
||||
yield query
|
||||
self.__OPT_WHERE.add_nest(query.get_where_obj())
|
||||
self.__OPT_WHERE.add_nest(query.get_where_obj(), logic)
|
||||
|
||||
def group(self, condition, params=()):
|
||||
'''
|
||||
|
||||
+60
-7
@@ -13,7 +13,7 @@ import shutil
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timedelta
|
||||
from hashlib import md5
|
||||
from typing import Optional, Tuple, List, Dict
|
||||
|
||||
@@ -87,6 +87,7 @@ class _SSLDatabase:
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
|
||||
ssl_db = _SSLDatabase()
|
||||
|
||||
|
||||
@@ -227,7 +228,8 @@ class SSLManger:
|
||||
def save_by_data(self, certificate: str,
|
||||
private_key: str,
|
||||
cloud_id: Optional[int] = None,
|
||||
other_data: Optional[Dict] = None) -> Dict:
|
||||
other_data: Optional[Dict] = None,
|
||||
log_file: Optional[str] = "") -> Dict:
|
||||
|
||||
if not certificate.startswith("-----BEGIN") or not private_key.startswith("-----BEGIN"):
|
||||
raise ValueError(public.lang("Certificate format error"))
|
||||
@@ -273,8 +275,59 @@ class SSLManger:
|
||||
|
||||
res_id = ssl_db.connection().insert(pdata)
|
||||
public.M('ssl_info').insert(pdata) # add default.db ssl_info table
|
||||
if isinstance(res_id, str) and res_id.startswith("error"):
|
||||
raise ValueError(public.lang("db write error"))
|
||||
# ======= save dns domain db ============
|
||||
try:
|
||||
from ssl_domainModelV2.model import DnsDomainProvider, DnsDomainSSL
|
||||
provider, account, token = auth_info.get("auth_to", "||").split("|")
|
||||
p_obj = DnsDomainProvider.objects.filter(
|
||||
name=provider, api_user=account, api_key=token
|
||||
).first()
|
||||
# upload cert maybe have no provider
|
||||
pid = p_obj.id if p_obj and provider != "" else 0
|
||||
# keep the same ssl cert unique
|
||||
# more detail => self.sub_all_cert(key_file, pem_file)
|
||||
for ssl in DnsDomainSSL.objects.filter(
|
||||
provider_id=pid,
|
||||
dns=info.get("dns", []),
|
||||
subject=info.get("subject", "")
|
||||
):
|
||||
if all([
|
||||
ssl.info.get("issuer") == info.get("issuer", ""),
|
||||
ssl.info.get("issuer_O") == info.get("issuer_O", ""),
|
||||
ssl.dns == info.get("dns", []),
|
||||
]):
|
||||
ssl.delete()
|
||||
try:
|
||||
date_time = datetime.strptime(info.get("notAfter"), "%Y-%m-%d")
|
||||
not_after_ts = int(time.mktime(date_time.timetuple())) * 1000
|
||||
except:
|
||||
not_after_ts = int(
|
||||
time.mktime(
|
||||
time.strptime(f"{datetime.now().date() + timedelta(days=30)}", "%Y-%m-%d")
|
||||
)
|
||||
) * 1000
|
||||
|
||||
DnsDomainSSL(**{
|
||||
"provider_id": pid,
|
||||
"hash": hash_data,
|
||||
"path": "{}/{}".format(SSL_SAVE_PATH, hash_data),
|
||||
"dns": info.get("dns", []),
|
||||
"subject": info.get("subject", ""),
|
||||
"info": info,
|
||||
"cloud_id": int(cloud_id),
|
||||
"not_after": info.get("notAfter", ""),
|
||||
"not_after_ts": not_after_ts,
|
||||
"auth_info": auth_info,
|
||||
"log": log_file,
|
||||
}).save()
|
||||
except Exception as e:
|
||||
import traceback
|
||||
public.print_log(traceback.format_exc())
|
||||
public.print_log("save dns domain db error: {}".format(e))
|
||||
# ======= end dns domain db ===========
|
||||
|
||||
# if isinstance(res_id, str) and res_id.startswith("error"):
|
||||
# raise ValueError(public.lang("db write error"))
|
||||
|
||||
pdata["id"] = res_id
|
||||
if not os.path.exists(pdata["path"]):
|
||||
@@ -401,7 +454,7 @@ class SSLManger:
|
||||
all_ids = ssl_db.connection().field("id").select()
|
||||
for ssl_id in all_ids:
|
||||
if ssl_id["id"] not in change_set:
|
||||
ssl_db.connection().where("id = ?", (ssl_id["id"], )).update({"cloud_id": -1})
|
||||
ssl_db.connection().where("id = ?", (ssl_id["id"],)).update({"cloud_id": -1})
|
||||
|
||||
# 从本地收集证书
|
||||
def _get_ssl_by_local_data(self): # 从本地获取可用证书
|
||||
@@ -472,7 +525,7 @@ class SSLManger:
|
||||
target["dns"] = json.loads(target["dns"])
|
||||
target["info"] = json.loads(target["info"])
|
||||
target['endtime'] = int((datetime.strptime(target['not_after'], "%Y-%m-%d").timestamp()
|
||||
- datetime.today().timestamp()) / (60 * 60 * 24))
|
||||
- datetime.today().timestamp()) / (60 * 60 * 24))
|
||||
return target
|
||||
|
||||
@classmethod
|
||||
@@ -624,7 +677,7 @@ class SSLManger:
|
||||
res_data = json.loads(res_text)
|
||||
if res_data["status"] is True:
|
||||
cloud_id = int(res_data["data"].get("id"))
|
||||
ssl_db.connection().where("id = ?", (target["id"], )).update({"cloud_id": cloud_id})
|
||||
ssl_db.connection().where("id = ?", (target["id"],)).update({"cloud_id": cloud_id})
|
||||
|
||||
return res_data
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user