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:
Jack
2025-03-26 09:04:29 +08:00
parent bfa378592b
commit 6f212ce3d6
1114 changed files with 68134 additions and 19111 deletions
+353 -113
View File
@@ -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
View File
@@ -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')
+1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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):
+10
View File
@@ -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__
+362
View File
@@ -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
+634
View File
@@ -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"
)
+260
View File
@@ -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
View File
@@ -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
+9 -4
View File
@@ -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)
+2 -2
View File
@@ -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
View File
@@ -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: